{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader, Sampler\nfrom torchvision import models, transforms\nfrom PIL import Image\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import train_test_split\nfrom tqdm.auto import tqdm\n\n# ------------------------------------------------------------------\n# CONFIG\n# ------------------------------------------------------------------\nMODEL_NAME = \"convnext_base\"\n\nINPUT_DIRS = [\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-negative-part1\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-negative-part2\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-negative-part3\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-pngs-positive-part-1\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-pngs-part3\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-pngs\",\n]\n\nTRAIN_CSV = (\"/kaggle/input/competitions/rsna-intracranial-hemorrhage-detection/\"\n             \"rsna-intracranial-hemorrhage-detection/stage_2_train.csv\")\n\nMIN_FILE_SIZE_KB = 40\nMIN_FILE_SIZE_BYTES = MIN_FILE_SIZE_KB * 1024\n\nLABEL_COLS = [\"any\", \"epidural\", \"intraparenchymal\",\n              \"intraventricular\", \"subarachnoid\", \"subdural\"]\n\nBATCH_SIZE = 32\nNUM_EPOCHS = 15\nLR = 1e-4\n\nVAL_SPLIT = 0.15\nTEST_SPLIT = 0.15\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nUSE_AMP = DEVICE.type == \"cuda\"\n\nEARLY_STOPPING_PATIENCE = 5\nCHECKPOINT_DIR = \"/kaggle/working/checkpoints\"\nTEST_IDS_PATH = \"/kaggle/working/test_ids.csv\"\n\nNUM_WORKERS = 4\n\n\ndef load_labels(csv_path):\n    y = pd.read_csv(csv_path)\n    id_split = y.ID.str.rsplit(\"_\", n=1, expand=True)\n    y = pd.concat([id_split, y.Label], axis=1)\n    y.columns = [\"id\", \"sub_type\", \"label\"]\n    y = y.drop_duplicates(subset=[\"id\", \"sub_type\"])\n    df = y.pivot(index=\"id\", columns=\"sub_type\", values=\"label\")\n    return df\n\n\nclass ICHDataset(Dataset):\n    def __init__(self, df, id_to_path, transform=None):\n        self.df = df\n        self.id_to_path = id_to_path\n        self.transform = transform\n        self.ids = df.index.values\n        self.labels = df[LABEL_COLS].values.astype(np.float32)\n\n    def __len__(self):\n        return len(self.ids)\n\n    def __getitem__(self, idx):\n        img_id = self.ids[idx]\n        path = self.id_to_path[img_id]\n        img = Image.open(path).convert(\"RGB\")\n        if self.transform:\n            img = self.transform(img)\n        label = torch.tensor(self.labels[idx], dtype=torch.float32)\n        return img, label\n\n\nclass BalancedRandomSampler(Sampler):\n    def __init__(self, labels_any):\n        self.pos_idx = np.where(labels_any == 1)[0]\n        self.neg_idx = np.where(labels_any == 0)[0]\n\n    def __iter__(self):\n        n = min(len(self.neg_idx), max(len(self.pos_idx), 1))\n        neg_sample = np.random.choice(self.neg_idx, n, replace=False)\n        pos_sample = (np.random.choice(self.pos_idx, n, replace=False)\n                      if len(self.pos_idx) > 0 else np.array([], dtype=int))\n        ids = np.concatenate([pos_sample, neg_sample])\n        np.random.shuffle(ids)\n        return iter(ids.tolist())\n\n    def __len__(self):\n        return min(len(self.neg_idx), max(len(self.pos_idx), 1)) * 2\n\n\ndef build_convnext_base(num_classes=len(LABEL_COLS)):\n    model = models.convnext_base(weights=models.ConvNeXt_Base_Weights.IMAGENET1K_V1)\n    for param in model.parameters():\n        param.requires_grad = False\n    # unfreeze only the last stage (features[-1]), same \"last-block-only\" pattern as ViT/ResNet\n    for param in model.features[-1].parameters():\n        param.requires_grad = True\n    in_features = model.classifier[2].in_features\n    model.classifier[2] = nn.Sequential(\n        nn.Dropout(0.3),\n        nn.Linear(in_features, num_classes),\n    )\n    for param in model.classifier.parameters():\n        param.requires_grad = True\n    return model.to(DEVICE)\n\n\nMODEL_BUILDERS = {\"convnext_base\": build_convnext_base}\n\nOPTIMIZER_CONFIG = {\n    \"convnext_base\": {\"type\": \"adamw\", \"lr\": 1e-4, \"weight_decay\": 1e-4},\n}\n\nSCHEDULER_CONFIG = {\n    \"convnext_base\": {\"warmup_epochs\": 1},\n}\n\n\ndef build_optimizer(model, model_name):\n    cfg = OPTIMIZER_CONFIG[model_name]\n    trainable_params = filter(lambda p: p.requires_grad, model.parameters())\n    if cfg[\"type\"] == \"sgd\":\n        return torch.optim.SGD(trainable_params, lr=cfg[\"lr\"], momentum=cfg[\"momentum\"])\n    elif cfg[\"type\"] == \"adamw\":\n        return torch.optim.AdamW(trainable_params, lr=cfg[\"lr\"],\n                                  weight_decay=cfg.get(\"weight_decay\", 0.0))\n    else:\n        raise ValueError(f\"Unknown optimizer type: {cfg['type']}\")\n\n\ndef build_scheduler(optimizer, model_name, num_epochs):\n    warmup_epochs = SCHEDULER_CONFIG[model_name][\"warmup_epochs\"]\n    warmup_epochs = min(warmup_epochs, max(num_epochs - 1, 0))\n\n    if warmup_epochs > 0:\n        warmup = torch.optim.lr_scheduler.LinearLR(\n            optimizer, start_factor=0.1, end_factor=1.0, total_iters=warmup_epochs\n        )\n        cosine = torch.optim.lr_scheduler.CosineAnnealingLR(\n            optimizer, T_max=max(num_epochs - warmup_epochs, 1)\n        )\n        scheduler = torch.optim.lr_scheduler.SequentialLR(\n            optimizer, schedulers=[warmup, cosine], milestones=[warmup_epochs]\n        )\n    else:\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)\n\n    return scheduler\n\n\ndef train_one_epoch(model, loader, optimizer, criterion, scaler, epoch, num_epochs):\n    model.train()\n    total_loss = 0.0\n    pbar = tqdm(loader, desc=f\"train epoch {epoch+1}/{num_epochs}\", leave=False)\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(DEVICE, non_blocking=True), labels.to(DEVICE, non_blocking=True)\n        optimizer.zero_grad()\n\n        with torch.autocast(device_type=DEVICE.type, enabled=USE_AMP):\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        total_loss += loss.item() * imgs.size(0)\n        pbar.set_postfix(loss=f\"{loss.item():.4f}\")\n    return total_loss / len(loader.dataset)\n\n\n@torch.no_grad()\ndef evaluate(model, loader, criterion, epoch, num_epochs):\n    model.eval()\n    total_loss = 0.0\n    all_preds, all_labels = [], []\n    pbar = tqdm(loader, desc=f\"val epoch {epoch+1}/{num_epochs}\", leave=False)\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(DEVICE, non_blocking=True), labels.to(DEVICE, non_blocking=True)\n\n        with torch.autocast(device_type=DEVICE.type, enabled=USE_AMP):\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n\n        total_loss += loss.item() * imgs.size(0)\n        all_preds.append(torch.sigmoid(outputs.float()).cpu().numpy())\n        all_labels.append(labels.cpu().numpy())\n\n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n\n    aucs = {}\n    for i, name in enumerate(LABEL_COLS):\n        if len(np.unique(all_labels[:, i])) > 1:\n            aucs[name] = roc_auc_score(all_labels[:, i], all_preds[:, i])\n        else:\n            aucs[name] = float(\"nan\")\n\n    return total_loss / len(loader.dataset), aucs\n\n\ndef run_training(model_name, train_ds, val_ds, y_train):\n    print(f\"\\n{'='*60}\")\n    print(f\"FULL TRAINING RUN: {model_name}\")\n    print(f\"{'='*60}\")\n\n    model = MODEL_BUILDERS[model_name]()\n    criterion = nn.BCEWithLogitsLoss()\n    optimizer = build_optimizer(model, model_name)\n    scheduler = build_scheduler(optimizer, model_name, NUM_EPOCHS)\n    scaler = torch.amp.GradScaler(device=DEVICE.type, enabled=USE_AMP)\n\n    sampler = BalancedRandomSampler(y_train[\"any\"].values)\n    train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n                               num_workers=NUM_WORKERS, pin_memory=True,\n                               persistent_workers=True, prefetch_factor=4)\n    val_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False,\n                             num_workers=NUM_WORKERS, pin_memory=True,\n                             persistent_workers=True, prefetch_factor=4)\n\n    os.makedirs(CHECKPOINT_DIR, exist_ok=True)\n    checkpoint_path = os.path.join(CHECKPOINT_DIR, f\"{model_name}_best.pt\")\n\n    best_val_loss = float(\"inf\")\n    epochs_no_improve = 0\n\n    for epoch in range(NUM_EPOCHS):\n        current_lr = optimizer.param_groups[0][\"lr\"]\n\n        train_loss = train_one_epoch(model, train_loader, optimizer, criterion,\n                                      scaler, epoch, NUM_EPOCHS)\n        val_loss, val_aucs = evaluate(model, val_loader, criterion, epoch, NUM_EPOCHS)\n\n        auc_str = \", \".join(f\"{k}={v:.3f}\" for k, v in val_aucs.items())\n        print(f\"[{model_name}] Epoch {epoch+1}/{NUM_EPOCHS} | lr={current_lr:.2e} | \"\n              f\"train_loss={train_loss:.4f} | val_loss={val_loss:.4f} | {auc_str}\")\n\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            epochs_no_improve = 0\n            torch.save({\n                \"model_name\": model_name,\n                \"epoch\": epoch + 1,\n                \"model_state_dict\": model.state_dict(),\n                \"optimizer_state_dict\": optimizer.state_dict(),\n                \"val_loss\": val_loss,\n                \"val_aucs\": val_aucs,\n            }, checkpoint_path)\n            print(f\"[{model_name}]   -> val loss improved, checkpoint saved to {checkpoint_path}\")\n        else:\n            epochs_no_improve += 1\n            print(f\"[{model_name}]   -> no improvement ({epochs_no_improve}/{EARLY_STOPPING_PATIENCE})\")\n\n        scheduler.step()\n\n        if epochs_no_improve >= EARLY_STOPPING_PATIENCE:\n            print(f\"[{model_name}] Early stopping triggered at epoch {epoch+1}.\")\n            break\n\n    print(f\"[{model_name}] Training complete. Best val_loss={best_val_loss:.4f}, \"\n          f\"checkpoint at {checkpoint_path}\")\n\n\ndef scan_input_dirs(input_dirs, min_size_bytes):\n    id_to_path = {}\n    skipped_small = 0\n    for d in input_dirs:\n        if not os.path.isdir(d):\n            print(f\"WARNING: directory not found, skipping: {d}\")\n            continue\n        print(f\"listing {d} ...\", flush=True)\n        filenames = [f for f in os.listdir(d) if f.endswith(\".png\")]\n        print(f\"  -> {len(filenames)} png files, scanning sizes...\", flush=True)\n        for f in tqdm(filenames, desc=f\"scanning {os.path.basename(d)}\"):\n            full_path = os.path.join(d, f)\n            if os.path.getsize(full_path) <= min_size_bytes:\n                skipped_small += 1\n                continue\n            img_id = f[:-4]\n            id_to_path[img_id] = full_path\n    print(f\"Scanned {len(input_dirs)} directories: {len(id_to_path)} usable PNGs \"\n          f\"found, {skipped_small} filtered out as low-information (<= \"\n          f\"{min_size_bytes} bytes / {MIN_FILE_SIZE_KB}KB).\")\n    return id_to_path\n\n\ndef make_splits(df):\n    train_df, temp_df = train_test_split(\n        df, test_size=VAL_SPLIT + TEST_SPLIT, stratify=df[\"any\"], random_state=42\n    )\n    val_df, test_df = train_test_split(\n        temp_df, test_size=TEST_SPLIT / (VAL_SPLIT + TEST_SPLIT),\n        stratify=temp_df[\"any\"], random_state=42\n    )\n    return train_df, val_df, test_df\n\n\ndef main():\n    print(f\"*** FULL TRAINING RUN: {MODEL_NAME} ***\")\n    print(f\"Using device: {DEVICE} | AMP enabled: {USE_AMP}\")\n\n    id_to_path = scan_input_dirs(INPUT_DIRS, MIN_FILE_SIZE_BYTES)\n\n    if len(id_to_path) == 0:\n        print(\"\\nNo usable PNGs found across INPUT_DIRS. Check the paths are \"\n              \"correct and the datasets are actually attached as Input.\")\n        return\n\n    df = load_labels(TRAIN_CSV)\n    df = df[df.index.isin(id_to_path.keys())]\n    print(f\"Matched {len(df)} of those to labels. \"\n          f\"Positive rate: {df['any'].mean():.4f}\")\n\n    train_df, val_df, test_df = make_splits(df)\n    print(f\"Full run sizes — train: {len(train_df)}, val: {len(val_df)}, \"\n          f\"test: {len(test_df)}\")\n\n    os.makedirs(os.path.dirname(TEST_IDS_PATH), exist_ok=True)\n    test_df.index.to_series(name=\"id\").to_csv(TEST_IDS_PATH, index=False)\n    print(f\"Saved test set ids to {TEST_IDS_PATH}\")\n\n    train_transform = transforms.Compose([\n        transforms.Resize((224, 224)),\n        transforms.RandomHorizontalFlip(),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                              std=[0.229, 0.224, 0.225]),\n    ])\n    eval_transform = transforms.Compose([\n        transforms.Resize((224, 224)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                              std=[0.229, 0.224, 0.225]),\n    ])\n\n    global train_df_g, val_df_g, test_df_g, id_to_path_g\n    train_df_g, val_df_g, test_df_g, id_to_path_g = train_df, val_df, test_df, id_to_path\n\n    train_ds = ICHDataset(train_df, id_to_path, transform=train_transform)\n    val_ds = ICHDataset(val_df, id_to_path, transform=eval_transform)\n\n    global val_ds_g\n    val_ds_g = val_ds\n\n    run_training(MODEL_NAME, train_ds, val_ds, train_df)\n\n\nmain()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-07-28T13:33:20.088254Z","iopub.execute_input":"2026-07-28T13:33:20.088977Z","iopub.status.idle":"2026-07-28T17:38:43.147858Z","shell.execute_reply.started":"2026-07-28T13:33:20.088947Z","shell.execute_reply":"2026-07-28T17:38:43.141934Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\nfrom torchvision import models, transforms\nfrom sklearn.metrics import (\n    roc_auc_score, average_precision_score, precision_recall_curve,\n    confusion_matrix, f1_score,\n)\n\nMODEL_NAME = \"convnext_base\"\nCHECKPOINT_PATH = f\"/kaggle/working/checkpoints/{MODEL_NAME}_best.pt\"\n\n\ndef build_convnext_base(num_classes=len(LABEL_COLS)):\n    model = models.convnext_base(weights=models.ConvNeXt_Base_Weights.IMAGENET1K_V1)\n    for param in model.parameters():\n        param.requires_grad = False\n    for param in model.features[-1].parameters():\n        param.requires_grad = True\n    in_features = model.classifier[2].in_features\n    model.classifier[2] = nn.Sequential(\n        nn.Dropout(0.3),\n        nn.Linear(in_features, num_classes),\n    )\n    for param in model.classifier.parameters():\n        param.requires_grad = True\n    return model.to(DEVICE)\n\n\nMODEL_BUILDERS = {\"convnext_base\": build_convnext_base}\n\n# Reuse from Cell 1 if present, else rebuild\nif \"val_df_g\" in globals() and \"test_df_g\" in globals() and \"id_to_path_g\" in globals():\n    print(\"Reusing splits from training cell.\")\n    val_df, test_df, id_to_path = val_df_g, test_df_g, id_to_path_g\n    eval_transform = transforms.Compose([\n        transforms.Resize((224, 224)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                              std=[0.229, 0.224, 0.225]),\n    ])\n    val_ds = ICHDataset(val_df, id_to_path, transform=eval_transform)\n    test_ds = ICHDataset(test_df, id_to_path, transform=eval_transform)\nelse:\n    print(\"Splits not found in memory — rebuilding (same random_state=42).\")\n    id_to_path = scan_input_dirs(INPUT_DIRS, MIN_FILE_SIZE_BYTES)\n    df = load_labels(TRAIN_CSV)\n    df = df[df.index.isin(id_to_path.keys())]\n    _, val_df, test_df = make_splits(df)\n    eval_transform = transforms.Compose([\n        transforms.Resize((224, 224)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                              std=[0.229, 0.224, 0.225]),\n    ])\n    val_ds = ICHDataset(val_df, id_to_path, transform=eval_transform)\n    test_ds = ICHDataset(test_df, id_to_path, transform=eval_transform)\n\nval_loader = DataLoader(val_ds, batch_size=32, shuffle=False, num_workers=2, pin_memory=True)\ntest_loader = DataLoader(test_ds, batch_size=32, shuffle=False, num_workers=2, pin_memory=True)\n\nmodel = MODEL_BUILDERS[MODEL_NAME]()\ncheckpoint = torch.load(CHECKPOINT_PATH, map_location=DEVICE, weights_only=False)\nmodel.load_state_dict(checkpoint[\"model_state_dict\"])\nmodel.eval()\nprint(f\"Loaded checkpoint: epoch {checkpoint['epoch']}, val_loss={checkpoint['val_loss']:.4f}\")\n\n\n@torch.no_grad()\ndef run_inference(loader):\n    all_probs, all_labels = [], []\n    for imgs, labels in loader:\n        imgs = imgs.to(DEVICE)\n        with torch.autocast(device_type=DEVICE.type, enabled=USE_AMP):\n            outputs = model(imgs)\n        probs = torch.sigmoid(outputs.float()).cpu().numpy()\n        all_probs.append(probs)\n        all_labels.append(labels.numpy())\n    return np.concatenate(all_probs), np.concatenate(all_labels)\n\n\nprint(\"Running inference on val split (threshold selection only)...\")\nval_probs, val_labels = run_inference(val_loader)\nprint(\"Running inference on test split (reported numbers)...\")\ntest_probs, test_labels = run_inference(test_loader)\nprint(f\"Done: {len(val_probs)} val images, {len(test_probs)} test images.\")\n\n\ndef best_f1_threshold(y_true, y_prob):\n    precisions, recalls, thresholds = precision_recall_curve(y_true, y_prob)\n    f1s = 2 * precisions * recalls / (precisions + recalls + 1e-12)\n    best_idx = np.nanargmax(f1s[:-1])\n    return thresholds[best_idx] if len(thresholds) > 0 else 0.5\n\n\nrows = []\nconfusion_matrices = {}\n\nfor i, name in enumerate(LABEL_COLS):\n    y_true_val = val_labels[:, i]\n    y_true_test = test_labels[:, i]\n    y_prob_test = test_probs[:, i]\n\n    if len(np.unique(y_true_val)) < 2 or len(np.unique(y_true_test)) < 2:\n        continue\n\n    thresh = best_f1_threshold(y_true_val, val_probs[:, i])\n    y_pred_test = (y_prob_test >= thresh).astype(int)\n\n    roc_auc = roc_auc_score(y_true_test, y_prob_test)\n    pr_auc = average_precision_score(y_true_test, y_prob_test)\n\n    tn, fp, fn, tp = confusion_matrix(y_true_test, y_pred_test, labels=[0, 1]).ravel()\n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0.0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0.0\n    specificity = tn / (tn + fp) if (tn + fp) > 0 else 0.0\n    f1 = f1_score(y_true_test, y_pred_test, zero_division=0)\n\n    confusion_matrices[name] = np.array([[tn, fp], [fn, tp]])\n\n    rows.append({\n        \"class\": name, \"roc_auc\": roc_auc, \"pr_auc\": pr_auc,\n        \"threshold_from_val\": thresh, \"precision\": precision,\n        \"recall_sensitivity\": recall, \"specificity\": specificity, \"f1\": f1,\n        \"n_positive\": int(y_true_test.sum()), \"n_total\": len(y_true_test),\n    })\n\nresults_df = pd.DataFrame(rows).set_index(\"class\")\nmacro_auc = roc_auc_score(test_labels, test_probs, average=\"macro\")\nmicro_auc = roc_auc_score(test_labels, test_probs, average=\"micro\")\n\nprint(\"\\n\" + \"=\" * 90)\nprint(f\"PER-CLASS METRICS (test split) — {MODEL_NAME}, checkpoint epoch {checkpoint['epoch']}\")\nprint(\"=\" * 90)\nprint(results_df.round(4).to_string())\nprint(f\"\\nMacro-average ROC-AUC: {macro_auc:.4f}\")\nprint(f\"Micro-average ROC-AUC: {micro_auc:.4f}\")\n\nprint(\"\\n\" + \"=\" * 90)\nprint(\"CONFUSION MATRICES (test split; rows=true, cols=predicted, order=[neg, pos])\")\nprint(\"=\" * 90)\nfor name, cm in confusion_matrices.items():\n    print(f\"\\n{name}:\")\n    print(f\"           pred_neg  pred_pos\")\n    print(f\"true_neg   {cm[0,0]:>8}  {cm[0,1]:>8}\")\n    print(f\"true_pos   {cm[1,0]:>8}  {cm[1,1]:>8}\")\n\nout_path = f\"/kaggle/working/{MODEL_NAME}_eval_metrics.csv\"\nresults_df.to_csv(out_path)\nprint(f\"\\nFull metrics table saved to {out_path}\")\n\nglobal val_probs_g, val_labels_g, test_probs_g, test_labels_g\nval_probs_g, val_labels_g, test_probs_g, test_labels_g = val_probs, val_labels, test_probs, test_labels","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T17:46:03.283028Z","iopub.execute_input":"2026-07-28T17:46:03.283547Z","iopub.status.idle":"2026-07-28T17:54:59.618715Z","shell.execute_reply.started":"2026-07-28T17:46:03.283515Z","shell.execute_reply":"2026-07-28T17:54:59.617928Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom sklearn.metrics import roc_curve, precision_recall_curve, auc\n\nn_classes = len(LABEL_COLS)\nfig, axes = plt.subplots(2, n_classes, figsize=(4 * n_classes, 8))\n\nfor i, name in enumerate(LABEL_COLS):\n    y_true = test_labels[:, i]\n    y_prob = test_probs[:, i]\n\n    if len(np.unique(y_true)) < 2:\n        axes[0, i].set_visible(False)\n        axes[1, i].set_visible(False)\n        continue\n\n    fpr, tpr, _ = roc_curve(y_true, y_prob)\n    roc_auc_val = auc(fpr, tpr)\n    ax_roc = axes[0, i]\n    ax_roc.plot(fpr, tpr, label=f\"AUC = {roc_auc_val:.3f}\")\n    ax_roc.plot([0, 1], [0, 1], \"k--\", alpha=0.4)\n    ax_roc.set_title(f\"{name}\\nROC curve\")\n    ax_roc.set_xlabel(\"False Positive Rate\")\n    ax_roc.set_ylabel(\"True Positive Rate\")\n    ax_roc.legend(fontsize=8, loc=\"lower right\")\n\n    precision, recall, _ = precision_recall_curve(y_true, y_prob)\n    pr_auc_val = average_precision_score(y_true, y_prob)\n    base_rate = y_true.mean()\n    ax_pr = axes[1, i]\n    ax_pr.plot(recall, precision, label=f\"AP = {pr_auc_val:.3f}\")\n    ax_pr.axhline(base_rate, color=\"k\", linestyle=\"--\", alpha=0.4,\n                  label=f\"baseline = {base_rate:.3f}\")\n    ax_pr.set_title(f\"{name}\\nPR curve\")\n    ax_pr.set_xlabel(\"Recall\")\n    ax_pr.set_ylabel(\"Precision\")\n    ax_pr.legend(fontsize=8, loc=\"upper right\")\n\nplt.suptitle(f\"ROC & PR curves — {MODEL_NAME}, test split (n={len(test_labels)})\", y=1.02)\nplt.tight_layout()\nplt.savefig(f\"/kaggle/working/{MODEL_NAME}_roc_pr_curves.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\n\nprint(f\"Saved ROC/PR curve grid to /kaggle/working/{MODEL_NAME}_roc_pr_curves.png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T18:16:58.157718Z","iopub.execute_input":"2026-07-28T18:16:58.158003Z","iopub.status.idle":"2026-07-28T18:17:01.224598Z","shell.execute_reply.started":"2026-07-28T18:16:58.15798Z","shell.execute_reply":"2026-07-28T18:17:01.223711Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import roc_auc_score, average_precision_score, brier_score_loss\n\nN_BINS = 15\n\n\nclass TemperatureScaler(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.log_T = nn.Parameter(torch.zeros(1))\n\n    def forward(self, logits):\n        return logits / torch.exp(self.log_T)\n\n\n@torch.no_grad()\ndef get_logits(loader):\n    all_logits, all_labels = [], []\n    for imgs, labels in loader:\n        imgs = imgs.to(DEVICE)\n        with torch.autocast(device_type=DEVICE.type, enabled=USE_AMP):\n            outputs = model(imgs)\n        all_logits.append(outputs.float().cpu())\n        all_labels.append(labels)\n    return torch.cat(all_logits), torch.cat(all_labels)\n\n\nprint(\"Running inference (logits) on val split (for fitting T)...\")\nval_logits, val_labels_t = get_logits(val_loader)\nprint(\"Running inference (logits) on test split (for reporting)...\")\ntest_logits, test_labels_t = get_logits(test_loader)\n\nscaler = TemperatureScaler()\noptimizer = torch.optim.LBFGS([scaler.log_T], lr=0.05, max_iter=100)\nbce = nn.BCEWithLogitsLoss()\n\ndef closure():\n    optimizer.zero_grad()\n    loss = bce(scaler(val_logits), val_labels_t)\n    loss.backward()\n    return loss\n\noptimizer.step(closure)\nT = torch.exp(scaler.log_T).item()\nprint(f\"\\nFitted temperature: T = {T:.4f}  \"\n      f\"({'softens' if T > 1 else 'sharpens'} the model's confidence)\")\n\n\ndef expected_and_max_calibration_error(y_true, y_prob, n_bins=15):\n    bin_edges = np.linspace(0, 1, n_bins + 1)\n    ece, mce = 0.0, 0.0\n    bin_accs, bin_confs, bin_counts = [], [], []\n    for lo, hi in zip(bin_edges[:-1], bin_edges[1:]):\n        mask = (y_prob > lo) & (y_prob <= hi)\n        count = mask.sum()\n        if count == 0:\n            bin_accs.append(np.nan)\n            bin_confs.append((lo + hi) / 2)\n            bin_counts.append(0)\n            continue\n        acc = y_true[mask].mean()\n        conf = y_prob[mask].mean()\n        gap = abs(acc - conf)\n        ece += (count / len(y_prob)) * gap\n        mce = max(mce, gap)\n        bin_accs.append(acc)\n        bin_confs.append(conf)\n        bin_counts.append(count)\n    return ece, mce, bin_edges, bin_accs, bin_confs, bin_counts\n\n\ndef nll(y_true, y_prob):\n    eps = 1e-12\n    p = np.clip(y_prob, eps, 1 - eps)\n    return -np.mean(y_true * np.log(p) + (1 - y_true) * np.log(1 - p))\n\n\ntest_labels_np = test_labels_t.numpy()\nprobs_before = torch.sigmoid(test_logits).numpy()\nprobs_after = torch.sigmoid(test_logits / T).numpy()\n\nrows = []\nfor i, name in enumerate(LABEL_COLS):\n    y_true = test_labels_np[:, i]\n    if len(np.unique(y_true)) < 2:\n        continue\n\n    p_before, p_after = probs_before[:, i], probs_after[:, i]\n    ece_b, mce_b, *_ = expected_and_max_calibration_error(y_true, p_before, N_BINS)\n    ece_a, mce_a, *_ = expected_and_max_calibration_error(y_true, p_after, N_BINS)\n\n    rows.append({\n        \"class\": name,\n        \"roc_auc\": roc_auc_score(y_true, p_before),\n        \"pr_auc\": average_precision_score(y_true, p_before),\n        \"ece_before\": ece_b, \"ece_after\": ece_a,\n        \"mce_before\": mce_b, \"mce_after\": mce_a,\n        \"brier_before\": brier_score_loss(y_true, p_before),\n        \"brier_after\": brier_score_loss(y_true, p_after),\n        \"nll_before\": nll(y_true, p_before),\n        \"nll_after\": nll(y_true, p_after),\n    })\n\ncalib_results_df = pd.DataFrame(rows).set_index(\"class\")\nprint(\"\\n\" + \"=\" * 100)\nprint(f\"CALIBRATION METRICS (test split, n={len(test_labels_np)}) — {MODEL_NAME}, T={T:.4f}\")\nprint(\"=\" * 100)\nprint(calib_results_df.round(4).to_string())\n\nmacro_ece_before = calib_results_df[\"ece_before\"].mean()\nmacro_ece_after = calib_results_df[\"ece_after\"].mean()\nmacro_brier_before = calib_results_df[\"brier_before\"].mean()\nmacro_brier_after = calib_results_df[\"brier_after\"].mean()\nprint(f\"\\nMacro-avg ECE:   before={macro_ece_before:.4f} -> after={macro_ece_after:.4f}\")\nprint(f\"Macro-avg Brier: before={macro_brier_before:.4f} -> after={macro_brier_after:.4f}\")\n\nany_idx = LABEL_COLS.index(\"any\")\ny_true_any = test_labels_np[:, any_idx]\np_before_any = probs_before[:, any_idx]\np_after_any = probs_after[:, any_idx]\n\nece_b, mce_b, edges_b, accs_b, confs_b, counts_b = expected_and_max_calibration_error(\n    y_true_any, p_before_any, N_BINS)\nece_a, mce_a, edges_a, accs_a, confs_a, counts_a = expected_and_max_calibration_error(\n    y_true_any, p_after_any, N_BINS)\n\nfig, axes = plt.subplots(2, 2, figsize=(11, 9))\nbin_centers = (edges_b[:-1] + edges_b[1:]) / 2\nbin_width = 1 / N_BINS\n\nfor col, (accs, counts, ece, mce, label) in enumerate([\n    (accs_b, counts_b, ece_b, mce_b, \"Before temp scaling (T=1.0)\"),\n    (accs_a, counts_a, ece_a, mce_a, f\"After temp scaling (T={T:.3f})\"),\n]):\n    valid = [not np.isnan(a) for a in accs]\n    ax_rel = axes[0, col]\n    ax_rel.bar(np.array(bin_centers)[valid], np.array(accs)[valid],\n               width=bin_width, edgecolor=\"black\", alpha=0.7, label=\"Model\")\n    ax_rel.plot([0, 1], [0, 1], \"k--\", label=\"Perfect calibration\")\n    ax_rel.set_xlabel(\"Predicted probability\")\n    ax_rel.set_ylabel(\"Observed frequency\")\n    ax_rel.set_title(f\"{label}\\nECE={ece:.4f}  MCE={mce:.4f}\")\n    ax_rel.legend(fontsize=8)\n\n    ax_hist = axes[1, col]\n    ax_hist.bar(bin_centers, counts, width=bin_width, edgecolor=\"black\", alpha=0.7, color=\"gray\")\n    ax_hist.set_xlabel(\"Predicted probability\")\n    ax_hist.set_ylabel(\"Count\")\n    ax_hist.set_title(\"Confidence histogram\")\n\nplt.suptitle(f\"Calibration — {MODEL_NAME}, class 'any' (test, n={len(test_labels_np)})\", y=1.00)\nplt.tight_layout()\nplt.savefig(f\"/kaggle/working/{MODEL_NAME}_calibration_reliability.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\n\nout_path = f\"/kaggle/working/{MODEL_NAME}_calibration_metrics.csv\"\ncalib_results_df.to_csv(out_path)\nwith open(f\"/kaggle/working/{MODEL_NAME}_temperature.txt\", \"w\") as f:\n    f.write(f\"T={T:.6f}\\n\")\nprint(f\"\\nSaved calibration table to {out_path}\")\nprint(f\"Saved fitted T to /kaggle/working/{MODEL_NAME}_temperature.txt\")\nprint(f\"Saved reliability diagram to /kaggle/working/{MODEL_NAME}_calibration_reliability.png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T17:55:57.808986Z","iopub.execute_input":"2026-07-28T17:55:57.8095Z","iopub.status.idle":"2026-07-28T18:02:55.845702Z","shell.execute_reply.started":"2026-07-28T17:55:57.809464Z","shell.execute_reply":"2026-07-28T18:02:55.844923Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install grad-cam -q\n\nimport os\nimport pickle\nimport numpy as np\nimport pandas as pd\nimport torch\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nfrom pytorch_grad_cam import GradCAM\nfrom pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\nfrom pytorch_grad_cam.utils.image import show_cam_on_image\n\n# ------------------------------------------------------------------\n# GRAD-CAM — ConvNeXt-Base. Target layer: model.features[-1][-1], the\n# last block of the final stage. ConvNeXt keeps standard 4D conv\n# feature maps (unlike ViT), so plain GradCAM applies without a\n# reshape_transform.\n# ------------------------------------------------------------------\nRANDOM_SEED = 42\nnp.random.seed(RANDOM_SEED)\n\nN_CAM_SAMPLES = 30\nN_VISUALIZE_PER_CLASS = 2\n\nGRADCAM_OUT_DIR = \"/kaggle/working/gradcam_examples\"\nCAM_CACHE_PATH = \"/kaggle/working/cam_cache.pkl\"\nos.makedirs(GRADCAM_OUT_DIR, exist_ok=True)\n\nCHANNEL_NAMES = [\"brain\", \"subdural\", \"bone\"]\n\ntarget_layers = [model.features[-1][-1]]\ncam = GradCAM(model=model, target_layers=target_layers)\n\ncam_cache = {}\nimg_cache = {}\nbone_corr_records = []\n\n\ndef load_image_for_cam(path):\n    img = Image.open(path).convert(\"RGB\").resize((224, 224))\n    img_np = np.array(img).astype(np.float32) / 255.0\n    input_tensor = eval_transform(img).unsqueeze(0).to(DEVICE)\n    return input_tensor, img_np\n\n\ndef bone_shortcut_correlation(grayscale_cam, bone_channel):\n    cam_flat = grayscale_cam.flatten()\n    bone_flat = bone_channel.flatten()\n    if cam_flat.std() < 1e-8 or bone_flat.std() < 1e-8:\n        return np.nan\n    return float(np.corrcoef(cam_flat, bone_flat)[0, 1])\n\n\ndef compute_gradcam(img_id, class_idx, class_name, id_to_path):\n    path = id_to_path[img_id]\n    input_tensor, img_np = load_image_for_cam(path)\n    targets = [ClassifierOutputTarget(class_idx)]\n    grayscale_cam = cam(input_tensor=input_tensor, targets=targets)[0, :]\n\n    cam_cache.setdefault(img_id, {})[class_name] = grayscale_cam\n    img_cache[img_id] = img_np\n\n    bone_channel = img_np[:, :, 2]\n    corr = bone_shortcut_correlation(grayscale_cam, bone_channel)\n    bone_corr_records.append({\"img_id\": img_id, \"class\": class_name, \"bone_corr\": corr})\n\n    visualization = show_cam_on_image(img_np, grayscale_cam, use_rgb=True)\n    return img_np, visualization, grayscale_cam, corr\n\n\ndef show_and_save_gradcam_panel(img_id, class_idx, class_name, id_to_path, fig_axes, tag=\"\"):\n    img_np, visualization, grayscale_cam, corr = compute_gradcam(\n        img_id, class_idx, class_name, id_to_path\n    )\n    ax_brain, ax_subdural, ax_bone, ax_composite, ax_cam = fig_axes\n\n    ax_brain.imshow(img_np[:, :, 0], cmap=\"gray\"); ax_brain.set_title(\"brain\", fontsize=8); ax_brain.axis(\"off\")\n    ax_subdural.imshow(img_np[:, :, 1], cmap=\"gray\"); ax_subdural.set_title(\"subdural\", fontsize=8); ax_subdural.axis(\"off\")\n    ax_bone.imshow(img_np[:, :, 2], cmap=\"gray\"); ax_bone.set_title(\"bone\", fontsize=8); ax_bone.axis(\"off\")\n    ax_composite.imshow(img_np); ax_composite.set_title(f\"{img_id}\\ncomposite\", fontsize=8); ax_composite.axis(\"off\")\n    ax_cam.imshow(visualization); ax_cam.set_title(f\"{class_name}{tag}\\nbone-corr={corr:.2f}\", fontsize=8); ax_cam.axis(\"off\")\n\n    out_path = os.path.join(GRADCAM_OUT_DIR, f\"{class_name}_{img_id}{tag.replace(' ', '_')}.png\")\n    Image.fromarray(visualization).save(out_path)\n\n\ntest_ids_array = test_df.index.values\nclass_preds = {}\nfor class_name in results_df.index:\n    class_idx = LABEL_COLS.index(class_name)\n    thresh = results_df.loc[class_name, \"threshold_from_val\"]\n    class_preds[class_name] = (test_probs[:, class_idx] >= thresh)\n\nclasses_to_show = [\"any\", \"epidural\", \"intraparenchymal\",\n                    \"intraventricular\", \"subarachnoid\", \"subdural\"]\n\nn_cols = N_VISUALIZE_PER_CLASS * 5\nfig, axes = plt.subplots(len(classes_to_show), n_cols,\n                          figsize=(3 * n_cols / 2, 3.2 * len(classes_to_show)))\n\nsample_summary = []\n\nfor row, class_name in enumerate(classes_to_show):\n    class_idx = LABEL_COLS.index(class_name)\n    y_true = test_labels[:, class_idx]\n    y_pred = class_preds[class_name]\n\n    true_pos_mask = (y_true == 1) & (y_pred == 1)\n    true_pos_ids = test_ids_array[true_pos_mask]\n    true_pos_ids = np.array([i for i in true_pos_ids if i in id_to_path])\n\n    n_available = len(true_pos_ids)\n    n_sample = min(N_CAM_SAMPLES, n_available)\n    sample_summary.append({\"class\": class_name, \"n_available\": n_available, \"n_sampled\": n_sample})\n\n    if n_sample == 0:\n        for col in range(n_cols):\n            axes[row, col].set_visible(False)\n        print(f\"WARNING: no confirmed true positives with resolvable paths for '{class_name}'.\")\n        continue\n\n    sample_ids = np.random.choice(true_pos_ids, size=n_sample, replace=False)\n\n    print(f\"[{class_name}] computing CAMs for {n_sample} images \"\n          f\"({n_available} available)...\")\n    for i, img_id in enumerate(tqdm(sample_ids, desc=class_name, leave=False)):\n        if i < N_VISUALIZE_PER_CLASS:\n            panel_axes = axes[row, i * 5:(i + 1) * 5]\n            show_and_save_gradcam_panel(img_id, class_idx, class_name, id_to_path,\n                                         panel_axes, tag=\" (confirmed TP)\")\n        else:\n            compute_gradcam(img_id, class_idx, class_name, id_to_path)\n\n    for col in range(min(N_VISUALIZE_PER_CLASS, n_sample) * 5, n_cols):\n        axes[row, col].set_visible(False)\n\nplt.suptitle(f\"Grad-CAM — {MODEL_NAME}, sample visualization\\n\"\n             f\"(full cache: up to {N_CAM_SAMPLES}/class, cam_cache size = {len(cam_cache)})\",\n             y=1.001)\nplt.tight_layout()\nplt.savefig(f\"/kaggle/working/{MODEL_NAME}_gradcam_sample_visualization.png\",\n            dpi=200, bbox_inches=\"tight\")\nplt.show()\n\nprint(f\"\\nSample sizes used per class:\")\nprint(pd.DataFrame(sample_summary).to_string(index=False))\n\nwith open(CAM_CACHE_PATH, \"wb\") as f:\n    pickle.dump({\"cam_cache\": cam_cache, \"img_cache\": img_cache}, f)\nprint(f\"\\nSaved cam_cache + img_cache ({len(cam_cache)} images) to {CAM_CACHE_PATH}\")\n\nbone_corr_df = pd.DataFrame(bone_corr_records)\nprint(\"\\n\" + \"=\" * 70)\nprint(\"BONE-CHANNEL SHORTCUT CHECK (full sampled set)\")\nprint(\"=\" * 70)\nprint(bone_corr_df.groupby(\"class\")[\"bone_corr\"].agg([\"mean\", \"std\", \"count\"]).round(3))\n\nbone_corr_out = f\"/kaggle/working/{MODEL_NAME}_bone_shortcut_correlations.csv\"\nbone_corr_df.to_csv(bone_corr_out, index=False)\nprint(f\"\\nSaved per-example bone-correlation records to {bone_corr_out}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T18:17:54.195355Z","iopub.execute_input":"2026-07-28T18:17:54.19617Z","iopub.status.idle":"2026-07-28T18:18:21.469107Z","shell.execute_reply.started":"2026-07-28T18:17:54.196141Z","shell.execute_reply":"2026-07-28T18:18:21.468277Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nfrom pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\n\n# ------------------------------------------------------------------\n# AOPC + MAX-SENSITIVITY — ConvNeXt-Base, Grad-CAM explainability validation\n# ------------------------------------------------------------------\n# Per-channel mean-fill baseline, 16x16 patch grid (14px patches).\n# ConvNeXt-Base has the same 32x total downsampling as ResNet50, so\n# the last stage's feature map is also 7x7 at 224x224 input — this\n# patch resolution is consistent across all three models.\n\nN_SAMPLES_PER_CLASS = 50\nN_AOPC_STEPS = 10\nGRID_SIZE = 16\nN_SENSITIVITY_REPEATS = 10\nNOISE_STD_FRACTION = 0.05\n\nnp.random.seed(RANDOM_SEED)\npatch_size = 224 // GRID_SIZE\n\n\ndef normalize_np_to_tensor(img_np):\n    mean = np.array([0.485, 0.456, 0.406], dtype=np.float32)\n    std = np.array([0.229, 0.224, 0.225], dtype=np.float32)\n    normed = (img_np - mean) / std\n    tensor = torch.from_numpy(normed.transpose(2, 0, 1)).unsqueeze(0).float().to(DEVICE)\n    return tensor\n\n\n@torch.no_grad()\ndef get_class_prob(img_np, class_idx):\n    input_tensor = normalize_np_to_tensor(img_np)\n    with torch.autocast(device_type=DEVICE.type, enabled=USE_AMP):\n        outputs = model(input_tensor)\n    return torch.sigmoid(outputs.float())[0, class_idx].item()\n\n\ndef cam_to_patch_ranking(grayscale_cam):\n    patch_scores = np.zeros((GRID_SIZE, GRID_SIZE))\n    for i in range(GRID_SIZE):\n        for j in range(GRID_SIZE):\n            patch = grayscale_cam[i * patch_size:(i + 1) * patch_size,\n                                   j * patch_size:(j + 1) * patch_size]\n            patch_scores[i, j] = patch.mean()\n    flat_order = np.argsort(-patch_scores.flatten())\n    return flat_order\n\n\ndef perturb_patches(img_np, patch_indices_to_remove):\n    perturbed = img_np.copy()\n    channel_means = img_np.reshape(-1, 3).mean(axis=0)\n    for flat_idx in patch_indices_to_remove:\n        i, j = divmod(flat_idx, GRID_SIZE)\n        perturbed[i * patch_size:(i + 1) * patch_size,\n                  j * patch_size:(j + 1) * patch_size, :] = channel_means\n    return perturbed\n\n\ndef compute_aopc(img_np, grayscale_cam, class_idx):\n    ranking = cam_to_patch_ranking(grayscale_cam)\n    total_patches = GRID_SIZE * GRID_SIZE\n    original_prob = get_class_prob(img_np, class_idx)\n\n    drops = []\n    for step in range(1, N_AOPC_STEPS + 1):\n        n_remove = int(total_patches * step / N_AOPC_STEPS)\n        perturbed_img = perturb_patches(img_np, ranking[:n_remove])\n        perturbed_prob = get_class_prob(perturbed_img, class_idx)\n        drops.append(original_prob - perturbed_prob)\n\n    return float(np.mean(drops))\n\n\ndef compute_max_sensitivity(img_id, class_idx, class_name, id_to_path):\n    path = id_to_path[img_id]\n    orig_tensor, orig_img_np = load_image_for_cam(path)\n    targets = [ClassifierOutputTarget(class_idx)]\n\n    orig_cam = cam(input_tensor=orig_tensor, targets=targets)[0, :]\n    noise_std = NOISE_STD_FRACTION * orig_img_np.std()\n\n    max_diff = 0.0\n    for _ in range(N_SENSITIVITY_REPEATS):\n        noise = np.random.normal(0, noise_std, orig_img_np.shape).astype(np.float32)\n        noisy_img = np.clip(orig_img_np + noise, 0, 1)\n        noisy_tensor = normalize_np_to_tensor(noisy_img)\n        noisy_cam = cam(input_tensor=noisy_tensor, targets=targets)[0, :]\n        diff = np.linalg.norm(orig_cam - noisy_cam)\n        max_diff = max(max_diff, diff)\n\n    return float(max_diff)\n\n\ntest_ids_array = test_df.index.values\nrecords = []\n\nfor class_name in results_df.index:\n    class_idx = LABEL_COLS.index(class_name)\n    y_true = test_labels[:, class_idx]\n    y_pred = class_preds[class_name]\n\n    true_pos_mask = (y_true == 1) & (y_pred == 1)\n    true_pos_ids = test_ids_array[true_pos_mask]\n    true_pos_ids = np.array([i for i in true_pos_ids if i in id_to_path])\n\n    n_sample = min(N_SAMPLES_PER_CLASS, len(true_pos_ids))\n    if n_sample == 0:\n        print(f\"Skipping '{class_name}' — no confirmed true positives with resolvable paths.\")\n        continue\n\n    sample_ids = np.random.choice(true_pos_ids, size=n_sample, replace=False)\n    print(f\"\\n{class_name}: running AOPC + Max-Sensitivity on {n_sample} confirmed TPs \"\n          f\"(available: {len(true_pos_ids)})\")\n\n    for img_id in tqdm(sample_ids, desc=class_name):\n        path = id_to_path[img_id]\n        _, img_np = load_image_for_cam(path)\n\n        if img_id in cam_cache and class_name in cam_cache[img_id]:\n            grayscale_cam = cam_cache[img_id][class_name]\n        else:\n            input_tensor, _ = load_image_for_cam(path)\n            targets = [ClassifierOutputTarget(class_idx)]\n            grayscale_cam = cam(input_tensor=input_tensor, targets=targets)[0, :]\n            cam_cache.setdefault(img_id, {})[class_name] = grayscale_cam\n\n        aopc = compute_aopc(img_np, grayscale_cam, class_idx)\n        max_sens = compute_max_sensitivity(img_id, class_idx, class_name, id_to_path)\n\n        records.append({\n            \"class\": class_name, \"img_id\": img_id,\n            \"aopc\": aopc, \"max_sensitivity\": max_sens,\n        })\n\nexplain_df = pd.DataFrame(records)\nprint(\"\\n\" + \"=\" * 90)\nprint(f\"AOPC + MAX-SENSITIVITY SUMMARY — {MODEL_NAME}\")\nprint(\"=\" * 90)\nprint(explain_df.groupby(\"class\")[[\"aopc\", \"max_sensitivity\"]].agg([\"mean\", \"std\", \"count\"]).round(4))\n\nout_path = f\"/kaggle/working/{MODEL_NAME}_aopc_max_sensitivity.csv\"\nexplain_df.to_csv(out_path, index=False)\nprint(f\"\\nSaved per-example AOPC/Max-Sensitivity records to {out_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T18:19:20.573555Z","iopub.execute_input":"2026-07-28T18:19:20.574222Z","iopub.status.idle":"2026-07-28T18:21:44.478344Z","shell.execute_reply.started":"2026-07-28T18:19:20.57419Z","shell.execute_reply":"2026-07-28T18:21:44.477375Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfor f in sorted(os.listdir(\"/kaggle/working\")):\n    print(f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T18:22:08.826966Z","iopub.execute_input":"2026-07-28T18:22:08.827622Z","iopub.status.idle":"2026-07-28T18:22:08.832991Z","shell.execute_reply.started":"2026-07-28T18:22:08.82759Z","shell.execute_reply":"2026-07-28T18:22:08.832029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\nimport os\n\n# Copy just what you want into a clean folder first, then zip that\nos.makedirs(\"/kaggle/working/export\", exist_ok=True)\nfiles_to_zip = [\n    \"cam_cache.pkl\",\n    \"convnext_base_aopc_max_sensitivity.csv\",\n    \"convnext_base_bone_shortcut_correlations.csv\",\n    \"convnext_base_calibration_metrics.csv\",\n    \"convnext_base_calibration_reliability.png\",\n    \"convnext_base_eval_metrics.csv\",\n    \"convnext_base_gradcam_sample_visualization.png\",\n    \"convnext_base_roc_pr_curves.png\",\n    \"convnext_base_temperature.txt\",\n    \"test_ids.csv\",\n]\nfor f in files_to_zip:\n    shutil.copy(f\"/kaggle/working/{f}\", f\"/kaggle/working/export/{f}\")\n\nshutil.copytree(\"/kaggle/working/gradcam_examples\", \"/kaggle/working/export/gradcam_examples\")\n\nshutil.make_archive(\"/kaggle/working/convnext_base_all_outputs\", \"zip\", \"/kaggle/working/export\")\nprint(\"Zipped to /kaggle/working/convnext_base_all_outputs.zip\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-09T04:12:41.625108Z","iopub.execute_input":"2026-09-09T04:12:41.625945Z","iopub.status.idle":"2026-09-09T04:12:41.872933Z","shell.execute_reply.started":"2026-09-09T04:12:41.62591Z","shell.execute_reply":"2026-09-09T04:12:41.871775Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install grad-cam -q\n\nimport os\nimport pickle\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models, transforms\nfrom PIL import Image\nfrom tqdm.auto import tqdm\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import precision_recall_curve\n\nfrom pytorch_grad_cam import GradCAM\nfrom pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\n\nprint(\"Imports complete.\")\nprint(\"Device:\", \"CUDA\" if torch.cuda.is_available() else \"CPU\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-09T04:12:46.844629Z","iopub.execute_input":"2026-09-09T04:12:46.845453Z","iopub.status.idle":"2026-09-09T04:12:58.2962Z","shell.execute_reply.started":"2026-09-09T04:12:46.845422Z","shell.execute_reply":"2026-09-09T04:12:58.295404Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MODEL_NAME = \"convnext_base\"\n\nINPUT_DIRS = [\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-negative-part1\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-negative-part2\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-negative-part3\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-pngs-positive-part-1\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-pngs-part3\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-pngs\",\n]\n\nTRAIN_CSV = (\n    \"/kaggle/input/competitions/\"\n    \"rsna-intracranial-hemorrhage-detection/\"\n    \"rsna-intracranial-hemorrhage-detection/\"\n    \"stage_2_train.csv\"\n)\n\nMIN_FILE_SIZE_KB = 40\nMIN_FILE_SIZE_BYTES = MIN_FILE_SIZE_KB * 1024\n\nLABEL_COLS = [\n    \"any\",\n    \"epidural\",\n    \"intraparenchymal\",\n    \"intraventricular\",\n    \"subarachnoid\",\n    \"subdural\",\n]\n\nBATCH_SIZE = 32\nVAL_SPLIT = 0.15\nTEST_SPLIT = 0.15\n\nDEVICE = torch.device(\n    \"cuda\" if torch.cuda.is_available() else \"cpu\"\n)\n\nUSE_AMP = DEVICE.type == \"cuda\"\nNUM_WORKERS = 2\n\nCHECKPOINT_PATH = (\n    \"/kaggle/working/checkpoints/\"\n    f\"{MODEL_NAME}_best.pt\"\n)\n\n# ---- FAST SANITY-CHECK SETTINGS ----\nN_AOPC_STEPS = 5\nGRID_SIZE = 16\nN_SENSITIVITY_REPEATS = 3\nNOISE_STD_FRACTION = 0.05\n\n# Only evaluate this many images from each TP/FP/TN/FN group initially\nMAX_IMAGES_PER_GROUP = 10\n\nprint(\"Configuration loaded.\")\nprint(\"Checkpoint:\", CHECKPOINT_PATH)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-09T04:13:12.423302Z","iopub.execute_input":"2026-09-09T04:13:12.424399Z","iopub.status.idle":"2026-09-09T04:13:12.431601Z","shell.execute_reply.started":"2026-09-09T04:13:12.42436Z","shell.execute_reply":"2026-09-09T04:13:12.430708Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ICHDataset(Dataset):\n\n    def __init__(self, df, id_to_path, transform=None):\n        self.df = df\n        self.id_to_path = id_to_path\n        self.transform = transform\n\n        self.ids = df.index.values\n\n        self.labels = (\n            df[LABEL_COLS]\n            .values\n            .astype(np.float32)\n        )\n\n    def __len__(self):\n        return len(self.ids)\n\n    def __getitem__(self, idx):\n\n        img_id = self.ids[idx]\n        path = self.id_to_path[img_id]\n\n        img = (\n            Image.open(path)\n            .convert(\"RGB\")\n        )\n\n        if self.transform:\n            img = self.transform(img)\n\n        label = torch.tensor(\n            self.labels[idx],\n            dtype=torch.float32\n        )\n\n        return img, label\n\n\ndef load_labels(csv_path):\n\n    y = pd.read_csv(csv_path)\n\n    id_split = y.ID.str.rsplit(\n        \"_\",\n        n=1,\n        expand=True\n    )\n\n    y = pd.concat(\n        [\n            id_split,\n            y.Label\n        ],\n        axis=1\n    )\n\n    y.columns = [\n        \"id\",\n        \"sub_type\",\n        \"label\"\n    ]\n\n    y = y.drop_duplicates(\n        subset=[\"id\", \"sub_type\"]\n    )\n\n    df = y.pivot(\n        index=\"id\",\n        columns=\"sub_type\",\n        values=\"label\"\n    )\n\n    return df\n\n\nprint(\"Dataset functions ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-09T04:13:17.934717Z","iopub.execute_input":"2026-09-09T04:13:17.935076Z","iopub.status.idle":"2026-09-09T04:13:17.943446Z","shell.execute_reply.started":"2026-09-09T04:13:17.935049Z","shell.execute_reply":"2026-09-09T04:13:17.942385Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def scan_input_dirs(\n    input_dirs,\n    min_file_size_bytes\n):\n\n    id_to_path = {}\n\n    print(\"Scanning image directories...\\n\")\n\n    for root_dir in input_dirs:\n\n        print(\"Scanning:\", root_dir)\n\n        if not os.path.exists(root_dir):\n            print(\"  WARNING: directory not found\")\n            continue\n\n        before = len(id_to_path)\n\n        for root, dirs, files in os.walk(root_dir):\n\n            for filename in files:\n\n                if not filename.lower().endswith(\n                    (\".png\", \".jpg\", \".jpeg\")\n                ):\n                    continue\n\n                path = os.path.join(\n                    root,\n                    filename\n                )\n\n                try:\n                    if os.path.getsize(path) < min_file_size_bytes:\n                        continue\n                except OSError:\n                    continue\n\n                img_id = os.path.splitext(filename)[0]\n\n                id_to_path[img_id] = path\n\n        print(\n            f\"  Added {len(id_to_path) - before:,} images\"\n        )\n\n    print(\"\\nTOTAL IMAGES FOUND:\", f\"{len(id_to_path):,}\")\n\n    return id_to_path\n\n\nid_to_path = scan_input_dirs(\n    INPUT_DIRS,\n    MIN_FILE_SIZE_BYTES\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-09T04:13:23.645941Z","iopub.execute_input":"2026-09-09T04:13:23.646608Z","iopub.status.idle":"2026-09-09T04:26:35.600303Z","shell.execute_reply.started":"2026-09-09T04:13:23.646572Z","shell.execute_reply":"2026-09-09T04:26:35.599347Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = load_labels(TRAIN_CSV)\n\nprint(\"Total labelled IDs:\", f\"{len(df):,}\")\n\ndf = df[\n    df.index.isin(id_to_path.keys())\n]\n\nprint(\n    \"Labelled images with valid paths:\",\n    f\"{len(df):,}\"\n)\n\n\ndef make_splits(df):\n\n    train_val_df, test_df = train_test_split(\n        df,\n        test_size=TEST_SPLIT,\n        random_state=42,\n        stratify=df[\"any\"]\n    )\n\n    train_df, val_df = train_test_split(\n        train_val_df,\n        test_size=VAL_SPLIT / (1.0 - TEST_SPLIT),\n        random_state=42,\n        stratify=train_val_df[\"any\"]\n    )\n\n    return train_df, val_df, test_df\n\n\ntrain_df, val_df, test_df = make_splits(df)\n\nprint(\"\\nSPLIT:\")\nprint(\"Train:\", f\"{len(train_df):,}\")\nprint(\"Validation:\", f\"{len(val_df):,}\")\nprint(\"Test:\", f\"{len(test_df):,}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-09T04:28:46.572201Z","iopub.execute_input":"2026-09-09T04:28:46.572527Z","iopub.status.idle":"2026-09-09T04:29:06.67864Z","shell.execute_reply.started":"2026-09-09T04:28:46.572503Z","shell.execute_reply":"2026-09-09T04:29:06.677818Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"eval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    ),\n])\n\n\nval_ds = ICHDataset(\n    val_df,\n    id_to_path,\n    transform=eval_transform\n)\n\ntest_ds = ICHDataset(\n    test_df,\n    id_to_path,\n    transform=eval_transform\n)\n\n\nval_loader = DataLoader(\n    val_ds,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=True\n)\n\ntest_loader = DataLoader(\n    test_ds,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=True\n)\n\nprint(\"DataLoaders ready.\")\nprint(\"Validation batches:\", len(val_loader))\nprint(\"Test batches:\", len(test_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-09T04:35:26.383911Z","iopub.execute_input":"2026-09-09T04:35:26.384723Z","iopub.status.idle":"2026-09-09T04:35:26.396082Z","shell.execute_reply.started":"2026-09-09T04:35:26.384689Z","shell.execute_reply":"2026-09-09T04:35:26.395032Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_convnext_base(\n    num_classes=len(LABEL_COLS)\n):\n\n    model = models.convnext_base(\n        weights=models.ConvNeXt_Base_Weights.IMAGENET1K_V1\n    )\n\n    # Freeze everything\n    for param in model.parameters():\n        param.requires_grad = False\n\n    # Unfreeze final stage\n    for param in model.features[-1].parameters():\n        param.requires_grad = True\n\n    # Replace classifier\n    in_features = model.classifier[2].in_features\n\n    model.classifier[2] = nn.Sequential(\n        nn.Dropout(0.3),\n        nn.Linear(\n            in_features,\n            num_classes\n        )\n    )\n\n    # Trainable classifier\n    for param in model.classifier.parameters():\n        param.requires_grad = True\n\n    return model.to(DEVICE)\n\n\nif not os.path.exists(CHECKPOINT_PATH):\n    raise FileNotFoundError(\n        f\"Checkpoint not found:\\n{CHECKPOINT_PATH}\"\n    )\n\n\nprint(\"Building ConvNeXt...\")\nmodel = build_convnext_base()\n\nprint(\"Loading checkpoint...\")\ncheckpoint = torch.load(\n    CHECKPOINT_PATH,\n    map_location=DEVICE,\n    weights_only=False\n)\n\nmodel.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\nmodel.eval()\n\nprint(\"\\nMODEL LOADED\")\nprint(\"Epoch:\", checkpoint[\"epoch\"])\nprint(\"Val loss:\", f\"{checkpoint['val_loss']:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-09T04:35:29.709534Z","iopub.execute_input":"2026-09-09T04:35:29.709977Z","iopub.status.idle":"2026-09-09T04:35:33.880889Z","shell.execute_reply.started":"2026-09-09T04:35:29.709948Z","shell.execute_reply":"2026-09-09T04:35:33.880037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef run_inference(loader, description):\n\n    all_probs = []\n    all_labels = []\n\n    for imgs, labels in tqdm(\n        loader,\n        desc=description\n    ):\n\n        imgs = imgs.to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        with torch.autocast(\n            device_type=DEVICE.type,\n            enabled=USE_AMP\n        ):\n            outputs = model(imgs)\n\n        probs = torch.sigmoid(\n            outputs.float()\n        ).cpu().numpy()\n\n        all_probs.append(probs)\n        all_labels.append(\n            labels.numpy()\n        )\n\n    return (\n        np.concatenate(all_probs),\n        np.concatenate(all_labels)\n    )\n\n\nprint(\"Running validation inference...\")\n\nval_probs, val_labels = run_inference(\n    val_loader,\n    \"Validation inference\"\n)\n\nprint(\n    \"\\nValidation complete:\",\n    len(val_probs)\n)\n\n\nprint(\"\\nRunning test inference...\")\n\ntest_probs, test_labels = run_inference(\n    test_loader,\n    \"Test inference\"\n)\n\nprint(\n    \"\\nTest complete:\",\n    len(test_probs)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-09T04:37:35.680892Z","iopub.execute_input":"2026-09-09T04:37:35.681873Z","iopub.status.idle":"2026-09-09T04:40:22.721168Z","shell.execute_reply.started":"2026-09-09T04:37:35.68184Z","shell.execute_reply":"2026-09-09T04:40:22.720182Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def best_f1_threshold(y_true, y_prob):\n\n    precisions, recalls, thresholds = (\n        precision_recall_curve(\n            y_true,\n            y_prob\n        )\n    )\n\n    f1s = (\n        2 * precisions * recalls\n        /\n        (\n            precisions\n            + recalls\n            + 1e-12\n        )\n    )\n\n    best_idx = np.nanargmax(\n        f1s[:-1]\n    )\n\n    return (\n        thresholds[best_idx]\n        if len(thresholds) > 0\n        else 0.5\n    )\n\n\nrows = []\n\nfor i, class_name in enumerate(LABEL_COLS):\n\n    threshold = best_f1_threshold(\n        val_labels[:, i],\n        val_probs[:, i]\n    )\n\n    rows.append({\n        \"class\": class_name,\n        \"threshold_from_val\": threshold\n    })\n\n\nresults_df = (\n    pd.DataFrame(rows)\n    .set_index(\"class\")\n)\n\nprint(\"Validation-derived thresholds:\\n\")\nprint(results_df.round(4))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_preds = {}\n\nfor class_name in LABEL_COLS:\n\n    class_idx = LABEL_COLS.index(\n        class_name\n    )\n\n    threshold = results_df.loc[\n        class_name,\n        \"threshold_from_val\"\n    ]\n\n    class_preds[class_name] = (\n        test_probs[:, class_idx]\n        >= threshold\n    ).astype(int)\n\n\ntest_ids_array = test_df.index.values\n\nassert len(test_ids_array) == len(\n    test_labels\n)\n\n\nprint(\"=\" * 70)\nprint(\"TEST CONFUSION GROUPS\")\nprint(\"=\" * 70)\n\n\nall_group_ids = {}\n\n\nfor class_name in LABEL_COLS:\n\n    class_idx = LABEL_COLS.index(\n        class_name\n    )\n\n    y_true = (\n        test_labels[:, class_idx]\n        .astype(int)\n    )\n\n    y_pred = (\n        class_preds[class_name]\n        .astype(int)\n    )\n\n    group_masks = {\n\n        \"TP\": (\n            (y_true == 1)\n            & (y_pred == 1)\n        ),\n\n        \"FP\": (\n            (y_true == 0)\n            & (y_pred == 1)\n        ),\n\n        \"TN\": (\n            (y_true == 0)\n            & (y_pred == 0)\n        ),\n\n        \"FN\": (\n            (y_true == 1)\n            & (y_pred == 0)\n        ),\n    }\n\n    all_group_ids[class_name] = {}\n\n    print(f\"\\n{class_name}\")\n\n    for group_name, mask in group_masks.items():\n\n        ids = test_ids_array[mask]\n\n        ids = np.array([\n            x for x in ids\n            if x in id_to_path\n        ])\n\n        # SMALL SANITY CHECK\n        ids = ids[:MAX_IMAGES_PER_GROUP]\n\n        all_group_ids[class_name][group_name] = ids\n\n        print(\n            f\"  {group_name}: {len(ids)} images\"\n        )","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"target_layers = [\n    model.features[-1][-1]\n]\n\ncam = GradCAM(\n    model=model,\n    target_layers=target_layers\n)\n\nprint(\"Grad-CAM ready.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_image_for_cam(img_id):\n\n    path = id_to_path[img_id]\n\n    img = (\n        Image.open(path)\n        .convert(\"RGB\")\n        .resize((224, 224))\n    )\n\n    img_np = (\n        np.asarray(img)\n        .astype(np.float32)\n        / 255.0\n    )\n\n    input_tensor = (\n        eval_transform(img)\n        .unsqueeze(0)\n        .to(DEVICE)\n    )\n\n    return input_tensor, img_np\n\n\ndef normalize_np_to_tensor(img_np):\n\n    mean = np.array(\n        [0.485, 0.456, 0.406],\n        dtype=np.float32\n    )\n\n    std = np.array(\n        [0.229, 0.224, 0.225],\n        dtype=np.float32\n    )\n\n    normed = (img_np - mean) / std\n\n    tensor = (\n        torch.from_numpy(\n            normed.transpose(2, 0, 1)\n        )\n        .unsqueeze(0)\n        .float()\n        .to(DEVICE)\n    )\n\n    return tensor\n\n\n@torch.no_grad()\ndef get_class_prob(img_np, class_idx):\n\n    input_tensor = normalize_np_to_tensor(\n        img_np\n    )\n\n    with torch.autocast(\n        device_type=DEVICE.type,\n        enabled=USE_AMP\n    ):\n        outputs = model(input_tensor)\n\n    probs = torch.sigmoid(\n        outputs.float()\n    )\n\n    return float(\n        probs[0, class_idx].item()\n    )\n\n\ndef get_gradcam(img_np, class_idx):\n\n    input_tensor = normalize_np_to_tensor(\n        img_np\n    )\n\n    targets = [\n        ClassifierOutputTarget(\n            class_idx\n        )\n    ]\n\n    grayscale_cam = cam(\n        input_tensor=input_tensor,\n        targets=targets\n    )[0, :]\n\n    return grayscale_cam\n\n\nprint(\"Explainability functions ready.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"patch_size = 224 // GRID_SIZE\n\n\ndef cam_to_patch_ranking(grayscale_cam):\n\n    patch_scores = np.zeros(\n        (GRID_SIZE, GRID_SIZE),\n        dtype=np.float32\n    )\n\n    for i in range(GRID_SIZE):\n\n        for j in range(GRID_SIZE):\n\n            patch = grayscale_cam[\n                i * patch_size:(i + 1) * patch_size,\n                j * patch_size:(j + 1) * patch_size\n            ]\n\n            patch_scores[i, j] = patch.mean()\n\n    ranking = np.argsort(\n        -patch_scores.flatten()\n    )\n\n    return ranking\n\n\ndef perturb_patches(\n    img_np,\n    patch_indices_to_remove\n):\n\n    perturbed = img_np.copy()\n\n    channel_means = (\n        img_np\n        .reshape(-1, 3)\n        .mean(axis=0)\n    )\n\n    for flat_idx in patch_indices_to_remove:\n\n        i, j = divmod(\n            int(flat_idx),\n            GRID_SIZE\n        )\n\n        perturbed[\n            i * patch_size:(i + 1) * patch_size,\n            j * patch_size:(j + 1) * patch_size,\n            :\n        ] = channel_means\n\n    return perturbed\n\n\nprint(\"Patch functions ready.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compute_aopc(\n    img_np,\n    grayscale_cam,\n    class_idx\n):\n\n    ranking = cam_to_patch_ranking(\n        grayscale_cam\n    )\n\n    total_patches = (\n        GRID_SIZE * GRID_SIZE\n    )\n\n    original_prob = get_class_prob(\n        img_np,\n        class_idx\n    )\n\n    drops = []\n\n    for step in range(\n        1,\n        N_AOPC_STEPS + 1\n    ):\n\n        n_remove = int(\n            total_patches\n            * step\n            / N_AOPC_STEPS\n        )\n\n        perturbed_img = perturb_patches(\n            img_np,\n            ranking[:n_remove]\n        )\n\n        perturbed_prob = get_class_prob(\n            perturbed_img,\n            class_idx\n        )\n\n        drops.append(\n            original_prob - perturbed_prob\n        )\n\n    return float(\n        np.mean(drops)\n    )","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compute_max_sensitivity(\n    img_np,\n    original_cam,\n    class_idx\n):\n\n    noise_std = (\n        NOISE_STD_FRACTION\n        * img_np.std()\n    )\n\n    max_diff = 0.0\n\n    for repeat in range(\n        N_SENSITIVITY_REPEATS\n    ):\n\n        noise = (\n            np.random.normal(\n                loc=0.0,\n                scale=noise_std,\n                size=img_np.shape\n            )\n            .astype(np.float32)\n        )\n\n        noisy_img = np.clip(\n            img_np + noise,\n            0.0,\n            1.0\n        )\n\n        noisy_tensor = normalize_np_to_tensor(\n            noisy_img\n        )\n\n        noisy_cam = cam(\n            input_tensor=noisy_tensor,\n            targets=[\n                ClassifierOutputTarget(\n                    class_idx\n                )\n            ]\n        )[0, :]\n\n        diff = np.linalg.norm(\n            original_cam - noisy_cam\n        )\n\n        max_diff = max(\n            max_diff,\n            float(diff)\n        )\n\n    return float(max_diff)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"records = []\n\ntotal_jobs = sum(\n    len(ids)\n    for class_groups in all_group_ids.values()\n    for ids in class_groups.values()\n)\n\ncompleted = 0\n\nprint(\"=\" * 70)\nprint(\"STARTING EXPLAINABILITY ANALYSIS\")\nprint(\"=\" * 70)\nprint(\"Total image/class/group evaluations:\", total_jobs)\nprint(\n    f\"AOPC steps: {N_AOPC_STEPS} | \"\n    f\"Sensitivity repeats: {N_SENSITIVITY_REPEATS}\"\n)\n\n\nfor class_name in LABEL_COLS:\n\n    class_idx = LABEL_COLS.index(\n        class_name\n    )\n\n    print(\n        f\"\\n{'=' * 70}\\n\"\n        f\"CLASS: {class_name}\\n\"\n        f\"{'=' * 70}\"\n    )\n\n    for group_name in [\n        \"TP\",\n        \"FP\",\n        \"TN\",\n        \"FN\"\n    ]:\n\n        group_ids = all_group_ids[\n            class_name\n        ][group_name]\n\n        print(\n            f\"\\n{group_name}: \"\n            f\"{len(group_ids)} images\"\n        )\n\n        for img_id in tqdm(\n            group_ids,\n            desc=f\"{class_name}/{group_name}\"\n        ):\n\n            try:\n\n                _, img_np = load_image_for_cam(\n                    img_id\n                )\n\n                # Grad-CAM\n                grayscale_cam = get_gradcam(\n                    img_np,\n                    class_idx\n                )\n\n                # AOPC\n                aopc = compute_aopc(\n                    img_np,\n                    grayscale_cam,\n                    class_idx\n                )\n\n                # Max Sensitivity\n                max_sens = compute_max_sensitivity(\n                    img_np,\n                    grayscale_cam,\n                    class_idx\n                )\n\n                # Original probability\n                probability = get_class_prob(\n                    img_np,\n                    class_idx\n                )\n\n                records.append({\n                    \"img_id\": img_id,\n                    \"class\": class_name,\n                    \"group\": group_name,\n                    \"probability\": probability,\n                    \"aopc\": aopc,\n                    \"max_sensitivity\": max_sens\n                })\n\n                completed += 1\n\n            except Exception as e:\n\n                print(\n                    f\"\\nWARNING: failed on \"\n                    f\"{img_id}: {e}\"\n                )\n\n        print(\n            f\"Completed: {completed}/{total_jobs}\"\n        )\n\n\nprint(\"\\nDONE.\")\nprint(\n    \"Successful evaluations:\",\n    len(records)\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"full_explainability_df = pd.DataFrame(\n    records\n)\n\nprint(\"=\" * 70)\nprint(\"FULL EXPLAINABILITY RESULTS\")\nprint(\"=\" * 70)\n\nprint(\n    \"Successful evaluations:\",\n    len(full_explainability_df)\n)\n\n\nsummary_full = (\n    full_explainability_df\n    .groupby([\"class\", \"group\"])\n    .agg(\n        n=(\"aopc\", \"count\"),\n\n        aopc_mean=(\"aopc\", \"mean\"),\n        aopc_median=(\"aopc\", \"median\"),\n        aopc_std=(\"aopc\", \"std\"),\n\n        max_sens_mean=(\n            \"max_sensitivity\",\n            \"mean\"\n        ),\n\n        max_sens_median=(\n            \"max_sensitivity\",\n            \"median\"\n        ),\n\n        max_sens_std=(\n            \"max_sensitivity\",\n            \"std\"\n        ),\n    )\n    .reset_index()\n)\n\nprint(\n    summary_full\n    .round(4)\n    .to_string(index=False)\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"full_out = (\n    \"/kaggle/working/\"\n    \"convnext_base_sanity_AOPC_MaxSensitivity.csv\"\n)\n\nsummary_out = (\n    \"/kaggle/working/\"\n    \"convnext_base_sanity_AOPC_MaxSensitivity_summary.csv\"\n)\n\nfull_explainability_df.to_csv(\n    full_out,\n    index=False\n)\n\nsummary_full.to_csv(\n    summary_out,\n    index=False\n)\n\nprint(\"Saved:\")\nprint(full_out)\nprint(summary_out)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}