{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# CUB ResNet50 Fine-tune + Grad-CAM Localization\n\nFine-tune an ImageNet-pretrained ResNet50 on CUB-200-2011, then evaluate Grad-CAM weakly-supervised localization on the official CUB test split.\n\nTraining uses only image-level class labels. Bounding boxes are used only for localization evaluation.","metadata":{}},{"cell_type":"markdown","source":"## 1. Setup repository and dependencies","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\nimport os\nimport subprocess\nimport sys\n\nON_KAGGLE = Path(\"/kaggle/working\").exists()\n\nif ON_KAGGLE:\n    REPO_DIR = Path(\"/kaggle/working/pytorch-grad-cam\")\n    if not (REPO_DIR / \"pytorch_grad_cam\").exists():\n        subprocess.check_call([\n            \"git\", \"clone\",\n            \"https://github.com/jacobgil/pytorch-grad-cam.git\",\n            str(REPO_DIR),\n        ])\n    os.chdir(REPO_DIR)\n    subprocess.check_call([\n        sys.executable, \"-m\", \"pip\", \"install\", \"-q\", \"-r\", \"requirements.txt\"\n    ])\nelse:\n    REPO_DIR = Path.cwd()\n    print(\"Local mode: assuming current directory is the pytorch-grad-cam repo root.\")\n\nif str(REPO_DIR) not in sys.path:\n    sys.path.insert(0, str(REPO_DIR))\n\nprint(f\"Repository directory: {REPO_DIR}\")\nprint(f\"Current directory: {Path.cwd()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-08T12:21:05.133925Z","iopub.execute_input":"2026-06-08T12:21:05.134196Z","iopub.status.idle":"2026-06-08T12:21:15.803158Z","shell.execute_reply.started":"2026-06-08T12:21:05.134173Z","shell.execute_reply":"2026-06-08T12:21:15.802207Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Imports and configuration\n\nFor a quick Kaggle smoke test, set `EPOCHS = 1`, `N_TRAIN = 200`, `N_VAL = 50`, and `N_TEST = 20`.","metadata":{}},{"cell_type":"code","source":"import csv\nimport json\nimport random\nimport time\nfrom collections import defaultdict\n\nimport cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport torch\nfrom matplotlib.patches import Rectangle\nfrom PIL import Image\nfrom torch import nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import models, transforms\nfrom torchvision.models import ResNet50_Weights\nfrom tqdm.auto import tqdm\n\nfrom pytorch_grad_cam import GradCAM\nfrom pytorch_grad_cam.utils.image import preprocess_image, show_cam_on_image\nfrom pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nUSE_AMP = torch.cuda.is_available()\n\n# Set this manually if your Kaggle dataset folder has a different name.\nDATA_ROOT = \"/kaggle/input/datasets/wenewone/cub2002011/CUB_200_2011\"\n\nNUM_CLASSES = 200\nIMAGE_SIZE = 448\nBATCH_SIZE = 32\nNUM_WORKERS = 2\nEPOCHS = 20\nLR_BACKBONE = 1e-5\nLR_HEAD = 1e-4\nWEIGHT_DECAY = 1e-4\nVAL_FRACTION = 0.10\n\n# Optional smoke-test limits. Use None for full data.\nN_TRAIN = None\nN_VAL = None\nN_TEST = None\n\nTHRESHOLD_RATIO = 0.1\nIOU_THRESHOLD = 0.4\nNUM_VIS_SAMPLES = 6\nSAVE_PREDICTIONS_CSV = True\nSAVE_ALL_VIS_IMAGES = True\n\nCHECKPOINT_PATH = Path(\"/kaggle/working/cub_resnet50_best.pth\") if ON_KAGGLE else REPO_DIR / \"cub_resnet50_best.pth\"\nPREDICTIONS_CSV_PATH = Path(\"/kaggle/working/cub_resnet50_gradcam_test_predictions.csv\") if ON_KAGGLE else REPO_DIR / \"cub_resnet50_gradcam_test_predictions.csv\"\nVIS_OUTPUT_DIR = Path(\"/kaggle/working/cub_gradcam_visualizations\") if ON_KAGGLE else REPO_DIR / \"cub_gradcam_visualizations\"\n\nprint(f\"Using device: {DEVICE}\")\nprint(f\"AMP enabled: {USE_AMP}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-08T12:21:15.805414Z","iopub.execute_input":"2026-06-08T12:21:15.805663Z","iopub.status.idle":"2026-06-08T12:21:26.23818Z","shell.execute_reply.started":"2026-06-08T12:21:15.805639Z","shell.execute_reply":"2026-06-08T12:21:26.2373Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Resolve CUB data root","metadata":{}},{"cell_type":"code","source":"def find_cub_root():\n    if DATA_ROOT is not None:\n        candidate = Path(DATA_ROOT)\n        if not (candidate / \"images.txt\").exists():\n            raise FileNotFoundError(f\"DATA_ROOT does not look like CUB_200_2011: {candidate}\")\n        return candidate\n\n    candidates = [\n        Path(\"/kaggle/input/cub-200-2011/CUB_200_2011\"),\n        Path(\"/kaggle/input/cub2002011/CUB_200_2011\"),\n        Path(\"/kaggle/input/CUB-200-2011/CUB_200_2011\"),\n        Path(\"/kaggle/input/cub-200-2011\"),\n        REPO_DIR / \"CUB-200-2011\" / \"CUB_200_2011\",\n        Path.cwd() / \"CUB-200-2011\" / \"CUB_200_2011\",\n    ]\n    for candidate in candidates:\n        if (candidate / \"images.txt\").exists() and (candidate / \"bounding_boxes.txt\").exists():\n            return candidate\n    raise FileNotFoundError(\"Could not find CUB_200_2011. Set DATA_ROOT manually.\")\n\n\nCUB_ROOT = find_cub_root()\nIMAGE_ROOT = CUB_ROOT / \"images\"\nprint(f\"CUB_ROOT: {CUB_ROOT}\")\nprint(f\"IMAGE_ROOT exists: {IMAGE_ROOT.exists()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-08T12:21:26.239146Z","iopub.execute_input":"2026-06-08T12:21:26.2395Z","iopub.status.idle":"2026-06-08T12:21:26.253825Z","shell.execute_reply.started":"2026-06-08T12:21:26.239477Z","shell.execute_reply":"2026-06-08T12:21:26.252785Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Parse CUB metadata and split train/val/test","metadata":{}},{"cell_type":"code","source":"def read_id_value_file(path, value_parser=str):\n    values = {}\n    with open(path, \"r\", encoding=\"utf-8\") as f:\n        for line in f:\n            line = line.strip()\n            if not line:\n                continue\n            image_id, value = line.split(maxsplit=1)\n            values[int(image_id)] = value_parser(value)\n    return values\n\n\ndef parse_classes(path):\n    class_id_to_name = {}\n    with open(path, \"r\", encoding=\"utf-8\") as f:\n        for line in f:\n            class_id, class_name = line.strip().split(maxsplit=1)\n            class_id_to_name[int(class_id)] = class_name\n    return class_id_to_name\n\n\ndef parse_bboxes(path):\n    boxes = {}\n    with open(path, \"r\", encoding=\"utf-8\") as f:\n        for line in f:\n            parts = line.strip().split()\n            image_id = int(parts[0])\n            x, y, w, h = [float(v) for v in parts[1:]]\n            boxes[image_id] = [x, y, x + w - 1.0, y + h - 1.0]\n    return boxes\n\n\nclass_id_to_name = parse_classes(CUB_ROOT / \"classes.txt\")\nimage_id_to_path = read_id_value_file(CUB_ROOT / \"images.txt\")\nimage_id_to_class_id = read_id_value_file(CUB_ROOT / \"image_class_labels.txt\", int)\nimage_id_to_is_train = read_id_value_file(CUB_ROOT / \"train_test_split.txt\", int)\nimage_id_to_box = parse_bboxes(CUB_ROOT / \"bounding_boxes.txt\")\n\nassert len(image_id_to_path) == 11788\nassert len(class_id_to_name) == 200\nassert len(image_id_to_class_id) == 11788\nassert len(image_id_to_box) == 11788\n\nrecords = []\nfor image_id, rel_path in image_id_to_path.items():\n    class_id = image_id_to_class_id[image_id]\n    records.append({\n        \"image_id\": image_id,\n        \"image_path\": str(IMAGE_ROOT / rel_path),\n        \"relative_path\": rel_path,\n        \"class_id\": class_id,\n        \"label_idx\": class_id - 1,\n        \"class_name\": class_id_to_name[class_id],\n        \"bbox\": image_id_to_box[image_id],\n        \"is_train\": bool(image_id_to_is_train[image_id]),\n    })\n\nofficial_train = [r for r in records if r[\"is_train\"]]\ntest_records = [r for r in records if not r[\"is_train\"]]\n\ndef stratified_train_val_split(train_records, val_fraction=0.10, seed=42):\n    rng = random.Random(seed)\n    by_class = defaultdict(list)\n    for record in train_records:\n        by_class[record[\"class_id\"]].append(record)\n\n    train_split = []\n    val_split = []\n    for class_id in sorted(by_class):\n        class_records = by_class[class_id]\n        rng.shuffle(class_records)\n        val_count = max(1, int(round(len(class_records) * val_fraction)))\n        val_split.extend(class_records[:val_count])\n        train_split.extend(class_records[val_count:])\n\n    rng.shuffle(train_split)\n    rng.shuffle(val_split)\n    return train_split, val_split\n\n\ntrain_records, val_records = stratified_train_val_split(official_train, VAL_FRACTION, SEED)\nif N_TRAIN is not None:\n    train_records = train_records[:N_TRAIN]\nif N_VAL is not None:\n    val_records = val_records[:N_VAL]\nif N_TEST is not None:\n    test_records = test_records[:N_TEST]\n\nprint(f\"All images: {len(records)}\")\nprint(f\"Official train: {len(official_train)}\")\nprint(f\"Train split: {len(train_records)}\")\nprint(f\"Val split: {len(val_records)}\")\nprint(f\"Official test selected: {len(test_records)}\")\nprint(f\"Example: {train_records[0]}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-08T12:21:26.254907Z","iopub.execute_input":"2026-06-08T12:21:26.255532Z","iopub.status.idle":"2026-06-08T12:21:26.471353Z","shell.execute_reply.started":"2026-06-08T12:21:26.255507Z","shell.execute_reply":"2026-06-08T12:21:26.470654Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Dataset and transforms","metadata":{}},{"cell_type":"code","source":"IMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD = [0.229, 0.224, 0.225]\n\ntrain_transform = transforms.Compose([\n    transforms.Resize(int(IMAGE_SIZE * 1.15)),\n    transforms.RandomResizedCrop(IMAGE_SIZE, scale=(0.75, 1.0), ratio=(0.75, 1.3333)),\n    transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n])\n\neval_transform = transforms.Compose([\n    transforms.Resize(int(IMAGE_SIZE * 1.15)),\n    transforms.CenterCrop(IMAGE_SIZE),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n])\n\nclass CUBClassificationDataset(Dataset):\n    def __init__(self, records, transform=None):\n        self.records = records\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.records)\n\n    def __getitem__(self, idx):\n        record = self.records[idx]\n        image = Image.open(record[\"image_path\"]).convert(\"RGB\")\n        if self.transform is not None:\n            image = self.transform(image)\n        label = int(record[\"label_idx\"])\n        return image, label\n\n\ntrain_dataset = CUBClassificationDataset(train_records, train_transform)\nval_dataset = CUBClassificationDataset(val_records, eval_transform)\n\npin_memory = torch.cuda.is_available()\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=NUM_WORKERS,\n    pin_memory=pin_memory,\n)\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=pin_memory,\n)\n\nprint(f\"Train batches: {len(train_loader)}\")\nprint(f\"Val batches: {len(val_loader)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-08T12:21:26.472513Z","iopub.execute_input":"2026-06-08T12:21:26.472948Z","iopub.status.idle":"2026-06-08T12:21:26.483123Z","shell.execute_reply.started":"2026-06-08T12:21:26.472905Z","shell.execute_reply":"2026-06-08T12:21:26.482315Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Build ResNet50 and optimizer","metadata":{}},{"cell_type":"code","source":"weights = ResNet50_Weights.DEFAULT\nmodel = models.resnet50(weights=weights)\nin_features = model.fc.in_features\nmodel.fc = nn.Linear(in_features, NUM_CLASSES)\nmodel = model.to(DEVICE)\n\nbackbone_params = []\nhead_params = []\nfor name, param in model.named_parameters():\n    if name.startswith(\"fc.\"):\n        head_params.append(param)\n    else:\n        backbone_params.append(param)\n\noptimizer = torch.optim.AdamW(\n    [\n        {\"params\": backbone_params, \"lr\": LR_BACKBONE},\n        {\"params\": head_params, \"lr\": LR_HEAD},\n    ],\n    weight_decay=WEIGHT_DECAY,\n)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\ncriterion = nn.CrossEntropyLoss()\n\nAMP_DEVICE_TYPE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nscaler = torch.amp.GradScaler(AMP_DEVICE_TYPE, enabled=USE_AMP)\n\n# Match the ImageNet notebook target layer for Grad-CAM.\ntarget_layers = [model.layer4]\n\nprint(weights)\nprint(model.fc)\nprint(target_layers)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-08T12:21:26.484104Z","iopub.execute_input":"2026-06-08T12:21:26.484608Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Fine-tune","metadata":{}},{"cell_type":"code","source":"def accuracy_topk(logits, targets, topk=(1, 5)):\n    max_k = max(topk)\n    _, pred = logits.topk(max_k, dim=1)\n    pred = pred.t()\n    correct = pred.eq(targets.reshape(1, -1).expand_as(pred))\n    batch_size = targets.size(0)\n    accuracies = []\n    for k in topk:\n        correct_k = correct[:k].reshape(-1).float().sum(0)\n        accuracies.append(float(correct_k.item() / batch_size))\n    return accuracies\n\n\ndef run_train_epoch(model, loader):\n    model.train()\n    running_loss = 0.0\n    running_top1 = 0.0\n    running_top5 = 0.0\n    seen = 0\n\n    for images, labels in tqdm(loader, desc=\"train\", leave=False):\n        images = images.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n        batch_size = labels.size(0)\n\n        optimizer.zero_grad(set_to_none=True)\n        with torch.amp.autocast(device_type=AMP_DEVICE_TYPE, enabled=USE_AMP):\n            logits = model(images)\n            loss = criterion(logits, labels)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        top1, top5 = accuracy_topk(logits.detach(), labels, topk=(1, 5))\n        running_loss += float(loss.item()) * batch_size\n        running_top1 += top1 * batch_size\n        running_top5 += top5 * batch_size\n        seen += batch_size\n\n    return {\n        \"loss\": running_loss / max(seen, 1),\n        \"top1_acc\": running_top1 / max(seen, 1),\n        \"top5_acc\": running_top5 / max(seen, 1),\n    }\n\n\n@torch.no_grad()\ndef run_eval_epoch(model, loader):\n    model.eval()\n    running_loss = 0.0\n    running_top1 = 0.0\n    running_top5 = 0.0\n    seen = 0\n\n    for images, labels in tqdm(loader, desc=\"val\", leave=False):\n        images = images.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n        batch_size = labels.size(0)\n\n        with torch.amp.autocast(device_type=AMP_DEVICE_TYPE, enabled=USE_AMP):\n            logits = model(images)\n            loss = criterion(logits, labels)\n\n        top1, top5 = accuracy_topk(logits, labels, topk=(1, 5))\n        running_loss += float(loss.item()) * batch_size\n        running_top1 += top1 * batch_size\n        running_top5 += top5 * batch_size\n        seen += batch_size\n\n    return {\n        \"loss\": running_loss / max(seen, 1),\n        \"top1_acc\": running_top1 / max(seen, 1),\n        \"top5_acc\": running_top5 / max(seen, 1),\n    }\n\n\nhistory = []\nbest_val_top1 = -1.0\nCHECKPOINT_PATH.parent.mkdir(parents=True, exist_ok=True)\nstart_time = time.time()\n\nfor epoch in range(1, EPOCHS + 1):\n    train_metrics = run_train_epoch(model, train_loader)\n    val_metrics = run_eval_epoch(model, val_loader)\n    scheduler.step()\n\n    row = {\n        \"epoch\": epoch,\n        \"train_loss\": train_metrics[\"loss\"],\n        \"train_top1_acc\": train_metrics[\"top1_acc\"],\n        \"train_top5_acc\": train_metrics[\"top5_acc\"],\n        \"val_loss\": val_metrics[\"loss\"],\n        \"val_top1_acc\": val_metrics[\"top1_acc\"],\n        \"val_top5_acc\": val_metrics[\"top5_acc\"],\n        \"lr_backbone\": optimizer.param_groups[0][\"lr\"],\n        \"lr_head\": optimizer.param_groups[1][\"lr\"],\n    }\n    history.append(row)\n\n    if val_metrics[\"top1_acc\"] > best_val_top1:\n        best_val_top1 = val_metrics[\"top1_acc\"]\n        torch.save({\n            \"model_state_dict\": model.state_dict(),\n            \"epoch\": epoch,\n            \"best_val_top1\": best_val_top1,\n            \"class_id_to_name\": class_id_to_name,\n            \"config\": {\n                \"image_size\": IMAGE_SIZE,\n                \"num_classes\": NUM_CLASSES,\n                \"lr_backbone\": LR_BACKBONE,\n                \"lr_head\": LR_HEAD,\n            },\n        }, CHECKPOINT_PATH)\n\n    print(\n        f\"Epoch {epoch:02d}/{EPOCHS} | \"\n        f\"train loss {train_metrics['loss']:.4f} top1 {train_metrics['top1_acc']:.4f} top5 {train_metrics['top5_acc']:.4f} | \"\n        f\"val loss {val_metrics['loss']:.4f} top1 {val_metrics['top1_acc']:.4f} top5 {val_metrics['top5_acc']:.4f} | \"\n        f\"best val top1 {best_val_top1:.4f}\"\n    )\n\nelapsed = time.time() - start_time\nprint(f\"Training complete in {elapsed / 60:.1f} minutes\")\nprint(f\"Best checkpoint: {CHECKPOINT_PATH}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Load best checkpoint","metadata":{}},{"cell_type":"code","source":"checkpoint = torch.load(CHECKPOINT_PATH, map_location=DEVICE)\nmodel.load_state_dict(checkpoint[\"model_state_dict\"])\nmodel.eval()\nprint(f\"Loaded checkpoint from epoch {checkpoint['epoch']} with val top1 {checkpoint['best_val_top1']:.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 9. Grad-CAM utilities","metadata":{}},{"cell_type":"code","source":"def load_rgb_float(image_path):\n    bgr_img = cv2.imread(str(image_path), cv2.IMREAD_COLOR)\n    if bgr_img is None:\n        raise FileNotFoundError(f\"Could not read image: {image_path}\")\n    rgb_img = cv2.cvtColor(bgr_img, cv2.COLOR_BGR2RGB)\n    rgb_float = np.float32(rgb_img) / 255.0\n    return rgb_img, rgb_float\n\n\ndef make_full_image_input_tensor(rgb_float):\n    return preprocess_image(\n        rgb_float,\n        mean=IMAGENET_MEAN,\n        std=IMAGENET_STD,\n    ).to(DEVICE)\n\n\ndef clamp_box(box, width, height):\n    if box is None:\n        return None\n    x_min, y_min, x_max, y_max = box\n    x_min = max(0.0, min(float(x_min), float(width - 1)))\n    x_max = max(0.0, min(float(x_max), float(width - 1)))\n    y_min = max(0.0, min(float(y_min), float(height - 1)))\n    y_max = max(0.0, min(float(y_max), float(height - 1)))\n    if x_max < x_min or y_max < y_min:\n        return None\n    return [x_min, y_min, x_max, y_max]\n\n\ndef iou_inclusive(box_a, box_b):\n    if box_a is None or box_b is None:\n        return 0.0\n    ax1, ay1, ax2, ay2 = box_a\n    bx1, by1, bx2, by2 = box_b\n    inter_x1 = max(ax1, bx1)\n    inter_y1 = max(ay1, by1)\n    inter_x2 = min(ax2, bx2)\n    inter_y2 = min(ay2, by2)\n    inter_w = max(0.0, inter_x2 - inter_x1 + 1.0)\n    inter_h = max(0.0, inter_y2 - inter_y1 + 1.0)\n    intersection = inter_w * inter_h\n    area_a = max(0.0, ax2 - ax1 + 1.0) * max(0.0, ay2 - ay1 + 1.0)\n    area_b = max(0.0, bx2 - bx1 + 1.0) * max(0.0, by2 - by1 + 1.0)\n    union = area_a + area_b - intersection\n    return float(intersection / union) if union > 0 else 0.0\n\n\ndef heatmap_to_largest_component_box(grayscale_cam, threshold_ratio=0.15):\n    max_value = float(np.max(grayscale_cam))\n    if max_value <= 0:\n        return None\n    mask = (grayscale_cam >= threshold_ratio * max_value).astype(np.uint8)\n    num_labels, labels, stats, centroids = cv2.connectedComponentsWithStats(mask, connectivity=8)\n    if num_labels <= 1:\n        return None\n    areas = stats[1:, cv2.CC_STAT_AREA]\n    largest_label = 1 + int(np.argmax(areas))\n    x = int(stats[largest_label, cv2.CC_STAT_LEFT])\n    y = int(stats[largest_label, cv2.CC_STAT_TOP])\n    w = int(stats[largest_label, cv2.CC_STAT_WIDTH])\n    h = int(stats[largest_label, cv2.CC_STAT_HEIGHT])\n    return [float(x), float(y), float(x + w - 1), float(y + h - 1)]\n\n\ndef cleanup_cam_state(cam):\n    if hasattr(cam, \"activations_and_grads\"):\n        cam.activations_and_grads.clear()\n    if hasattr(cam, \"outputs\"):\n        cam.outputs = None\n    model.zero_grad(set_to_none=True)\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\ndef format_percent(error_count, total):\n    return 100.0 * error_count / total if total else float(\"nan\")\n\n\ndef draw_cv2_box(rgb_uint8, box, color, label):\n    if box is None:\n        return rgb_uint8\n    image = rgb_uint8.copy()\n    x_min, y_min, x_max, y_max = [int(round(v)) for v in box]\n    cv2.rectangle(image, (x_min, y_min), (x_max, y_max), color, 2)\n    text_y = max(15, y_min - 5)\n    cv2.putText(\n        image,\n        label,\n        (x_min, text_y),\n        cv2.FONT_HERSHEY_SIMPLEX,\n        0.45,\n        color,\n        1,\n        cv2.LINE_AA,\n    )\n    return image\n\n\ndef save_gradcam_visualization(vis_payload, output_dir):\n    output_dir.mkdir(parents=True, exist_ok=True)\n    rgb_float = vis_payload[\"rgb_float\"]\n    result = vis_payload[\"result\"]\n    heatmap = vis_payload[\"top_heatmap\"]\n\n    original = np.clip(rgb_float * 255.0, 0, 255).astype(np.uint8)\n    overlay = show_cam_on_image(rgb_float, heatmap, use_rgb=True)\n\n    gt_box = result[\"gt_box\"]\n    pred_box = result[\"top5_boxes\"][0]\n    original = draw_cv2_box(original, gt_box, (0, 255, 0), \"GT\")\n    overlay = draw_cv2_box(overlay, gt_box, (0, 255, 0), \"GT\")\n    overlay = draw_cv2_box(overlay, pred_box, (255, 0, 0), \"Pred\")\n    gt_name = result[\"gt_class_name\"]\n    pred_name = result[\"top5_class_names\"][0]\n\n    combined = np.concatenate([original, overlay], axis=1)\n    header_h = 44\n    canvas = np.full((combined.shape[0] + header_h, combined.shape[1], 3), 255, dtype=np.uint8)\n    canvas[header_h:, :, :] = combined\n    pred_name = result[\"top5_class_names\"][0]\n    title = (\n        f\"id={result['image_id']} \"\n        f\"gt={result['gt_class_id']}:{gt_name} \"\n        f\"pred={result['top5_class_ids'][0]}:{pred_name} \"\n        f\"score={result['top5_scores'][0]:.3f} \"\n        f\"iou={result['top5_ious'][0]:.3f} \"\n        f\"loc={result['top1_loc_correct']}\"\n    )\n    cv2.putText(canvas, title[:180], (8, 28), cv2.FONT_HERSHEY_SIMPLEX, 0.55, (0, 0, 0), 1, cv2.LINE_AA)\n\n    output_path = output_dir / f\"{int(result['image_id']):06d}_{pred_name.replace('/', '_')}.jpg\"\n    cv2.imwrite(str(output_path), cv2.cvtColor(canvas, cv2.COLOR_RGB2BGR), [cv2.IMWRITE_JPEG_QUALITY, 92])\n    return output_path","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10. Evaluate Grad-CAM on CUB test split","metadata":{}},{"cell_type":"code","source":"def evaluate_cub_image(record, cam):\n    rgb_img, rgb_float = load_rgb_float(record[\"image_path\"])\n    height, width = rgb_float.shape[:2]\n    input_tensor = make_full_image_input_tensor(rgb_float)\n    gt_label_idx = int(record[\"label_idx\"])\n    gt_class_id = int(record[\"class_id\"])\n    gt_box = clamp_box(record[\"bbox\"], width, height)\n\n    with torch.no_grad():\n        logits = model(input_tensor)\n        probabilities = torch.softmax(logits, dim=1)\n        top_scores_tensor, top_indices_tensor = torch.topk(probabilities, k=5, dim=1)\n\n    top_indices = [int(v) for v in top_indices_tensor[0].detach().cpu().tolist()]\n    top_scores = [float(v) for v in top_scores_tensor[0].detach().cpu().tolist()]\n    top_class_ids = [idx + 1 for idx in top_indices]\n    top_class_names = [class_id_to_name[class_id] for class_id in top_class_ids]\n\n    top_boxes = []\n    top_ious = []\n    top_heatmap = None\n\n    for class_idx in top_indices:\n        targets = [ClassifierOutputTarget(class_idx)]\n        try:\n            grayscale_cam = cam(input_tensor=input_tensor, targets=targets)[0]\n            if top_heatmap is None:\n                top_heatmap = grayscale_cam.copy()\n            pred_box = heatmap_to_largest_component_box(grayscale_cam, THRESHOLD_RATIO)\n            pred_box = clamp_box(pred_box, width, height)\n            iou = iou_inclusive(pred_box, gt_box)\n        finally:\n            cleanup_cam_state(cam)\n\n        top_boxes.append(pred_box)\n        top_ious.append(float(iou))\n\n    top1_cls_correct = top_indices[0] == gt_label_idx\n    top5_cls_correct = gt_label_idx in top_indices\n    top1_loc_correct = top1_cls_correct and top_ious[0] >= IOU_THRESHOLD\n    top5_loc_correct = any(\n        class_idx == gt_label_idx and iou >= IOU_THRESHOLD\n        for class_idx, iou in zip(top_indices, top_ious)\n    )\n\n    result = {\n        \"image_id\": int(record[\"image_id\"]),\n        \"image_path\": record[\"relative_path\"],\n        \"width\": int(width),\n        \"height\": int(height),\n        \"gt_class_id\": gt_class_id,\n        \"gt_class_name\": record[\"class_name\"],\n        \"gt_box\": gt_box,\n        \"top5_class_ids\": top_class_ids,\n        \"top5_class_names\": top_class_names,\n        \"top5_scores\": top_scores,\n        \"top5_boxes\": top_boxes,\n        \"top5_ious\": top_ious,\n        \"top1_cls_correct\": bool(top1_cls_correct),\n        \"top5_cls_correct\": bool(top5_cls_correct),\n        \"top1_loc_correct\": bool(top1_loc_correct),\n        \"top5_loc_correct\": bool(top5_loc_correct),\n    }\n\n    vis_payload = {\n        \"rgb_float\": rgb_float,\n        \"top_heatmap\": top_heatmap,\n        \"result\": result,\n    }\n    return result, vis_payload\n\n\ntest_results = []\nvisual_samples = []\nstart_time = time.time()\n\nwith GradCAM(model=model, target_layers=target_layers) as cam:\n    for record in tqdm(test_records, desc=\"Grad-CAM test\"):\n        result, vis_payload = evaluate_cub_image(record, cam)\n        test_results.append(result)\n        if SAVE_ALL_VIS_IMAGES:\n            save_gradcam_visualization(vis_payload, VIS_OUTPUT_DIR)\n        if len(visual_samples) < NUM_VIS_SAMPLES:\n            visual_samples.append(vis_payload)\n\nelapsed = time.time() - start_time\ntotal = len(test_results)\ntop1_cls_errors = sum(not r[\"top1_cls_correct\"] for r in test_results)\ntop5_cls_errors = sum(not r[\"top5_cls_correct\"] for r in test_results)\ntop1_loc_errors = sum(not r[\"top1_loc_correct\"] for r in test_results)\ntop5_loc_errors = sum(not r[\"top5_loc_correct\"] for r in test_results)\n\nprint(f\"Evaluated test images: {total}\")\nif SAVE_ALL_VIS_IMAGES:\n    print(f\"Saved Grad-CAM visualizations to: {VIS_OUTPUT_DIR}\")\nprint(f\"Elapsed: {elapsed:.1f}s ({elapsed / max(total, 1):.2f}s/image)\")\nprint()\nprint(\"Metric                              Error %\")\nprint(\"----------------------------------  -------\")\nprint(f\"Top-1 classification error          {format_percent(top1_cls_errors, total):7.2f}\")\nprint(f\"Top-5 classification error          {format_percent(top5_cls_errors, total):7.2f}\")\nprint(f\"Top-1 localization error            {format_percent(top1_loc_errors, total):7.2f}\")\nprint(f\"Top-5 localization error            {format_percent(top5_loc_errors, total):7.2f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 11. Save test predictions","metadata":{}},{"cell_type":"code","source":"if SAVE_PREDICTIONS_CSV:\n    PREDICTIONS_CSV_PATH.parent.mkdir(parents=True, exist_ok=True)\n    fieldnames = [\n        \"image_id\",\n        \"image_path\",\n        \"width\",\n        \"height\",\n        \"gt_class_id\",\n        \"gt_class_name\",\n        \"gt_box\",\n        \"top5_class_ids\",\n        \"top5_class_names\",\n        \"top5_scores\",\n        \"top5_boxes\",\n        \"top5_ious\",\n        \"top1_cls_correct\",\n        \"top5_cls_correct\",\n        \"top1_loc_correct\",\n        \"top5_loc_correct\",\n    ]\n    with open(PREDICTIONS_CSV_PATH, \"w\", encoding=\"utf-8\", newline=\"\") as f:\n        writer = csv.DictWriter(f, fieldnames=fieldnames)\n        writer.writeheader()\n        for result in test_results:\n            row = result.copy()\n            for key in [\"gt_box\", \"top5_class_ids\", \"top5_class_names\", \"top5_scores\", \"top5_boxes\", \"top5_ious\"]:\n                row[key] = json.dumps(row[key])\n            writer.writerow(row)\n    print(f\"Saved predictions to: {PREDICTIONS_CSV_PATH}\")\nelse:\n    print(\"SAVE_PREDICTIONS_CSV is False; no CSV was written.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 12. Visualize examples\n\nGreen boxes are CUB ground truth. Red boxes are top-1 Grad-CAM predicted boxes.","metadata":{}},{"cell_type":"code","source":"def draw_box(ax, box, color, label):\n    if box is None:\n        return\n    x_min, y_min, x_max, y_max = box\n    rect = Rectangle(\n        (x_min, y_min),\n        x_max - x_min + 1,\n        y_max - y_min + 1,\n        linewidth=2,\n        edgecolor=color,\n        facecolor=\"none\",\n    )\n    ax.add_patch(rect)\n    ax.text(\n        x_min,\n        max(0, y_min - 4),\n        label,\n        color=\"white\",\n        fontsize=8,\n        bbox={\"facecolor\": color, \"alpha\": 0.8, \"pad\": 1, \"edgecolor\": \"none\"},\n    )\n\n\ndef show_visual_samples(samples):\n    if not samples:\n        print(\"No visual samples available.\")\n        return\n\n    rows = len(samples)\n    fig, axes = plt.subplots(rows, 2, figsize=(12, 5 * rows))\n    if rows == 1:\n        axes = np.expand_dims(axes, axis=0)\n\n    for row_idx, sample in enumerate(samples):\n        rgb_float = sample[\"rgb_float\"]\n        result = sample[\"result\"]\n        heatmap = sample[\"top_heatmap\"]\n        overlay = show_cam_on_image(rgb_float, heatmap, use_rgb=True)\n\n        top_class_name = result[\"top5_class_names\"][0]\n        top_class_id = result[\"top5_class_ids\"][0]\n        top_score = result[\"top5_scores\"][0]\n        top_box = result[\"top5_boxes\"][0]\n        top_iou = result[\"top5_ious\"][0]\n        gt_box = result[\"gt_box\"]\n\n        axes[row_idx, 0].imshow(rgb_float)\n        axes[row_idx, 0].set_title(\n            f\"GT: {result['gt_class_id']} {result['gt_class_name']}\", fontsize=10\n        )\n        axes[row_idx, 0].axis(\"off\")\n\n        title = (\n            f\"{result['image_id']} | pred: {top_class_id} {top_class_name} | \"\n            f\"score: {top_score:.3f} | IoU: {top_iou:.3f} | loc: {result['top1_loc_correct']}\"\n        )\n        axes[row_idx, 1].imshow(overlay)\n        axes[row_idx, 1].set_title(title, fontsize=10)\n        axes[row_idx, 1].axis(\"off\")\n\n        draw_box(axes[row_idx, 0], gt_box, \"lime\", \"GT\")\n        draw_box(axes[row_idx, 1], gt_box, \"lime\", \"GT\")\n        draw_box(axes[row_idx, 1], top_box, \"red\", \"Pred\")\n\n    plt.tight_layout()\n    plt.show()\n\n\nshow_visual_samples(visual_samples)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}