| import os |
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from PIL import Image, ImageOps |
| import torchvision.transforms as transforms |
| import numpy as np |
| import gradio as gr |
| from collections import OrderedDict |
|
|
| from torchvision.models import efficientnet_v2_s, EfficientNet_V2_S_Weights |
|
|
| |
| |
| |
| MODEL_WEIGHTS = "model_weights.pth" |
| CLASSES_FILE = "classes.txt" |
| NUM_CLASSES = 200 |
| DEVICE = "cuda" if torch.cuda.is_available() else "cpu" |
| DEFAULT_TOP_K = 5 |
|
|
| |
| |
| |
| def load_class_names(path=CLASSES_FILE, num_classes=NUM_CLASSES): |
| if not os.path.exists(path): |
| print(f"Class file '{path}' not found -> using generic names") |
| return [f"class_{i}" for i in range(num_classes)] |
| names = {} |
| with open(path, "r", encoding="utf-8") as f: |
| lines = [ln.strip() for ln in f if ln.strip()] |
| for i, ln in enumerate(lines): |
| parts = ln.split(maxsplit=1) |
| if len(parts) == 2 and parts[0].isdigit(): |
| idx = int(parts[0]) |
| names[idx] = parts[1].strip() |
| else: |
| names[i] = ln |
| max_idx = max(names.keys()) if names else -1 |
| size = max(num_classes, max_idx + 1) |
| class_list = [f"class_{i}" for i in range(size)] |
| for idx, nm in names.items(): |
| if 0 <= idx < size: |
| class_list[idx] = nm |
| if len(class_list) != num_classes: |
| if len(class_list) < num_classes: |
| class_list += [f"class_{i}" for i in range(len(class_list), num_classes)] |
| else: |
| class_list = class_list[:num_classes] |
| print(f"Loaded {len(names)} class names from '{path}' (final list length {len(class_list)})") |
| return class_list |
|
|
| class_names = load_class_names() |
|
|
| |
| |
| |
| def build_model(num_classes=NUM_CLASSES, pretrained=True): |
| if pretrained: |
| weights = EfficientNet_V2_S_Weights.IMAGENET1K_V1 |
| model = efficientnet_v2_s(weights=weights) |
| else: |
| model = efficientnet_v2_s(weights=None) |
|
|
| |
| in_features = model.classifier[-1].in_features |
| model.classifier = nn.Sequential( |
| nn.Dropout(p=0.2, inplace=True), |
| nn.Linear(in_features, 512), |
| nn.BatchNorm1d(512), |
| nn.ReLU(inplace=True), |
| nn.Dropout(0.4, inplace=True), |
| nn.Linear(512, 256), |
| nn.BatchNorm1d(256), |
| nn.ReLU(inplace=True), |
| nn.Dropout(0.3, inplace=True), |
| nn.Linear(256, num_classes), |
| ) |
| return model |
|
|
| def load_weights(model, path=MODEL_WEIGHTS, device=DEVICE): |
| if not os.path.exists(path): |
| print(f"Weight file '{path}' not found. App will run with random weights.") |
| return model |
|
|
| ckpt = torch.load(path, map_location=device) |
| |
| if isinstance(ckpt, dict): |
| if "model_state_dict" in ckpt: |
| state = ckpt["model_state_dict"] |
| elif "state_dict" in ckpt: |
| state = ckpt["state_dict"] |
| else: |
| state = ckpt |
| else: |
| state = ckpt |
|
|
| print(f"Loaded checkpoint object with {len(state)} keys") |
|
|
| for k in list(state.keys())[:25]: |
| v = state[k] |
| try: |
| shape = tuple(v.shape) |
| except Exception: |
| shape = type(v) |
| print(f" ckpt key: {k} shape/type: {shape}") |
|
|
| |
| new_state = {} |
| for k, v in state.items(): |
| nk = k |
| for prefix in ("module.", "model.", "efficientnet."): |
| if nk.startswith(prefix): |
| nk = nk[len(prefix):] |
| new_state[nk] = v |
|
|
| model_state = model.state_dict() |
|
|
| print("Example model keys (first 25):") |
| for k in list(model_state.keys())[:25]: |
| shape = tuple(model_state[k].shape) if hasattr(model_state[k], "shape") else type(model_state[k]) |
| print(f" model key: {k} shape/type: {shape}") |
|
|
| |
| dropped = [] |
| for k in list(new_state.keys()): |
| if k in model_state: |
| if hasattr(new_state[k], "shape") and hasattr(model_state[k], "shape"): |
| if tuple(new_state[k].shape) != tuple(model_state[k].shape): |
| dropped.append((k, tuple(new_state[k].shape), tuple(model_state[k].shape))) |
| new_state.pop(k) |
|
|
| if dropped: |
| print("Dropped checkpoint keys with mismatched shapes (usually final layer):") |
| for k, s_ckpt, s_model in dropped: |
| print(f" {k}: ckpt {s_ckpt} != model {s_model}") |
|
|
| res = model.load_state_dict(new_state, strict=False) |
| print(f"Loaded with strict=False. missing_keys: {len(res.missing_keys)}, unexpected_keys: {len(res.unexpected_keys)}") |
| if res.missing_keys: |
| print(" Examples of missing keys:", res.missing_keys[:10]) |
| if res.unexpected_keys: |
| print(" Examples of unexpected ckpt keys:", res.unexpected_keys[:10]) |
|
|
| return model |
|
|
| model = build_model(len(class_names)) |
| model = load_weights(model, MODEL_WEIGHTS, DEVICE) |
| model.to(DEVICE) |
| model.eval() |
|
|
| |
| |
| |
| weights = EfficientNet_V2_S_Weights.IMAGENET1K_V1 |
| transform = weights.transforms() |
|
|
| |
| |
| |
| def predict(image: Image.Image, top_k: int = DEFAULT_TOP_K): |
| if image is None: |
| return [] |
|
|
| img = transform(image).unsqueeze(0).to(DEVICE) |
| with torch.no_grad(): |
| logits = model(img) |
| probs = torch.softmax(logits, dim=1)[0] |
| top_k = min(top_k, probs.size(0)) |
| values, indices = torch.topk(probs, k=top_k) |
|
|
| |
| print("logits min/max/std:", logits.min().item(), logits.max().item(), logits.std().item()) |
| top_logits, top_idx = torch.topk(logits[0], k=top_k) |
| print("top logits:", [float(x) for x in top_logits], "indices:", [int(i) for i in top_idx]) |
| try: |
| final_linear = None |
| if hasattr(model, "classifier") and isinstance(model.classifier[-1], nn.Linear): |
| final_linear = model.classifier[-1] |
| if final_linear is not None: |
| print("final fc weight norm:", final_linear.weight.norm().item(), |
| "bias norm:", final_linear.bias.norm().item() if final_linear.bias is not None else None) |
| else: |
| print("No final linear layer found in model.classifier[-1]") |
| except Exception as e: |
| print("final layer inspect failed:", e) |
|
|
| rows = [] |
| for p, idx in zip(values.tolist(), indices.tolist()): |
| rows.append([class_names[idx], f"{p*100:.2f}%"]) |
|
|
| return rows |
|
|
| |
| |
| |
| def get_last_conv_layer(m: nn.Module): |
| last = None |
| for mod in m.modules(): |
| if isinstance(mod, nn.Conv2d): |
| last = mod |
| if last is None: |
| raise RuntimeError("No Conv2d layer found for Grad-CAM.") |
| return last |
|
|
| _gc_acts = None |
| _gc_grads = None |
|
|
| def _forward_hook(module, inp, out): |
| global _gc_acts |
| _gc_acts = out.detach() |
|
|
| def _backward_hook(module, grad_in, grad_out): |
| global _gc_grads |
| _gc_grads = grad_out[0].detach() |
|
|
| |
| _target_layer = get_last_conv_layer(model.features if hasattr(model, "features") else model) |
| _target_layer.register_forward_hook(_forward_hook) |
| |
| if hasattr(_target_layer, "register_full_backward_hook"): |
| _target_layer.register_full_backward_hook(_backward_hook) |
| else: |
| _target_layer.register_backward_hook(_backward_hook) |
|
|
| def gradcam(image: Image.Image, class_idx: int = None, alpha: float = 0.5): |
| """ |
| Returns a PIL image with Grad-CAM overlay for the chosen class (top-1 if None). |
| """ |
| if image is None: |
| return None |
|
|
| model.zero_grad() |
| orig_w, orig_h = image.size |
|
|
| x = transform(image).unsqueeze(0).to(DEVICE) |
| x.requires_grad_(True) |
| with torch.enable_grad(): |
| logits = model(x) |
| probs = torch.softmax(logits, dim=1) |
| if class_idx is None: |
| class_idx = int(torch.argmax(probs, dim=1).item()) |
|
|
| one_hot = torch.zeros_like(logits) |
| one_hot[0, class_idx] = 1.0 |
| logits.backward(gradient=one_hot, retain_graph=False) |
|
|
| if _gc_acts is None or _gc_grads is None: |
| print("Grad-CAM hooks did not capture activations/gradients.") |
| return None |
|
|
| |
| |
| weights_gc = _gc_grads.mean(dim=(2, 3), keepdim=True) |
| cam = (weights_gc * _gc_acts).sum(dim=1, keepdim=True) |
| cam = F.relu(cam) |
|
|
| |
| cam_min, cam_max = cam.min(), cam.max() |
| if float(cam_max - cam_min) > 1e-8: |
| cam = (cam - cam_min) / (cam_max - cam_min) |
| else: |
| cam = torch.zeros_like(cam) |
|
|
| |
| cam_up = F.interpolate(cam, size=(orig_h, orig_w), mode="bilinear", align_corners=False) |
| cam_np = cam_up[0, 0].detach().cpu().numpy() |
|
|
| |
| cam_img = Image.fromarray((cam_np * 255).astype("uint8"), mode="L") |
| heatmap = ImageOps.colorize(cam_img, black="black", white="red") |
|
|
| |
| base = image.convert("RGB") |
| overlay = Image.blend(base, heatmap.convert("RGB"), alpha=alpha) |
| return overlay |
|
|
| def predict_with_gradcam(image: Image.Image, top_k: int = DEFAULT_TOP_K, cam_on_top1: bool = True): |
| rows = predict(image, top_k=top_k) |
| if image is None: |
| return rows, None |
| |
| x = transform(image).unsqueeze(0).to(DEVICE) |
| with torch.no_grad(): |
| logits = model(x) |
| top1_idx = int(torch.argmax(logits, dim=1).item()) |
| cam_img = gradcam(image, class_idx=top1_idx if cam_on_top1 else None, alpha=0.5) |
| return rows, cam_img |
|
|
| |
| |
| |
| title = "CUB-200 EfficientNetV2-S Classifier + Grad-CAM" |
| description = ( |
| "Upload a bird image. Model fine-tuned on CUB-200. " |
| "Class names loaded from 'classes.txt'. Includes Grad-CAM visualization." |
| ) |
|
|
| def get_class_rows(): |
| return [[i, class_names[i]] for i in range(len(class_names))] |
|
|
| with gr.Blocks() as demo: |
| gr.Markdown(f"## {title}\n\n{description}") |
|
|
| with gr.Row(): |
| with gr.Column(scale=2): |
| inp_image = gr.Image(type="pil", label="Input Image") |
| topk_slider = gr.Slider(minimum=1, maximum=10, step=1, value=DEFAULT_TOP_K, label="Top-k") |
| with gr.Row(): |
| predict_btn = gr.Button("Predict") |
| predict_cam_btn = gr.Button("Predict + Grad-CAM") |
| with gr.Column(scale=1): |
| out_predictions = gr.Dataframe(headers=["Class", "Confidence"], label="Top-k predictions") |
| cam_image = gr.Image(type="pil", label="Grad-CAM Overlay") |
| gr.Markdown("### Classes (convenience)") |
| show_classes_btn = gr.Button("Show classes") |
| out_classes = gr.Dataframe(headers=["Index", "Class"], value=get_class_rows(), label="All classes") |
|
|
| |
| predict_btn.click(fn=predict, inputs=[inp_image, topk_slider], outputs=out_predictions) |
| predict_cam_btn.click(fn=predict_with_gradcam, inputs=[inp_image, topk_slider], outputs=[out_predictions, cam_image]) |
| show_classes_btn.click(fn=lambda: get_class_rows(), inputs=None, outputs=out_classes) |
|
|
| if __name__ == "__main__": |
| demo.launch(server_name="0.0.0.0", share=False) |
|
|