File size: 4,998 Bytes
1ea7ba6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
"""Side-by-side HiResCAM comparison across the 3 training regimes.



Picks 4 representative images, runs HiResCAM with each of:

  - best_original.pth      (raw baseline)

  - best_sam.pth           (SAM bg-removal)

  - best_strongaug.pth     (strong augmentation)



Builds a 4x4 grid (rows = images, cols = original | 3 CAMs).

"""
import sys, os
sys.path.insert(0, "/mnt/d/SpiceNet" if os.path.exists("/mnt/d/SpiceNet") else "D:/SpiceNet")

import numpy as np
import torch
import cv2
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from PIL import Image
import albumentations as A
from albumentations.pytorch import ToTensorV2

import config
from src.model import load_checkpoint
from src.gradcam import GradCAM
from src.features import extract_all as _extract_all


_RESIZE = int(config.IMG_SIZE * 256 / 224)
_CROP   = config.IMG_SIZE


def model_view(img_path):
    img = np.array(Image.open(img_path).convert("RGB"))
    crop_tf = A.Compose([A.Resize(_RESIZE, _RESIZE), A.CenterCrop(_CROP, _CROP)])
    display = crop_tf(image=img)["image"]
    norm_tf = A.Compose([A.Normalize(mean=config.IMG_MEAN, std=config.IMG_STD), ToTensorV2()])
    tensor = norm_tf(image=display)["image"].unsqueeze(0)
    tex_np, col_np = _extract_all(img)
    tex = torch.from_numpy(tex_np).unsqueeze(0).float()
    col = torch.from_numpy(col_np).unsqueeze(0).float()
    return tensor, display, tex, col


def compute_cam(model, tensor, tex, col, device, smooth_sigma=1.2):
    gcam = GradCAM(model, mode="fusion", algorithm="hirescam",
                   smooth_sigma=smooth_sigma)
    tensor = tensor.to(device); tex = tex.to(device); col = col.to(device)
    cam, pred = gcam(tensor, tex=tex, col=col)
    gcam.remove()
    return cam, pred


def overlay(display, cam, smooth_sigma):
    heat = cv2.resize(cam, (display.shape[1], display.shape[0]),
                      interpolation=cv2.INTER_CUBIC)
    if smooth_sigma > 0:
        k = max(3, int(2 * round(3 * smooth_sigma) + 1))
        heat = cv2.GaussianBlur(heat, (k, k), smooth_sigma)
        heat = heat - heat.min()
        if heat.max() > 0:
            heat = heat / heat.max()
    heat_u8 = cv2.applyColorMap(np.uint8(255 * heat), cv2.COLORMAP_JET)
    heat_u8 = cv2.cvtColor(heat_u8, cv2.COLOR_BGR2RGB)
    return cv2.addWeighted(display, 0.5, heat_u8, 0.5, 0)


def main():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    # 4 representative images: powder, scattered seeds, single piece, painted plate
    samples = [
        ("Spice_Spectrum/turmeric/turmeric_10.jpg",   "turmeric"),
        ("Spice_Spectrum/cumin/cumin_50.jpg",         "cumin"),
        ("Spice_Spectrum/cinnamon/cinnamon_100.jpg",  "cinnamon"),
        ("Spice_Spectrum/saffron/saffron_200.jpg",    "saffron"),
    ]

    pipelines = [
        ("Raw baseline",   "best_original.pth"),
        ("SAM bg-removal", "best_sam.pth"),
        ("Strong aug",     "best_strongaug.pth"),
    ]

    fig, axes = plt.subplots(len(samples), 1 + len(pipelines),
                             figsize=(3 * (1 + len(pipelines)), 3 * len(samples)))

    for row, (path, lbl) in enumerate(samples):
        tensor, display, tex, col = model_view(path)
        axes[row, 0].imshow(display)
        axes[row, 0].set_title("Input" if row == 0 else "", fontsize=10)
        axes[row, 0].set_ylabel(lbl, fontsize=10)
        axes[row, 0].set_xticks([]); axes[row, 0].set_yticks([])

        for col_idx, (pname, ckpt_name) in enumerate(pipelines, start=1):
            ckpt_path = config.CHECKPOINT_DIR / ckpt_name
            if not ckpt_path.exists():
                axes[row, col_idx].text(0.5, 0.5, f"missing: {ckpt_name}",
                                        ha="center", va="center")
                axes[row, col_idx].set_xticks([]); axes[row, col_idx].set_yticks([])
                continue
            model, *_ = load_checkpoint(str(ckpt_path), device)
            cam, pred_idx = compute_cam(model, tensor, tex, col, device)
            ov = overlay(display, cam, smooth_sigma=1.2)
            axes[row, col_idx].imshow(ov)
            pred_name = config.CLASSES[pred_idx]
            color = "green" if pred_name == lbl else "red"
            axes[row, col_idx].set_title(
                f"{pname}\npred: {pred_name}" if row == 0 else f"pred: {pred_name}",
                fontsize=9, color=color,
            )
            axes[row, col_idx].set_xticks([]); axes[row, col_idx].set_yticks([])
            del model
            if device.type == "cuda":
                torch.cuda.empty_cache()

    plt.suptitle("HiResCAM comparison — same images, 3 training regimes", fontsize=13)
    plt.tight_layout()
    out = config.OUTPUT_DIR / "cam_comparison_3regimes.png"
    plt.savefig(out, dpi=150, bbox_inches="tight")
    plt.close()
    print(f"Saved -> {out}")


if __name__ == "__main__":
    main()