{"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\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]","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-07-24T15:05:34.645155Z","iopub.execute_input":"2026-07-24T15:05:34.645833Z","iopub.status.idle":"2026-07-24T15:05:34.649356Z","shell.execute_reply.started":"2026-07-24T15:05:34.645805Z","shell.execute_reply":"2026-07-24T15:05:34.648732Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for d in INPUT_DIRS:\n    print(d, \"->\", \"EXISTS\" if os.path.isdir(d) else \"MISSING\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T15:05:44.249194Z","iopub.execute_input":"2026-07-24T15:05:44.249916Z","iopub.status.idle":"2026-07-24T15:05:44.268864Z","shell.execute_reply.started":"2026-07-24T15:05:44.249889Z","shell.execute_reply":"2026-07-24T15:05:44.268037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\n\nfor d in INPUT_DIRS:\n    if not os.path.isdir(d):\n        print(f\"{d} -> MISSING, skipping\")\n        continue\n    print(f\"listing {d} ...\")\n    t0 = time.time()\n    filenames = [f for f in os.listdir(d) if f.endswith(\".png\")]\n    print(f\"  -> {len(filenames)} png files found in {time.time() - t0:.1f}s\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T15:05:53.04894Z","iopub.execute_input":"2026-07-24T15:05:53.049642Z","iopub.status.idle":"2026-07-24T15:05:53.6448Z","shell.execute_reply.started":"2026-07-24T15:05:53.049615Z","shell.execute_reply":"2026-07-24T15:05:53.644158Z"}},"outputs":[],"execution_count":null},{"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 = \"vit_b16\"\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_resnet50(num_classes=len(LABEL_COLS)):\n    model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)\n    for param in model.parameters():\n        param.requires_grad = False\n    for param in model.layer4.parameters():\n        param.requires_grad = True\n    model.fc = nn.Sequential(\n        nn.Dropout(0.3),\n        nn.Linear(model.fc.in_features, num_classes),\n    )\n    return model.to(DEVICE)\n\n\ndef build_vit_b16(num_classes=len(LABEL_COLS)):\n    model = models.vit_b_16(weights=models.ViT_B_16_Weights.IMAGENET1K_V1)\n    for param in model.parameters():\n        param.requires_grad = False\n    for param in model.encoder.layers[-1].parameters():\n        param.requires_grad = True\n    model.heads = nn.Sequential(\n        nn.Dropout(0.3),\n        nn.Linear(model.hidden_dim, num_classes),\n    )\n    for param in model.heads.parameters():\n        param.requires_grad = True\n    return model.to(DEVICE)\n\n\nMODEL_BUILDERS = {\"resnet50\": build_resnet50, \"vit_b16\": build_vit_b16}\n\nOPTIMIZER_CONFIG = {\n    \"resnet50\": {\"type\": \"sgd\", \"lr\": 1e-3, \"momentum\": 0.9},\n    \"vit_b16\": {\"type\": \"adamw\", \"lr\": 1e-4, \"weight_decay\": 1e-4},\n}\n\nSCHEDULER_CONFIG = {\n    \"resnet50\": {\"warmup_epochs\": 0},\n    \"vit_b16\": {\"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), labels.to(DEVICE)\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), labels.to(DEVICE)\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    val_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False,\n                             num_workers=NUM_WORKERS, pin_memory=True)\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":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T15:07:32.821951Z","iopub.execute_input":"2026-07-24T15:07:32.822725Z","iopub.status.idle":"2026-07-24T17:49:43.377101Z","shell.execute_reply.started":"2026-07-24T15:07:32.822697Z","shell.execute_reply":"2026-07-24T17:49:43.375988Z"}},"outputs":[],"execution_count":null},{"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\nfrom torchvision import models, transforms\nfrom PIL import Image\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    roc_auc_score, average_precision_score, precision_recall_curve,\n    confusion_matrix, f1_score,\n)\nfrom tqdm.auto import tqdm\n\n# ------------------------------------------------------------------\n# CONFIG (same as training run)\n# ------------------------------------------------------------------\nMODEL_NAME = \"vit_b16\"\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nUSE_AMP = DEVICE.type == \"cuda\"\nprint(f\"Using device: {DEVICE}\")\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]\nTRAIN_CSV = (\"/kaggle/input/competitions/rsna-intracranial-hemorrhage-detection/\"\n             \"rsna-intracranial-hemorrhage-detection/stage_2_train.csv\")\nMIN_FILE_SIZE_KB = 40\nMIN_FILE_SIZE_BYTES = MIN_FILE_SIZE_KB * 1024\nVAL_SPLIT = 0.15\nTEST_SPLIT = 0.15\nLABEL_COLS = [\"any\", \"epidural\", \"intraparenchymal\",\n              \"intraventricular\", \"subarachnoid\", \"subdural\"]\nCHECKPOINT_PATH = f\"/kaggle/working/checkpoints/{MODEL_NAME}_best.pt\"\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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-27T09:39:38.590909Z","iopub.execute_input":"2026-07-27T09:39:38.591742Z","iopub.status.idle":"2026-07-27T09:39:38.600443Z","shell.execute_reply.started":"2026-07-27T09:39:38.591713Z","shell.execute_reply":"2026-07-27T09:39:38.599689Z"}},"outputs":[],"execution_count":null},{"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\nfrom torchvision import models, transforms\nfrom PIL import Image\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    roc_auc_score, average_precision_score, precision_recall_curve,\n    confusion_matrix, f1_score,\n)\nfrom tqdm.auto import tqdm\n\n# ------------------------------------------------------------------\n# CONFIG (same as training run)\n# ------------------------------------------------------------------\nMODEL_NAME = \"vit_b16\"\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nUSE_AMP = DEVICE.type == \"cuda\"\nprint(f\"Using device: {DEVICE}\")\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]\nTRAIN_CSV = (\"/kaggle/input/competitions/rsna-intracranial-hemorrhage-detection/\"\n             \"rsna-intracranial-hemorrhage-detection/stage_2_train.csv\")\nMIN_FILE_SIZE_KB = 40\nMIN_FILE_SIZE_BYTES = MIN_FILE_SIZE_KB * 1024\nVAL_SPLIT = 0.15\nTEST_SPLIT = 0.15\nLABEL_COLS = [\"any\", \"epidural\", \"intraparenchymal\",\n              \"intraventricular\", \"subarachnoid\", \"subdural\"]\nCHECKPOINT_PATH = f\"/kaggle/working/checkpoints/{MODEL_NAME}_best.pt\"\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\ndef build_resnet50(num_classes=len(LABEL_COLS)):\n    model = models.resnet50(weights=None)\n    model.fc = nn.Sequential(nn.Dropout(0.3), nn.Linear(model.fc.in_features, num_classes))\n    return model.to(DEVICE)\n\n\ndef build_vit_b16(num_classes=len(LABEL_COLS)):\n    model = models.vit_b_16(weights=None)\n    model.heads = nn.Sequential(nn.Dropout(0.3), nn.Linear(model.hidden_dim, num_classes))\n    return model.to(DEVICE)\n\n\nMODEL_BUILDERS = {\"resnet50\": build_resnet50, \"vit_b16\": build_vit_b16}\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    return y.pivot(index=\"id\", columns=\"sub_type\", values=\"label\")\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        filenames = [f for f in os.listdir(d) if f.endswith(\".png\")]\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            id_to_path[f[:-4]] = full_path\n    print(f\"Scanned {len(input_dirs)} directories: {len(id_to_path)} usable PNGs \"\n          f\"found, {skipped_small} filtered out.\")\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\n# ------------------------------------------------------------------\n# REBUILD: data + model\n# ------------------------------------------------------------------\nprint(\"Scanning input directories...\")\nid_to_path = scan_input_dirs(INPUT_DIRS, MIN_FILE_SIZE_BYTES)\ndf = load_labels(TRAIN_CSV)\ndf = df[df.index.isin(id_to_path.keys())]\ntrain_df, val_df, test_df = make_splits(df)\nprint(f\"Rebuilt splits — train: {len(train_df)}, val: {len(val_df)}, test: {len(test_df)}\")\n\nos.makedirs(\"/kaggle/working\", exist_ok=True)\ntest_df.index.to_series(name=\"id\").to_csv(\"/kaggle/working/test_ids.csv\", index=False)\n\neval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n])\nval_ds = ICHDataset(val_df, id_to_path, transform=eval_transform)\ntest_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# ------------------------------------------------------------------\n# INFERENCE\n# ------------------------------------------------------------------\n@torch.no_grad()\ndef run_inference(loader, desc=\"inference\"):\n    all_probs, all_labels = [], []\n    for imgs, labels in tqdm(loader, desc=desc):\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\nval_probs, val_labels = run_inference(val_loader, desc=\"val inference\")\ntest_probs, test_labels = run_inference(test_loader, desc=\"test inference\")\nprint(f\"Done: {len(val_probs)} val images, {len(test_probs)} test images.\")\n\nnp.savez(\"/kaggle/working/test_inference_cache.npz\",\n         val_probs=val_probs, val_labels=val_labels,\n         test_probs=test_probs, test_labels=test_labels)\nprint(\"Cached val/test probs+labels to /kaggle/working/test_inference_cache.npz\")\n\n\n# ------------------------------------------------------------------\n# METRICS TABLE\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\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}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-27T10:05:24.267423Z","iopub.execute_input":"2026-07-27T10:05:24.268061Z","iopub.status.idle":"2026-07-27T10:19:42.973044Z","shell.execute_reply.started":"2026-07-27T10:05:24.268012Z","shell.execute_reply":"2026-07-27T10:19:42.972012Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom sklearn.metrics import roc_curve, precision_recall_curve, auc\n\n# ------------------------------------------------------------------\n# ROC + PR CURVES — per class, test split\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-27T10:22:50.154816Z","iopub.execute_input":"2026-07-27T10:22:50.155367Z","iopub.status.idle":"2026-07-27T10:22:53.347293Z","shell.execute_reply.started":"2026-07-27T10:22:50.155323Z","shell.execute_reply":"2026-07-27T10:22:53.34629Z"}},"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-27T10:23:09.83382Z","iopub.execute_input":"2026-07-27T10:23:09.834237Z","iopub.status.idle":"2026-07-27T10:29:25.853718Z","shell.execute_reply.started":"2026-07-27T10:23:09.834209Z","shell.execute_reply":"2026-07-27T10:29:25.85302Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install grad-cam -q\n\nimport os\nimport json\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 (ViT) — scaled up sample, for stable AOPC/Max-Sensitivity input.\n# Computes+caches CAMs for N_CAM_SAMPLES confirmed TPs per class,\n# but only visualizes a small subset (grid would be unusable at N=30).\n# ------------------------------------------------------------------\nRANDOM_SEED = 42\nnp.random.seed(RANDOM_SEED)\n\nN_CAM_SAMPLES = 30          # per class, capped by availability (esp. epidural)\nN_VISUALIZE_PER_CLASS = 2   # kept small — just a sanity-check grid\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\n# --- ViT-specific: reshape the 196 patch tokens back into a 14x14 grid ---\n# 224px image / 16px patch = 14 patches per side -> 14*14 = 196 tokens,\n# plus 1 CLS token at index 0 which we drop before reshaping.\ndef vit_reshape_transform(tensor, height=14, width=14):\n    result = tensor[:, 1:, :].reshape(tensor.size(0), height, width, tensor.size(2))\n    # bring channels to dim 1, like a conv activation: [B, C, H, W]\n    result = result.transpose(2, 3).transpose(1, 2)\n    return result\n\ntarget_layers = [model.encoder.layers[-1].ln_1]\ncam = GradCAM(model=model, target_layers=target_layers,\n              reshape_transform=vit_reshape_transform)\n\ncam_cache = {}          # img_id -> {class_name: grayscale_cam}\nimg_cache = {}          # img_id -> img_np (3-window array), reused by AOPC/Max-Sensitivity\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\n# ------------------------------------------------------------------\n# Per-class predictions using val-derived thresholds\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\n# ------------------------------------------------------------------\n# Compute + cache CAMs for N_CAM_SAMPLES confirmed TPs per class.\n# Only the first N_VISUALIZE_PER_CLASS of each are plotted.\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-27T10:45:45.017985Z","iopub.execute_input":"2026-07-27T10:45:45.018719Z","iopub.status.idle":"2026-07-27T10:46:03.508253Z","shell.execute_reply.started":"2026-07-27T10:45:45.018671Z","shell.execute_reply":"2026-07-27T10:46:03.5074Z"}},"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\n\n# ------------------------------------------------------------------\n# AOPC + MAX-SENSITIVITY — ViT-B/16, Grad-CAM explainability validation\n# ------------------------------------------------------------------\n# AOPC (Samek et al., 2016): measures whether the CAM correctly\n# identifies pixels the model actually relies on. Progressively\n# perturbs the highest-ranked CAM regions and measures how fast the\n# predicted probability drops. A steep drop = a faithful explanation.\n#\n# Max-Sensitivity (Yeh et al., 2019): measures explanation robustness.\n# Adds small random noise to the input several times and recomputes\n# the CAM each time; reports the MAX heatmap change observed. Low\n# sensitivity = a robust/trustworthy explanation.\n#\n# Perturbation baseline: PER-CHANNEL MEAN of the image itself, not\n# zero-fill (same rationale as the ResNet run — 0 in a given window\n# channel is a specific clinical HU value, not \"no information\").\n# Perturbation operates on a 14x14 patch grid (16px patches), which\n# for ViT also happens to match the native patch tokenization — so\n# each perturbed patch corresponds to exactly one input token.\n\nN_SAMPLES_PER_CLASS = 50\nN_AOPC_STEPS = 10            # cumulative 10%, 20%, ..., 100% of patches perturbed\nGRID_SIZE = 16               # kept consistent with the ResNet run for comparability;\n                              # NOTE: ViT's native patch grid is 14x14 (16px patches on\n                              # a 224px image already == GRID_SIZE here, so this lines\n                              # up with the model's own tokenization, not just the CAM).\nN_SENSITIVITY_REPEATS = 10\nNOISE_STD_FRACTION = 0.05    # noise magnitude as a fraction of per-channel image std\n\nnp.random.seed(RANDOM_SEED)\npatch_size = 224 // GRID_SIZE\n\n\ndef normalize_np_to_tensor(img_np):\n    \"\"\"img_np: [H,W,3] float in [0,1] -> normalized [1,3,H,W] tensor, matching eval_transform.\"\"\"\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    \"\"\"Average CAM value per grid patch, return patch indices ranked descending by importance.\"\"\"\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())  # descending importance\n    return flat_order\n\n\ndef perturb_patches(img_np, patch_indices_to_remove):\n    \"\"\"Replace given patches (flat indices into GRID_SIZE x GRID_SIZE) with per-channel mean-fill.\"\"\"\n    perturbed = img_np.copy()\n    channel_means = img_np.reshape(-1, 3).mean(axis=0)  # per-channel mean over whole image\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    \"\"\"Recomputes the CAM under small random noise perturbations; returns max heatmap change.\"\"\"\n    path = id_to_path[img_id]\n    orig_tensor, orig_img_np = load_image_for_cam(path)\n    from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\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\n# ------------------------------------------------------------------\n# Run over confirmed true positives per class\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            from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\n            grayscale_cam = cam(input_tensor=input_tensor,\n                                 targets=[ClassifierOutputTarget(class_idx)])[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            \"img_id\": img_id, \"class\": class_name,\n            \"aopc\": aopc, \"max_sensitivity\": max_sens,\n        })\n\naopc_sens_df = pd.DataFrame(records)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"AOPC + MAX-SENSITIVITY SUMMARY (mean ± std, per class)\")\nprint(\"=\" * 80)\nsummary = aopc_sens_df.groupby(\"class\").agg(\n    aopc_mean=(\"aopc\", \"mean\"), aopc_std=(\"aopc\", \"std\"),\n    sens_mean=(\"max_sensitivity\", \"mean\"), sens_std=(\"max_sensitivity\", \"std\"),\n    n=(\"aopc\", \"count\"),\n).round(4)\nprint(summary.to_string())\n\nout_path = f\"/kaggle/working/{MODEL_NAME}_aopc_max_sensitivity.csv\"\naopc_sens_df.to_csv(out_path, index=False)\nsummary_out_path = f\"/kaggle/working/{MODEL_NAME}_aopc_max_sensitivity_summary.csv\"\nsummary.to_csv(summary_out_path)\nprint(f\"\\nSaved per-image results to {out_path}\")\nprint(f\"Saved per-class summary to {summary_out_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-27T10:46:13.532619Z","iopub.execute_input":"2026-07-27T10:46:13.533101Z","iopub.status.idle":"2026-07-27T10:49:33.712527Z","shell.execute_reply.started":"2026-07-27T10:46:13.533068Z","shell.execute_reply":"2026-07-27T10:49:33.711933Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\nimport os\nfrom IPython.display import FileLink\n\nzip_base_name = \"/kaggle/working/vit_b16_full_outputs\"\noutput_dir = \"/kaggle/working/vit_b16_full_outputs_files\"\nos.makedirs(output_dir, exist_ok=True)\n\nfor f in [\n    f\"/kaggle/working/checkpoints/{MODEL_NAME}_best.pt\",\n    f\"/kaggle/working/{MODEL_NAME}_eval_metrics.csv\",\n    f\"/kaggle/working/{MODEL_NAME}_roc_pr_curves.png\",\n    f\"/kaggle/working/{MODEL_NAME}_calibration_metrics.csv\",\n    f\"/kaggle/working/{MODEL_NAME}_calibration_reliability.png\",\n    f\"/kaggle/working/{MODEL_NAME}_temperature.txt\",\n    \"/kaggle/working/test_ids.csv\",\n    \"/kaggle/working/test_inference_cache.npz\",\n    f\"/kaggle/working/{MODEL_NAME}_gradcam_sample_visualization.png\",\n    f\"/kaggle/working/{MODEL_NAME}_bone_shortcut_correlations.csv\",\n    f\"/kaggle/working/{MODEL_NAME}_aopc_max_sensitivity.csv\",\n    f\"/kaggle/working/{MODEL_NAME}_aopc_max_sensitivity_summary.csv\",\n    \"/kaggle/working/cam_cache.pkl\",\n]:\n    if os.path.exists(f):\n        shutil.copy(f, output_dir)\n        print(f\"Added: {f}\")\n    else:\n        print(f\"Skipped (not found): {f}\")\n\nshutil.make_archive(zip_base_name, \"zip\", output_dir)\nprint(f\"\\nBacked up to {zip_base_name}.zip\")\n\nFileLink(f\"{zip_base_name}.zip\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-27T10:51:41.649722Z","iopub.execute_input":"2026-07-27T10:51:41.650548Z","iopub.status.idle":"2026-07-27T10:52:05.930625Z","shell.execute_reply.started":"2026-07-27T10:51:41.650515Z","shell.execute_reply":"2026-07-27T10:52:05.930011Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfor f in sorted(os.listdir(\"/kaggle/working\")):\n    print(f)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\nimport os\n\n# ================================================================\n# PACKAGE ALL ViT-B/16 OUTPUTS\n# ================================================================\n\nexport_dir = \"/kaggle/working/vit_b16_export\"\n\n# Start clean so stale ConvNeXt files cannot accidentally enter\n# the final archive.\nif os.path.exists(export_dir):\n    shutil.rmtree(export_dir)\n\nos.makedirs(export_dir, exist_ok=True)\n\n\nfiles_to_zip = [\n    \"cam_cache.pkl\",\n    \"cam_cache_multilayer.pkl\",\n\n    \"vit_b16_aopc_max_sensitivity.csv\",\n    \"vit_b16_aopc_max_sensitivity_summary.csv\",\n\n    \"vit_b16_bone_shortcut_correlations.csv\",\n    \"vit_b16_bone_shortcut_correlations_multilayer.csv\",\n\n    \"vit_b16_calibration_metrics.csv\",\n    \"vit_b16_calibration_reliability.png\",\n\n    \"vit_b16_eval_metrics.csv\",\n\n    \"vit_b16_gradcam_sample_visualization.png\",\n    \"vit_b16_gradcam_multilayer_sample_visualization.png\",\n\n    \"vit_b16_roc_pr_curves.png\",\n    \"vit_b16_temperature.txt\",\n\n    \"test_ids.csv\",\n]\n\n\nfor filename in files_to_zip:\n\n    source = f\"/kaggle/working/{filename}\"\n    destination = f\"{export_dir}/{filename}\"\n\n    if os.path.exists(source):\n        shutil.copy2(source, destination)\n        print(f\"Added: {filename}\")\n\n    else:\n        print(f\"WARNING — missing: {filename}\")\n\n\n# ------------------------------------------------\n# Grad-CAM example directories\n# ------------------------------------------------\n\ngradcam_dir = \"/kaggle/working/gradcam_examples\"\n\nif os.path.exists(gradcam_dir):\n\n    shutil.copytree(\n        gradcam_dir,\n        f\"{export_dir}/gradcam_examples\"\n    )\n\n    print(\"Added: gradcam_examples/\")\n\n\ngradcam_multilayer_dir = (\n    \"/kaggle/working/gradcam_examples_multilayer\"\n)\n\nif os.path.exists(gradcam_multilayer_dir):\n\n    shutil.copytree(\n        gradcam_multilayer_dir,\n        f\"{export_dir}/gradcam_examples_multilayer\"\n    )\n\n    print(\n        \"Added: gradcam_examples_multilayer/\"\n    )\n\n\n# ------------------------------------------------\n# Create archive\n# ------------------------------------------------\n\narchive_base = (\n    \"/kaggle/working/\"\n    \"vit_b16_full_outputs\"\n)\n\narchive_path = shutil.make_archive(\n    archive_base,\n    \"zip\",\n    export_dir\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"ViT-B/16 OUTPUT ARCHIVE CREATED\")\nprint(\"=\" * 80)\nprint(f\"Archive: {archive_path}\")\nprint(\"=\" * 80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T03:39:05.064468Z","iopub.execute_input":"2026-09-14T03:39:05.064723Z","iopub.status.idle":"2026-09-14T03:39:14.467829Z","shell.execute_reply.started":"2026-09-14T03:39:05.064699Z","shell.execute_reply":"2026-09-14T03:39:14.467089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#cell 13\n\nimport os\nimport numpy as np\nimport torch\n\nMODEL_NAME = \"vit_b16\"\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\n\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\n# ================================================================\n# ViT-B/16 EXPLAINABILITY SETTINGS\n# ================================================================\n\n# ViT-B/16 receives 224x224 images using 16x16 patches:\n#\n#   224 / 16 = 14\n#\n# Therefore:\n#\n#   14 x 14 = 196 spatial tokens\n#\n# plus one CLS token.\n\nGRID_SIZE = 14\nPATCH_SIZE = 16\n\nN_AOPC_STEPS = 10\n\nN_SENSITIVITY_REPEATS = 10\n\nNOISE_STD_FRACTION = 0.05\n\n\n# ================================================================\n# FULL POPULATION\n# ================================================================\n#\n# IMPORTANT:\n# Do NOT cap TP/FP/TN/FN groups.\n#\n# Every eligible test image is evaluated.\n#\n# Therefore MAX_IMAGES_PER_GROUP is deliberately absent.\n# ================================================================\n\n\nCHECKPOINT_PATH_EXPLAINABILITY = (\n    \"/kaggle/working/\"\n    \"vit_b16_population_checkpoint.csv\"\n)\n\nCHECKPOINT_EVERY = 200\n\n\nprint(\"Configuration loaded.\")\nprint(\"Model:\", MODEL_NAME)\nprint(\"Checkpoint:\", CHECKPOINT_PATH)\n\nprint(\n    f\"ViT patch grid: \"\n    f\"{GRID_SIZE} x {GRID_SIZE}\"\n)\n\nprint(\n    f\"Native patch size: \"\n    f\"{PATCH_SIZE} x {PATCH_SIZE}px\"\n)\n\nprint(\n    f\"AOPC steps: {N_AOPC_STEPS}\"\n)\n\nprint(\n    f\"Max-Sensitivity repeats: \"\n    f\"{N_SENSITIVITY_REPEATS}\"\n)\n\nprint(\"FULL POPULATION MODE: ENABLED\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T03:39:14.469718Z","iopub.execute_input":"2026-09-14T03:39:14.469938Z","iopub.status.idle":"2026-09-14T03:39:14.477876Z","shell.execute_reply.started":"2026-09-14T03:39:14.469917Z","shell.execute_reply":"2026-09-14T03:39:14.477034Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# DATASET\n# ================================================================\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom PIL import Image\nfrom torch.utils.data import Dataset, DataLoader\n\n\nclass ICHDataset(Dataset):\n\n    def __init__(\n        self,\n        df,\n        id_to_path,\n        transform=None\n    ):\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\n    def __len__(self):\n        return len(self.ids)\n\n\n    def __getitem__(self, idx):\n\n        img_id = self.ids[idx]\n\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=[\n            \"id\",\n            \"sub_type\"\n        ]\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-14T03:39:14.481731Z","iopub.execute_input":"2026-09-14T03:39:14.482093Z","iopub.status.idle":"2026-09-14T03:39:14.55639Z","shell.execute_reply.started":"2026-09-14T03:39:14.482036Z","shell.execute_reply":"2026-09-14T03:39:14.555629Z"}},"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-14T03:39:14.557417Z","iopub.execute_input":"2026-09-14T03:39:14.557745Z","iopub.status.idle":"2026-09-14T03:48:35.352693Z","shell.execute_reply.started":"2026-09-14T03:39:14.557712Z","shell.execute_reply":"2026-09-14T03:48:35.351812Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# LOAD LABELS + TRAIN / VALIDATION / TEST SPLITS\n# ================================================================\n\nimport pandas as pd\n\nfrom sklearn.model_selection import train_test_split\n\n\n# ------------------------------------------------\n# Load labels\n# ------------------------------------------------\n\ndf = load_labels(TRAIN_CSV)\n\nprint(\n    \"Total labelled IDs:\",\n    f\"{len(df):,}\"\n)\n\n\n# ------------------------------------------------\n# Keep only images for which we have a valid path\n# ------------------------------------------------\n\ndf = df[\n    df.index.isin(\n        id_to_path.keys()\n    )\n].copy()\n\nprint(\n    \"Labelled images with valid paths:\",\n    f\"{len(df):,}\"\n)\n\n\n# ------------------------------------------------\n# Create splits\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=(\n            VAL_SPLIT\n            / (1.0 - TEST_SPLIT)\n        ),\n        random_state=42,\n        stratify=train_val_df[\"any\"]\n    )\n\n    return (\n        train_df,\n        val_df,\n        test_df\n    )\n\n\ntrain_df, val_df, test_df = (\n    make_splits(df)\n)\n\n\n# ------------------------------------------------\n# Report split sizes\n# ------------------------------------------------\n\nprint(\"\\nSPLIT:\")\nprint(\n    \"Train:\",\n    f\"{len(train_df):,}\"\n)\nprint(\n    \"Validation:\",\n    f\"{len(val_df):,}\"\n)\nprint(\n    \"Test:\",\n    f\"{len(test_df):,}\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T03:48:35.353787Z","iopub.execute_input":"2026-09-14T03:48:35.354167Z","iopub.status.idle":"2026-09-14T03:48:54.578206Z","shell.execute_reply.started":"2026-09-14T03:48:35.35414Z","shell.execute_reply":"2026-09-14T03:48:54.577422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision import transforms\n\neval_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-14T03:48:54.580776Z","iopub.execute_input":"2026-09-14T03:48:54.580995Z","iopub.status.idle":"2026-09-14T03:48:54.590711Z","shell.execute_reply.started":"2026-09-14T03:48:54.580974Z","shell.execute_reply":"2026-09-14T03:48:54.589774Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# ViT-B/16 MODEL — MATCH EXISTING CHECKPOINT ARCHITECTURE\n# ================================================================\n\nimport os\nimport torch\nimport torch.nn as nn\nfrom torchvision import models\n\n\ndef build_vit_b16(\n    num_classes=len(LABEL_COLS)\n):\n\n    # ------------------------------------------------------------\n    # Load ImageNet-pretrained ViT-B/16 backbone\n    # ------------------------------------------------------------\n\n    model = models.vit_b_16(\n        weights=models.ViT_B_16_Weights.IMAGENET1K_V1\n    )\n\n\n    # ------------------------------------------------------------\n    # Freeze everything\n    # ------------------------------------------------------------\n\n    for param in model.parameters():\n        param.requires_grad = False\n\n\n    # ------------------------------------------------------------\n    # Unfreeze final encoder block\n    # ------------------------------------------------------------\n\n    for param in model.encoder.layers[-1].parameters():\n        param.requires_grad = True\n\n\n    # ------------------------------------------------------------\n    # IMPORTANT:\n    #\n    # The existing checkpoint contains:\n    #\n    #     heads.1.weight\n    #     heads.1.bias\n    #\n    # Therefore the classification head MUST be a direct\n    # Sequential under model.heads.\n    #\n    # Do NOT use model.heads.head.\n    # ------------------------------------------------------------\n\n    in_features = model.heads.head.in_features\n\n    model.heads = nn.Sequential(\n        nn.Dropout(0.3),\n        nn.Linear(\n            in_features,\n            num_classes\n        )\n    )\n\n\n    # ------------------------------------------------------------\n    # Make classification head trainable\n    # ------------------------------------------------------------\n\n    for param in model.heads.parameters():\n        param.requires_grad = True\n\n\n    return model.to(DEVICE)\n\n\n# ================================================================\n# LOAD EXISTING TRAINED ViT-B/16 CHECKPOINT\n# ================================================================\n\nif not os.path.exists(CHECKPOINT_PATH):\n\n    raise FileNotFoundError(\n        f\"Checkpoint not found:\\n\"\n        f\"{CHECKPOINT_PATH}\"\n    )\n\n\nprint(\"Building ViT-B/16...\")\n\nmodel = build_vit_b16()\n\n\n# ------------------------------------------------\n# Verify architecture BEFORE loading checkpoint\n# ------------------------------------------------\n\nprint(\"\\nClassification head:\")\nprint(model.heads)\n\nprint(\"\\nExpected checkpoint-compatible keys:\")\nprint(\"  heads.0.weight\")\nprint(\"  heads.0.bias\")\nprint(\"  heads.1.weight\")\nprint(\"  heads.1.bias\")\n\n\n# ------------------------------------------------\n# Load checkpoint\n# ------------------------------------------------\n\nprint(\"\\nLoading checkpoint...\")\n\ncheckpoint = torch.load(\n    CHECKPOINT_PATH,\n    map_location=DEVICE,\n    weights_only=False\n)\n\n\n# ------------------------------------------------\n# Load trained weights\n# ------------------------------------------------\n\nmodel.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\n\nmodel.eval()\n\n\n# ================================================================\n# CONFIRM\n# ================================================================\n\nprint(\"\\n\" + \"=\" * 60)\nprint(\"MODEL LOADED SUCCESSFULLY\")\nprint(\"=\" * 60)\n\nprint(\"Architecture: ViT-B/16\")\n\nprint(\n    \"Epoch:\",\n    checkpoint[\"epoch\"]\n)\n\nprint(\n    \"Val loss:\",\n    f\"{checkpoint['val_loss']:.4f}\"\n)\n\nprint(\n    \"Number of classes:\",\n    len(LABEL_COLS)\n)\n\nprint(\n    \"Labels:\",\n    LABEL_COLS\n)\n\nprint(\"=\" * 60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T03:48:54.591945Z","iopub.execute_input":"2026-09-14T03:48:54.592326Z","iopub.status.idle":"2026-09-14T03:48:58.108238Z","shell.execute_reply.started":"2026-09-14T03:48:54.592292Z","shell.execute_reply":"2026-09-14T03:48:58.107557Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 19 — ViT-B/16 VALIDATION + TEST INFERENCE\n# ================================================================\n#\n# FIX:\n# Use num_workers=0 and pin_memory=False for inference.\n#\n# The previous multiprocessing DataLoader configuration was causing\n# worker cleanup errors during test inference:\n#\n#   RuntimeError: cannot join current thread\n#   AssertionError: can only test a child process\n#\n# This does NOT change the model, checkpoint, preprocessing,\n# predictions, labels, or evaluation methodology.\n#\n# It only removes DataLoader multiprocessing from inference.\n# ================================================================\n\nimport os\nimport numpy as np\nimport torch\nfrom torch.utils.data import DataLoader\nfrom tqdm.auto import tqdm\n\n\n# ------------------------------------------------\n# Rebuild inference loaders safely\n# ------------------------------------------------\n\nval_loader = DataLoader(\n    val_ds,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=0,\n    pin_memory=False\n)\n\ntest_loader = DataLoader(\n    test_ds,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=0,\n    pin_memory=False\n)\n\nprint(\"Inference DataLoaders rebuilt.\")\nprint(\"Validation batches:\", len(val_loader))\nprint(\"Test batches:\", len(test_loader))\n\n\n# ================================================================\n# INFERENCE FUNCTION\n# ================================================================\n\n@torch.no_grad()\ndef run_inference(loader, description):\n\n    model.eval()\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=False\n        )\n\n        with torch.autocast(\n            device_type=DEVICE.type,\n            enabled=USE_AMP\n        ):\n            logits = model(imgs)\n\n        probs = torch.sigmoid(\n            logits.float()\n        )\n\n        all_probs.append(\n            probs.cpu().numpy()\n        )\n\n        all_labels.append(\n            labels.numpy()\n        )\n\n    return (\n        np.concatenate(all_probs, axis=0),\n        np.concatenate(all_labels, axis=0)\n    )\n\n\n# ================================================================\n# VALIDATION INFERENCE\n# ================================================================\n\nprint(\"=\" * 70)\nprint(\"RUNNING VALIDATION INFERENCE\")\nprint(\"=\" * 70)\n\nval_probs, val_labels = run_inference(\n    val_loader,\n    \"Validation inference\"\n)\n\nprint(\n    \"Validation probabilities:\",\n    val_probs.shape\n)\n\nprint(\n    \"Validation labels:\",\n    val_labels.shape\n)\n\n\n# ================================================================\n# TEST INFERENCE\n# ================================================================\n\nprint(\"=\" * 70)\nprint(\"RUNNING TEST INFERENCE\")\nprint(\"=\" * 70)\n\ntest_probs, test_labels = run_inference(\n    test_loader,\n    \"Test inference\"\n)\n\nprint(\n    \"Test probabilities:\",\n    test_probs.shape\n)\n\nprint(\n    \"Test labels:\",\n    test_labels.shape\n)\n\n\n# ================================================================\n# HARD ASSERTIONS\n# ================================================================\n\nassert val_probs.shape == (\n    len(val_df),\n    len(LABEL_COLS)\n)\n\nassert val_labels.shape == (\n    len(val_df),\n    len(LABEL_COLS)\n)\n\nassert test_probs.shape == (\n    len(test_df),\n    len(LABEL_COLS)\n)\n\nassert test_labels.shape == (\n    len(test_df),\n    len(LABEL_COLS)\n)\n\n\n# ================================================================\n# CACHE\n# ================================================================\n\ncache_path = (\n    \"/kaggle/working/\"\n    \"test_inference_cache.npz\"\n)\n\nnp.savez(\n    cache_path,\n    val_probs=val_probs,\n    val_labels=val_labels,\n    test_probs=test_probs,\n    test_labels=test_labels\n)\n\n\nprint(\"=\" * 70)\nprint(\"INFERENCE COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Validation images:\",\n    len(val_probs)\n)\n\nprint(\n    \"Test images:\",\n    len(test_probs)\n)\n\nprint(\"Saved:\", cache_path)\n\nprint(\"=\" * 70)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T03:48:58.109043Z","iopub.execute_input":"2026-09-14T03:48:58.109319Z","iopub.status.idle":"2026-09-14T04:08:12.818513Z","shell.execute_reply.started":"2026-09-14T03:48:58.10927Z","shell.execute_reply":"2026-09-14T04:08:12.817917Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 20 — VALIDATION-DERIVED CLASSIFICATION THRESHOLDS\n# ================================================================\n\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn.metrics import precision_recall_curve\n\n\n# ------------------------------------------------\n# Find threshold maximizing F1 on validation set\n# ------------------------------------------------\n\ndef best_f1_threshold(y_true, y_prob):\n\n    precisions, recalls, thresholds = precision_recall_curve(\n        y_true,\n        y_prob\n    )\n\n    # precision_recall_curve returns one extra precision/recall\n    # value compared with thresholds.\n    f1s = (\n        2.0 * precisions[:-1] * recalls[:-1]\n        /\n        (\n            precisions[:-1]\n            + recalls[:-1]\n            + 1e-12\n        )\n    )\n\n    if len(thresholds) == 0:\n        return 0.5\n\n    best_idx = int(np.nanargmax(f1s))\n\n    return float(thresholds[best_idx])\n\n\n# ------------------------------------------------\n# Calculate one threshold per hemorrhage class\n# ------------------------------------------------\n\nrows = []\n\nfor class_idx, class_name in enumerate(LABEL_COLS):\n\n    threshold = best_f1_threshold(\n        val_labels[:, class_idx],\n        val_probs[:, class_idx]\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\n\n# ------------------------------------------------\n# Display\n# ------------------------------------------------\n\nprint(\"=\" * 70)\nprint(\"VALIDATION-DERIVED CLASSIFICATION THRESHOLDS\")\nprint(\"=\" * 70)\n\nprint(\n    results_df.round(4).to_string()\n)\n\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T04:08:12.819675Z","iopub.execute_input":"2026-09-14T04:08:12.819957Z","iopub.status.idle":"2026-09-14T04:08:12.87636Z","shell.execute_reply.started":"2026-09-14T04:08:12.819933Z","shell.execute_reply":"2026-09-14T04:08:12.875702Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#cell21\n# ================================================================\n# CELL 21 — TEST CONFUSION GROUPS: TP / FP / TN / FN\n# ================================================================\n\nimport numpy as np\nimport pandas as pd\n\nassert \"test_probs\" in globals(), \"Cell 19 did not create test_probs. Re-run the UPDATED Cell 19 from this notebook.\"\nassert \"test_labels\" in globals(), \"Cell 19 did not create test_labels. Re-run the UPDATED Cell 19 from this notebook.\"\nassert \"results_df\" in globals(), \"Run Cell 20 first.\"\n\ntest_ids_array = test_df.index.to_numpy()\ngroup_records = []\n\nfor class_idx, class_name in enumerate(LABEL_COLS):\n\n    threshold = float(\n        results_df.loc[class_name, \"threshold_from_val\"]\n    )\n\n    y_true = test_labels[:, class_idx].astype(int)\n    y_prob = test_probs[:, class_idx]\n    y_pred = (y_prob >= threshold).astype(int)\n\n    for img_id, true_label, pred_label, prob in zip(\n        test_ids_array,\n        y_true,\n        y_pred,\n        y_prob\n    ):\n\n        if true_label == 1 and pred_label == 1:\n            group = \"TP\"\n        elif true_label == 0 and pred_label == 1:\n            group = \"FP\"\n        elif true_label == 1 and pred_label == 0:\n            group = \"FN\"\n        else:\n            group = \"TN\"\n\n        group_records.append({\n            \"img_id\": str(img_id),\n            \"class\": class_name,\n            \"true_label\": int(true_label),\n            \"pred_label\": int(pred_label),\n            \"probability\": float(prob),\n            \"threshold\": threshold,\n            \"group\": group,\n        })\n\nconfusion_groups_df = pd.DataFrame(group_records)\n\nprint(\"=\" * 80)\nprint(\"TEST CONFUSION GROUPS\")\nprint(\"=\" * 80)\n\ncounts = (\n    confusion_groups_df\n    .groupby([\"class\", \"group\"])\n    .size()\n    .unstack(fill_value=0)\n    .reindex(columns=[\"TN\", \"FP\", \"FN\", \"TP\"], fill_value=0)\n)\n\nprint(counts.to_string())\n\nexpected = len(test_df) * len(LABEL_COLS)\n\nprint(\"\\nTotal rows:\", len(confusion_groups_df))\nprint(\"Expected:\", expected)\n\nassert len(confusion_groups_df) == expected\n\nprint(\"=\" * 80)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T04:08:12.877336Z","iopub.execute_input":"2026-09-14T04:08:12.877697Z","iopub.status.idle":"2026-09-14T04:08:13.594657Z","shell.execute_reply.started":"2026-09-14T04:08:12.877673Z","shell.execute_reply":"2026-09-14T04:08:13.593942Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ================================================================\n# CELL 22 — BUILD FULL POPULATION CONFUSION GROUP IDS\n# ================================================================\n#\n# Creates:\n#\n#     all_group_ids[class_name][group_name]\n#\n# from the complete test-set confusion classification generated\n# in Cell 21.\n#\n# NO SAMPLING / NO CAP\n#\n# Every test image belongs to exactly one of:\n#\n#     TP / FP / TN / FN\n#\n# for every hemorrhage class.\n# ================================================================\n\nimport numpy as np\n\n\nGROUP_NAMES = [\n    \"TP\",\n    \"FP\",\n    \"TN\",\n    \"FN\",\n]\n\n\n# ================================================================\n# INITIALIZE STRUCTURE\n# ================================================================\n\nall_group_ids = {\n    class_name: {\n        group_name: []\n        for group_name in GROUP_NAMES\n    }\n    for class_name in LABEL_COLS\n}\n\n\n# ================================================================\n# POPULATE GROUPS\n# ================================================================\n\nfor class_name in LABEL_COLS:\n\n    class_df = confusion_groups_df[\n        confusion_groups_df[\"class\"] == class_name\n    ]\n\n    for group_name in GROUP_NAMES:\n\n        group_ids = (\n            class_df.loc[\n                class_df[\"group\"] == group_name,\n                \"img_id\"\n            ]\n            .astype(str)\n            .tolist()\n        )\n\n        all_group_ids[\n            class_name\n        ][\n            group_name\n        ] = group_ids\n\n\n# ================================================================\n# VALIDATION\n# ================================================================\n\nprint(\"=\" * 80)\nprint(\"FULL POPULATION CONFUSION GROUPS\")\nprint(\"=\" * 80)\n\ntotal_expected = len(test_df)\n\nfor class_name in LABEL_COLS:\n\n    counts = {\n        group_name: len(\n            all_group_ids[class_name][group_name]\n        )\n        for group_name in GROUP_NAMES\n    }\n\n    total = sum(counts.values())\n\n    print(\n        f\"{class_name:20s} \"\n        f\"TP={counts['TP']:,}  \"\n        f\"FP={counts['FP']:,}  \"\n        f\"TN={counts['TN']:,}  \"\n        f\"FN={counts['FN']:,}  \"\n        f\"TOTAL={total:,}\"\n    )\n\n    # Every test image must appear exactly once\n    # for this class.\n    assert total == total_expected, (\n        f\"{class_name}: expected {total_expected:,} \"\n        f\"images but found {total:,}\"\n    )\n\n\n# ================================================================\n# GLOBAL VALIDATION\n# ================================================================\n\nprint(\"\\nExpected images per class:\", f\"{total_expected:,}\")\nprint(\"Confusion groups validated successfully.\")\n\nassert set(all_group_ids.keys()) == set(LABEL_COLS)\n\nfor class_name in LABEL_COLS:\n\n    assert set(\n        all_group_ids[class_name].keys()\n    ) == set(GROUP_NAMES)\n\nprint(\"=\" * 80)\nprint(\"all_group_ids READY FOR CELL 28\")\nprint(\"=\" * 80)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T04:08:13.595762Z","iopub.execute_input":"2026-09-14T04:08:13.596067Z","iopub.status.idle":"2026-09-14T04:08:13.839721Z","shell.execute_reply.started":"2026-09-14T04:08:13.596027Z","shell.execute_reply":"2026-09-14T04:08:13.838991Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 22 — ViT-B/16 GRAD-CAM SETUP\n# ================================================================\n\n# Install Grad-CAM if it is not already available\nimport sys\nimport subprocess\nimport importlib.util\n\nif importlib.util.find_spec(\"pytorch_grad_cam\") is None:\n    print(\"Installing pytorch-grad-cam...\")\n    subprocess.check_call([\n        sys.executable, \"-m\", \"pip\", \"install\", \"-q\", \"grad-cam\"\n    ])\n    print(\"Installation complete.\")\n\n# Imports\nimport numpy as np\nimport torch\nfrom pytorch_grad_cam import GradCAM\n\n\n# ----------------------------------------------------------------\n# ViT reshape transform\n# Removes CLS token and converts 196 patch tokens into 14x14 grid\n# ----------------------------------------------------------------\n\ndef vit_reshape_transform(tensor, height=14, width=14):\n    # tensor shape: [batch, tokens, channels]\n    # ViT-B/16 has 197 tokens = 1 CLS + 196 spatial patches\n\n    result = tensor[:, 1:, :]  # Remove CLS token\n\n    result = result.reshape(\n        tensor.size(0),\n        height,\n        width,\n        tensor.size(2)\n    )\n\n    # Convert:\n    # [B, H, W, C] -> [B, C, H, W]\n    result = result.permute(0, 3, 1, 2)\n\n    return result\n\n\n# ----------------------------------------------------------------\n# Target layer\n# ----------------------------------------------------------------\n\ntarget_layers = [model.encoder.layers[-1].ln_1]\n\n\n# ----------------------------------------------------------------\n# Create Grad-CAM\n# ----------------------------------------------------------------\n\ncam = GradCAM(\n    model=model,\n    target_layers=target_layers,\n    reshape_transform=vit_reshape_transform\n)\n\n\n# ----------------------------------------------------------------\n# Verification\n# ----------------------------------------------------------------\n\nprint(\"=\" * 70)\nprint(\"ViT-B/16 Grad-CAM READY\")\nprint(\"=\" * 70)\nprint(\"Target layer: model.encoder.layers[-1].ln_1\")\nprint(\"Input: 224 x 224\")\nprint(\"Patch size: 16 x 16\")\nprint(\"Patch grid: 14 x 14\")\nprint(\"Spatial tokens: 196\")\nprint(\"CLS token: removed before spatial reshape\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T04:08:13.840679Z","iopub.execute_input":"2026-09-14T04:08:13.841273Z","iopub.status.idle":"2026-09-14T04:08:24.241207Z","shell.execute_reply.started":"2026-09-14T04:08:13.841248Z","shell.execute_reply":"2026-09-14T04:08:24.240553Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 23 — ViT IMAGE + GRAD-CAM HELPERS (BATCHED)\n# ================================================================\n\nimport numpy as np\nimport torch\nfrom PIL import Image\nfrom pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\n\nMEAN_ARR = np.array([0.485, 0.456, 0.406], dtype=np.float32)\nSTD_ARR = np.array([0.229, 0.224, 0.225], dtype=np.float32)\n\n\ndef load_image_for_cam(img_id):\n    \"\"\"Single-image loader kept for validation/back-compat.\"\"\"\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 load_images_batch(img_ids):\n    \"\"\"\n    Loads N images.\n\n    Returns a list of (224, 224, 3) float32 arrays\n    in [0, 1].\n    \"\"\"\n\n    img_nps = []\n\n    for img_id in img_ids:\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_nps.append(\n            np.asarray(img)\n            .astype(np.float32)\n            / 255.0\n        )\n\n    return img_nps, None\n\n\ndef normalize_np_to_tensor(img_np):\n    \"\"\"Single-image normalization.\"\"\"\n\n    normed = (\n        img_np - MEAN_ARR\n    ) / STD_ARR\n\n    return (\n        torch.from_numpy(\n            normed.transpose(2, 0, 1)\n        )\n        .unsqueeze(0)\n        .float()\n        .to(DEVICE)\n    )\n\n\ndef normalize_np_batch_to_tensor(img_np_batch):\n    \"\"\"\n    img_np_batch:\n        (N, 224, 224, 3)\n\n    Returns:\n        (N, 3, 224, 224)\n    \"\"\"\n\n    normed = (\n        img_np_batch - MEAN_ARR\n    ) / STD_ARR\n\n    return (\n        torch.from_numpy(\n            normed.transpose(0, 3, 1, 2)\n        )\n        .float()\n        .to(DEVICE)\n    )\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    return float(\n        torch.sigmoid(outputs.float())[\n            0, class_idx\n        ].item()\n    )\n\n\n@torch.no_grad()\ndef get_class_probs_batch(\n    img_np_batch,\n    class_idx\n):\n    \"\"\"\n    Returns one sigmoid probability per image.\n    \"\"\"\n\n    batch_tensor = (\n        normalize_np_batch_to_tensor(\n            img_np_batch\n        )\n    )\n\n    with torch.autocast(\n        device_type=DEVICE.type,\n        enabled=USE_AMP\n    ):\n        outputs = model(batch_tensor)\n\n    probs = torch.sigmoid(\n        outputs.float()\n    )[:, class_idx]\n\n    return (\n        probs\n        .detach()\n        .cpu()\n        .numpy()\n    )\n\n\ndef get_gradcam(img_np, class_idx):\n    \"\"\"Single-image Grad-CAM.\"\"\"\n\n    input_tensor = (\n        normalize_np_to_tensor(\n            img_np\n        )\n    )\n\n    return cam(\n        input_tensor=input_tensor,\n        targets=[\n            ClassifierOutputTarget(\n                class_idx\n            )\n        ]\n    )[0, :]\n\n\ndef get_gradcam_batch(\n    img_np_batch,\n    class_idx\n):\n    \"\"\"\n    Batched Grad-CAM.\n\n    All images share the same target class,\n    which is exactly what happens inside a\n    fixed class/group chunk.\n    \"\"\"\n\n    n = img_np_batch.shape[0]\n\n    batch_tensor = (\n        normalize_np_batch_to_tensor(\n            img_np_batch\n        )\n    )\n\n    targets = [\n        ClassifierOutputTarget(\n            class_idx\n        )\n        for _ in range(n)\n    ]\n\n    return cam(\n        input_tensor=batch_tensor,\n        targets=targets\n    )\n\n\nprint(\n    \"ViT explainability helpers ready (batched).\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T04:08:24.242245Z","iopub.execute_input":"2026-09-14T04:08:24.242811Z","iopub.status.idle":"2026-09-14T04:08:24.255919Z","shell.execute_reply.started":"2026-09-14T04:08:24.242784Z","shell.execute_reply":"2026-09-14T04:08:24.254995Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 25 — ViT NATIVE PATCH RANKING\n# ================================================================\n\nGRID_SIZE = 14\nPATCH_SIZE = 16\n\nassert GRID_SIZE * PATCH_SIZE == 224\n\n\ndef cam_to_patch_ranking(grayscale_cam):\n    \"\"\"\n    Mean CAM value per native 16x16 ViT patch.\n\n    Highest relevance is removed first\n    (MoRF ordering).\n    \"\"\"\n\n    patch_scores = (\n        grayscale_cam\n        .reshape(\n            GRID_SIZE,\n            PATCH_SIZE,\n            GRID_SIZE,\n            PATCH_SIZE\n        )\n        .mean(axis=(1, 3))\n    )\n\n    return np.argsort(\n        -patch_scores.reshape(-1)\n    )\n\n\ndef perturb_patches(\n    img_np,\n    patch_indices\n):\n    \"\"\"\n    Replace selected patches with the\n    image's per-channel mean.\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:\n\n        row, col = divmod(\n            int(flat_idx),\n            GRID_SIZE\n        )\n\n        perturbed[\n            row * PATCH_SIZE:\n            (row + 1) * PATCH_SIZE,\n\n            col * PATCH_SIZE:\n            (col + 1) * PATCH_SIZE,\n\n            :\n        ] = channel_means\n\n    return perturbed\n\n\ndef build_aopc_perturbation_batch(\n    img_np,\n    ranking,\n    n_steps\n):\n    \"\"\"\n    Builds the cumulative MoRF perturbations.\n\n    Returns:\n        (n_steps, 224, 224, 3)\n    \"\"\"\n\n    total_patches = (\n        GRID_SIZE * GRID_SIZE\n    )\n\n    steps = []\n\n    for step in range(\n        1,\n        n_steps + 1\n    ):\n\n        n_remove = max(\n            1,\n            int(\n                total_patches\n                * step\n                / n_steps\n            )\n        )\n\n        steps.append(\n            perturb_patches(\n                img_np,\n                ranking[:n_remove]\n            )\n        )\n\n    return np.stack(\n        steps,\n        axis=0\n    )\n\n\nprint(\n    \"=\" * 70\n)\nprint(\n    \"NATIVE ViT PATCH PERTURBATION READY\"\n)\nprint(\n    \"=\" * 70\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T04:08:24.256738Z","iopub.execute_input":"2026-09-14T04:08:24.257027Z","iopub.status.idle":"2026-09-14T04:08:24.278961Z","shell.execute_reply.started":"2026-09-14T04:08:24.256993Z","shell.execute_reply":"2026-09-14T04:08:24.278346Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 26 — BATCHED AOPC FOR ViT-B/16\n# ================================================================\n\nN_AOPC_STEPS = 10\n\n\ndef compute_aopc(\n    img_np,\n    grayscale_cam,\n    class_idx\n):\n    \"\"\"\n    Single-image AOPC.\n\n    Kept for validation against the\n    batched implementation.\n    \"\"\"\n\n    ranking = cam_to_patch_ranking(\n        grayscale_cam\n    )\n\n    original_prob = get_class_prob(\n        img_np,\n        class_idx\n    )\n\n    perturbed = (\n        build_aopc_perturbation_batch(\n            img_np,\n            ranking,\n            N_AOPC_STEPS\n        )\n    )\n\n    batch_tensor = (\n        normalize_np_batch_to_tensor(\n            perturbed\n        )\n    )\n\n    with torch.no_grad():\n\n        with torch.autocast(\n            device_type=DEVICE.type,\n            enabled=USE_AMP\n        ):\n\n            outputs = model(\n                batch_tensor\n            )\n\n        probs = (\n            torch.sigmoid(\n                outputs.float()\n            )[:, class_idx]\n            .cpu()\n            .numpy()\n        )\n\n    return float(\n        (\n            original_prob\n            - probs\n        ).mean()\n    )\n\n\ndef compute_aopc_batch(\n    img_np_batch,\n    grayscale_cam_batch,\n    class_idx,\n    original_probs\n):\n    \"\"\"\n    Cross-image batched AOPC.\n\n    Returns:\n        (B,) AOPC value per image.\n    \"\"\"\n\n    B = img_np_batch.shape[0]\n\n    all_perturbed = []\n\n    for b in range(B):\n\n        ranking = (\n            cam_to_patch_ranking(\n                grayscale_cam_batch[b]\n            )\n        )\n\n        all_perturbed.append(\n            build_aopc_perturbation_batch(\n                img_np_batch[b],\n                ranking,\n                N_AOPC_STEPS\n            )\n        )\n\n    stacked = np.concatenate(\n        all_perturbed,\n        axis=0\n    )\n\n    batch_tensor = (\n        normalize_np_batch_to_tensor(\n            stacked\n        )\n    )\n\n    with torch.no_grad():\n\n        with torch.autocast(\n            device_type=DEVICE.type,\n            enabled=USE_AMP\n        ):\n\n            outputs = model(\n                batch_tensor\n            )\n\n        probs = (\n            torch.sigmoid(\n                outputs.float()\n            )[:, class_idx]\n        )\n\n    probs = (\n        probs\n        .detach()\n        .cpu()\n        .numpy()\n        .reshape(\n            B,\n            N_AOPC_STEPS\n        )\n    )\n\n    drops = (\n        original_probs[:, None]\n        - probs\n    )\n\n    return drops.mean(\n        axis=1\n    )\n\n\nprint(\n    \"=\" * 70\n)\nprint(\n    \"BATCHED ViT AOPC READY (cross-image)\"\n)\nprint(\n    \"=\" * 70\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T04:08:24.279932Z","iopub.execute_input":"2026-09-14T04:08:24.280195Z","iopub.status.idle":"2026-09-14T04:08:24.297322Z","shell.execute_reply.started":"2026-09-14T04:08:24.280166Z","shell.execute_reply":"2026-09-14T04:08:24.296532Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 27 — BATCHED MAX-SENSITIVITY FOR ViT-B/16\n# ================================================================\n\nN_SENSITIVITY_REPEATS = 10\nNOISE_STD_FRACTION = 0.05\n\n\ndef compute_max_sensitivity(\n    img_id,\n    class_idx\n):\n    \"\"\"\n    Original single-image version.\n\n    Kept for validation only.\n    \"\"\"\n\n    _, orig_img_np = (\n        load_image_for_cam(\n            img_id\n        )\n    )\n\n    orig_tensor = (\n        normalize_np_to_tensor(\n            orig_img_np\n        )\n    )\n\n    targets = [\n        ClassifierOutputTarget(\n            class_idx\n        )\n    ]\n\n    orig_cam = cam(\n        input_tensor=orig_tensor,\n        targets=targets\n    )[0, :]\n\n    channel_std = (\n        orig_img_np\n        .reshape(-1, 3)\n        .std(axis=0)\n    )\n\n    noisy_imgs = []\n\n    for _ in range(\n        N_SENSITIVITY_REPEATS\n    ):\n\n        noise = np.random.normal(\n            0.0,\n            NOISE_STD_FRACTION\n            * channel_std,\n            orig_img_np.shape\n        ).astype(\n            np.float32\n        )\n\n        noisy_imgs.append(\n            np.clip(\n                orig_img_np + noise,\n                0.0,\n                1.0\n            )\n        )\n\n    batch_np = np.stack(\n        noisy_imgs,\n        axis=0\n    )\n\n    batch_tensor = (\n        normalize_np_batch_to_tensor(\n            batch_np\n        )\n    )\n\n    targets = [\n        ClassifierOutputTarget(\n            class_idx\n        )\n        for _ in range(\n            N_SENSITIVITY_REPEATS\n        )\n    ]\n\n    noisy_cams = cam(\n        input_tensor=batch_tensor,\n        targets=targets\n    )\n\n    diffs = np.linalg.norm(\n        (\n            noisy_cams\n            - orig_cam\n        ).reshape(\n            N_SENSITIVITY_REPEATS,\n            -1\n        ),\n        axis=1\n    )\n\n    return float(\n        diffs.max()\n    )\n\n\ndef compute_max_sensitivity_batch(\n    img_np_batch,\n    clean_cam_batch,\n    class_idx,\n    rng\n):\n    \"\"\"\n    Cross-image batched Max-Sensitivity.\n\n    clean_cam_batch is the clean CAM\n    already calculated for the images.\n\n    Returns:\n        (B,) max-sensitivity value per image.\n    \"\"\"\n\n    B, H, W, _ = (\n        img_np_batch.shape\n    )\n\n    all_noisy = []\n\n    for b in range(B):\n\n        channel_std = (\n            img_np_batch[b]\n            .reshape(-1, 3)\n            .std(axis=0)\n        )\n\n        noise = rng.normal(\n            0.0,\n            NOISE_STD_FRACTION\n            * channel_std,\n            size=(\n                N_SENSITIVITY_REPEATS,\n            )\n            + img_np_batch[b].shape\n        ).astype(\n            np.float32\n        )\n\n        noisy = np.clip(\n            img_np_batch[b][None, ...]\n            + noise,\n            0.0,\n            1.0\n        )\n\n        all_noisy.append(\n            noisy\n        )\n\n    stacked = np.concatenate(\n        all_noisy,\n        axis=0\n    )\n\n    batch_tensor = (\n        normalize_np_batch_to_tensor(\n            stacked\n        )\n    )\n\n    targets = [\n        ClassifierOutputTarget(\n            class_idx\n        )\n        for _ in range(\n            B * N_SENSITIVITY_REPEATS\n        )\n    ]\n\n    noisy_cams = cam(\n        input_tensor=batch_tensor,\n        targets=targets\n    )\n\n    noisy_cams = noisy_cams.reshape(\n        B,\n        N_SENSITIVITY_REPEATS,\n        H,\n        W\n    )\n\n    diffs = np.linalg.norm(\n        (\n            noisy_cams\n            - clean_cam_batch[:, None, :, :]\n        ).reshape(\n            B,\n            N_SENSITIVITY_REPEATS,\n            -1\n        ),\n        axis=2\n    )\n\n    return diffs.max(\n        axis=1\n    )\n\n\nprint(\n    \"=\" * 70\n)\nprint(\n    \"BATCHED ViT MAX-SENSITIVITY READY\"\n)\nprint(\n    \"=\" * 70\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T04:08:24.298331Z","iopub.execute_input":"2026-09-14T04:08:24.298659Z","iopub.status.idle":"2026-09-14T04:08:24.315724Z","shell.execute_reply.started":"2026-09-14T04:08:24.298618Z","shell.execute_reply":"2026-09-14T04:08:24.315143Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# VALIDATION CELL\n# RUN THIS BEFORE THE FULL POPULATION LOOP\n# ================================================================\n\nDET_ATOL = 1e-4\nDET_RTOL = 1e-3\nCAM_ATOL = 1e-3\n\nVALIDATION_PAIRS = [\n    (\n        all_group_ids[\"any\"][\"TP\"][0],\n        \"any\"\n    ),\n    (\n        all_group_ids[\"any\"][\"FP\"][0],\n        \"any\"\n    ),\n    (\n        all_group_ids[\"epidural\"][\"TN\"][0],\n        \"epidural\"\n    ),\n    (\n        all_group_ids[\"epidural\"][\"FN\"][0],\n        \"epidural\"\n    ),\n    (\n        all_group_ids[\"subdural\"][\"TP\"][0],\n        \"subdural\"\n    ),\n]\n\n\nprint(\n    f\"{'img_id':>14} \"\n    f\"{'class':>16} \"\n    f\"{'prob_old':>10} \"\n    f\"{'prob_new':>10} \"\n    f\"{'aopc_old':>10} \"\n    f\"{'aopc_new':>10} \"\n    f\"{'cam_absdiff':>12}\"\n)\n\nall_ok = True\n\n\nfor img_id, class_name in VALIDATION_PAIRS:\n\n    class_idx = LABEL_COLS.index(\n        class_name\n    )\n\n    # ----------------------------\n    # OLD / SINGLE IMAGE\n    # ----------------------------\n\n    _, img_np_old = (\n        load_image_for_cam(\n            img_id\n        )\n    )\n\n    old_cam = get_gradcam(\n        img_np_old,\n        class_idx\n    )\n\n    old_prob = get_class_prob(\n        img_np_old,\n        class_idx\n    )\n\n    old_aopc = compute_aopc(\n        img_np_old,\n        old_cam,\n        class_idx\n    )\n\n    # ----------------------------\n    # NEW / BATCHED\n    # ----------------------------\n\n    img_nps_new, _ = (\n        load_images_batch(\n            [img_id]\n        )\n    )\n\n    img_np_batch = np.stack(\n        img_nps_new,\n        axis=0\n    )\n\n    new_cams = get_gradcam_batch(\n        img_np_batch,\n        class_idx\n    )\n\n    new_probs = get_class_probs_batch(\n        img_np_batch,\n        class_idx\n    )\n\n    new_aopc = compute_aopc_batch(\n        img_np_batch,\n        new_cams,\n        class_idx,\n        new_probs\n    )[0]\n\n    # ----------------------------\n    # COMPARE\n    # ----------------------------\n\n    cam_diff = float(\n        np.abs(\n            old_cam\n            - new_cams[0]\n        ).max()\n    )\n\n    prob_ok = bool(\n        np.isclose(\n            old_prob,\n            new_probs[0],\n            atol=DET_ATOL,\n            rtol=DET_RTOL\n        )\n    )\n\n    aopc_ok = bool(\n        np.isclose(\n            old_aopc,\n            new_aopc,\n            atol=DET_ATOL,\n            rtol=DET_RTOL\n        )\n    )\n\n    cam_ok = (\n        cam_diff < CAM_ATOL\n    )\n\n    all_ok &= (\n        prob_ok\n        and aopc_ok\n        and cam_ok\n    )\n\n    print(\n        f\"{img_id:>14} \"\n        f\"{class_name:>16} \"\n        f\"{old_prob:>10.6f} \"\n        f\"{new_probs[0]:>10.6f} \"\n        f\"{old_aopc:>10.6f} \"\n        f\"{new_aopc:>10.6f} \"\n        f\"{cam_diff:>12.2e}\"\n        + (\n            \"\"\n            if (\n                prob_ok\n                and aopc_ok\n                and cam_ok\n            )\n            else \"  <-- CHECK\"\n        )\n    )\n\n\nprint()\n\nprint(\n    \"Deterministic validation:\",\n    \"PASSED\"\n    if all_ok\n    else \"FAILED — investigate before full run\"\n)\n\n\n# ================================================================\n# MAX-SENSITIVITY DISTRIBUTION CHECK\n# ================================================================\n\nN_TRIALS = 20\n\nprint(\n    f\"\\nMax-Sensitivity stochastic check — \"\n    f\"{N_TRIALS} repeated trials:\"\n)\n\n\nfor img_id, class_name in (\n    VALIDATION_PAIRS[:2]\n):\n\n    class_idx = LABEL_COLS.index(\n        class_name\n    )\n\n    _, img_np_old = (\n        load_image_for_cam(\n            img_id\n        )\n    )\n\n    old_cam_for_sens = get_gradcam(\n        img_np_old,\n        class_idx\n    )\n\n    # Old implementation\n    old_vals = [\n        compute_max_sensitivity(\n            img_id,\n            class_idx\n        )\n        for _ in range(\n            N_TRIALS\n        )\n    ]\n\n    # New implementation\n    img_np_batch = np.stack(\n        [img_np_old],\n        axis=0\n    )\n\n    rng_val = np.random.default_rng(\n        0\n    )\n\n    new_vals = [\n        compute_max_sensitivity_batch(\n            img_np_batch,\n            old_cam_for_sens[None, ...],\n            class_idx,\n            rng_val\n        )[0]\n        for _ in range(\n            N_TRIALS\n        )\n    ]\n\n    print(\n        f\"{img_id} / {class_name}: \"\n        f\"old mean={np.mean(old_vals):.4f} \"\n        f\"std={np.std(old_vals):.4f}  |  \"\n        f\"new mean={np.mean(new_vals):.4f} \"\n        f\"std={np.std(new_vals):.4f}\"\n    )\n\n\nprint(\n    \"\\nFor Max-Sensitivity, the exact values \"\n    \"do not need to match because the noise \"\n    \"is stochastic.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T04:08:24.316545Z","iopub.execute_input":"2026-09-14T04:08:24.316989Z","iopub.status.idle":"2026-09-14T04:08:43.447056Z","shell.execute_reply.started":"2026-09-14T04:08:24.316957Z","shell.execute_reply":"2026-09-14T04:08:43.446287Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 28 — FULL POPULATION ViT AOPC + MAX-SENSITIVITY (BATCHED)\n# OOM-HARDENED VERSION\n# ================================================================\n#\n# Same evaluation as before — nothing about what's measured changes.\n# Two additions to fix the CUDA OOM you hit in `intraventricular / TN`:\n#   1. BATCH_SIZE_EVAL dropped 32 -> 16 (320 -> 160 images per\n#      AOPC/sensitivity sub-batch).\n#   2. Explicit torch.cuda.empty_cache() + gc.collect() calls, both\n#      on the OOM fallback path and periodically during normal\n#      operation, so freed memory actually gets returned to the\n#      allocator's pool instead of sitting reserved-but-unused.\n# ================================================================\n\nimport os\nimport gc\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom tqdm.auto import tqdm\n\n# Put this as early as possible in a fresh session (before the model\n# is built) for best effect — harmless to also set here.\nos.environ.setdefault(\"PYTORCH_CUDA_ALLOC_CONF\", \"expandable_segments:True\")\n\nRANDOM_SEED = 42\nnp.random.seed(RANDOM_SEED)\nrng = np.random.default_rng(RANDOM_SEED)\n\nCHECKPOINT_EVERY = 200\n\n# Was 32. Halved for this GPU/session after the OOM in\n# intraventricular/TN — 16 -> 160 images per AOPC/sensitivity\n# sub-batch instead of 320. Drop to 8 if it still OOMs.\nBATCH_SIZE_EVAL = 16\n\n# How often (in chunks) to force a cache clear during normal,\n# non-failing operation. Every chunk would add unnecessary overhead;\n# every 10 keeps the cost negligible while still preventing the slow\n# creep back toward full memory.\nEMPTY_CACHE_EVERY_N_CHUNKS = 10\n\nPOPULATION_CHECKPOINT = \"/kaggle/working/vit_b16_population_checkpoint.csv\"\nGROUP_NAMES = [\"TP\", \"FP\", \"TN\", \"FN\"]\n\n# ---- build all_group_ids (unchanged from the original cell) ----\nall_group_ids = {\n    class_name: {group_name: [] for group_name in GROUP_NAMES}\n    for class_name in LABEL_COLS\n}\nfor class_name in LABEL_COLS:\n    class_df = confusion_groups_df[confusion_groups_df[\"class\"] == class_name]\n    for group_name in GROUP_NAMES:\n        all_group_ids[class_name][group_name] = (\n            class_df.loc[class_df[\"group\"] == group_name, \"img_id\"].astype(str).tolist()\n        )\n\nexpected_images = len(test_df)\nfor class_name in LABEL_COLS:\n    total = sum(len(all_group_ids[class_name][g]) for g in GROUP_NAMES)\n    assert total == expected_images, f\"{class_name}: expected {expected_images}, found {total}\"\n\n# ---- load existing checkpoint (unchanged logic — nothing invalidated) ----\nif os.path.exists(POPULATION_CHECKPOINT):\n    population_df = pd.read_csv(POPULATION_CHECKPOINT)\n    completed_keys = set(\n        zip(\n            population_df[\"img_id\"].astype(str),\n            population_df[\"class\"].astype(str),\n            population_df[\"group\"].astype(str),\n        )\n    )\n    records = population_df.to_dict(orient=\"records\")\n    print(\"Resuming checkpoint:\", len(records), \"completed rows\")\nelse:\n    records = []\n    completed_keys = set()\n    print(\"Starting full population analysis.\")\n\ntotal_jobs = sum(\n    len(ids) for class_groups in all_group_ids.values() for ids in class_groups.values()\n)\ncompleted = len(records)\nfailed = 0\n\nprint(\"=\" * 80)\nprint(\"STARTING / RESUMING FULL ViT POPULATION ANALYSIS (BATCHED, OOM-HARDENED)\")\nprint(\"=\" * 80)\nprint(f\"Total image x class evaluations: {total_jobs:,}\")\nprint(f\"Already completed: {completed:,}\")\nprint(f\"Remaining: {max(0, total_jobs - completed):,}\")\nprint(f\"Batch size (images per GPU call): {BATCH_SIZE_EVAL}\")\nprint(f\"AOPC steps: {N_AOPC_STEPS}\")\nprint(f\"Max-Sensitivity repeats: {N_SENSITIVITY_REPEATS}\")\nprint(f\"Checkpoint path: {POPULATION_CHECKPOINT}\")\nprint(\"=\" * 80)\n\n\ndef process_one_image(img_id, class_name, class_idx, group_name):\n    \"\"\"\n    Single-image fallback path, used only when a whole batch throws.\n    Uses the exact same batched helpers with N=1.\n    \"\"\"\n    img_nps, _ = load_images_batch([img_id])\n    img_np_batch = np.stack(img_nps, axis=0)\n    cams = get_gradcam_batch(img_np_batch, class_idx)\n    probs = get_class_probs_batch(img_np_batch, class_idx)\n    aopc = compute_aopc_batch(img_np_batch, cams, class_idx, probs)[0]\n    sens = compute_max_sensitivity_batch(img_np_batch, cams, class_idx, rng)[0]\n    return {\n        \"img_id\": img_id,\n        \"class\": class_name,\n        \"group\": group_name,\n        \"probability\": float(probs[0]),\n        \"aopc\": float(aopc),\n        \"max_sensitivity\": float(sens),\n    }\n\n\nfor class_name in LABEL_COLS:\n    class_idx = LABEL_COLS.index(class_name)\n    print(f\"\\n{'=' * 80}\\nCLASS: {class_name}\\n{'=' * 80}\")\n\n    for group_name in GROUP_NAMES:\n        group_ids = all_group_ids[class_name][group_name]\n        print(f\"\\n{group_name}: {len(group_ids):,} images in this group\")\n        if len(group_ids) == 0:\n            continue\n\n        remaining_ids = [\n            str(img_id) for img_id in group_ids\n            if (str(img_id), str(class_name), str(group_name)) not in completed_keys\n        ]\n        skipped = len(group_ids) - len(remaining_ids)\n        if skipped > 0:\n            print(f\"Skipping {skipped:,} already-completed evaluations.\")\n        if len(remaining_ids) == 0:\n            print(\"Nothing remaining in this group.\")\n            continue\n\n        n_chunks = (len(remaining_ids) + BATCH_SIZE_EVAL - 1) // BATCH_SIZE_EVAL\n        pbar = tqdm(total=len(remaining_ids), desc=f\"{class_name} / {group_name}\")\n\n        for chunk_idx in range(n_chunks):\n            chunk_ids = remaining_ids[\n                chunk_idx * BATCH_SIZE_EVAL: (chunk_idx + 1) * BATCH_SIZE_EVAL\n            ]\n\n            try:\n                img_nps, _ = load_images_batch(chunk_ids)\n                img_np_batch = np.stack(img_nps, axis=0)\n\n                clean_cams = get_gradcam_batch(img_np_batch, class_idx)\n                probs = get_class_probs_batch(img_np_batch, class_idx)\n                aopc_vals = compute_aopc_batch(img_np_batch, clean_cams, class_idx, probs)\n                sens_vals = compute_max_sensitivity_batch(img_np_batch, clean_cams, class_idx, rng)\n\n                for i, img_id in enumerate(chunk_ids):\n                    records.append({\n                        \"img_id\": img_id,\n                        \"class\": class_name,\n                        \"group\": group_name,\n                        \"probability\": float(probs[i]),\n                        \"aopc\": float(aopc_vals[i]),\n                        \"max_sensitivity\": float(sens_vals[i]),\n                    })\n                    completed_keys.add((img_id, class_name, group_name))\n                    completed += 1\n\n                pbar.update(len(chunk_ids))\n\n                # Periodic cache clear during normal operation, to stop\n                # the slow creep back toward full memory that led to\n                # the earlier OOM. Cheap relative to the GPU work above.\n                if (chunk_idx + 1) % EMPTY_CACHE_EVERY_N_CHUNKS == 0:\n                    torch.cuda.empty_cache()\n\n            except torch.cuda.OutOfMemoryError as chunk_exc:\n                print(f\"\\nWARNING: batch OOM ({chunk_exc}); clearing cache and \"\n                      f\"retrying images individually.\")\n                # Free whatever this failed attempt was holding before\n                # retrying — this is the step that was missing before.\n                del img_nps, img_np_batch\n                torch.cuda.empty_cache()\n                gc.collect()\n\n                for img_id in chunk_ids:\n                    try:\n                        row = process_one_image(img_id, class_name, class_idx, group_name)\n                        records.append(row)\n                        completed_keys.add((img_id, class_name, group_name))\n                        completed += 1\n                    except torch.cuda.OutOfMemoryError as img_exc:\n                        failed += 1\n                        print(f\"WARNING: still OOM on single image {img_id} / \"\n                              f\"{class_name} / {group_name}: {img_exc}\")\n                        torch.cuda.empty_cache()\n                        gc.collect()\n                    except Exception as img_exc:\n                        failed += 1\n                        print(f\"WARNING: failed on {img_id} / {class_name} / {group_name}: {img_exc}\")\n                    pbar.update(1)\n\n            except Exception as chunk_exc:\n                print(f\"\\nWARNING: batch failed ({chunk_exc}); retrying images individually.\")\n                for img_id in chunk_ids:\n                    try:\n                        row = process_one_image(img_id, class_name, class_idx, group_name)\n                        records.append(row)\n                        completed_keys.add((img_id, class_name, group_name))\n                        completed += 1\n                    except Exception as img_exc:\n                        failed += 1\n                        print(f\"WARNING: failed on {img_id} / {class_name} / {group_name}: {img_exc}\")\n                    pbar.update(1)\n\n            if completed % CHECKPOINT_EVERY < BATCH_SIZE_EVAL:\n                pd.DataFrame(records).to_csv(POPULATION_CHECKPOINT, index=False)\n                print(f\"\\n[CHECKPOINT] Saved {len(records):,} records \"\n                      f\"({completed:,} completed, {failed:,} failed so far)\")\n\n        pbar.close()\n        pd.DataFrame(records).to_csv(POPULATION_CHECKPOINT, index=False)\n        torch.cuda.empty_cache()  # clear at group boundaries too — cheap, and this is a natural pause point\n        print(f\"Completed so far: {completed:,} / {total_jobs:,}\")\n        print(f\"Failed so far: {failed:,}\")\n        print(f\"Checkpoint updated: {POPULATION_CHECKPOINT} ({len(records):,} records)\")\n\npopulation_df = pd.DataFrame(records)\npopulation_df.to_csv(POPULATION_CHECKPOINT, index=False)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"FULL POPULATION ANALYSIS COMPLETE\")\nprint(\"=\" * 80)\nprint(f\"Successful evaluations: {completed:,}\")\nprint(f\"Failed evaluations: {failed:,}\")\nprint(f\"Total records stored: {len(population_df):,}\")\nprint(f\"Expected evaluations: {total_jobs:,}\")\nprint(f\"Checkpoint: {POPULATION_CHECKPOINT}\")\nprint(\"=\" * 80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T04:08:43.448062Z","iopub.execute_input":"2026-09-14T04:08:43.448409Z","iopub.status.idle":"2026-09-14T05:08:32.301891Z","shell.execute_reply.started":"2026-09-14T04:08:43.448384Z","shell.execute_reply":"2026-09-14T05:08:32.301115Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#cell29\n# ================================================================\n# CELL 29 — SAVE FULL POPULATION RESULTS\n# ================================================================\n\nimport pandas as pd\n\nassert \"population_df\" in globals()\n\npopulation_out = (\n    \"/kaggle/working/\"\n    \"vit_b16_FULL_population_AOPC_MaxSensitivity.csv\"\n)\n\nsummary_out = (\n    \"/kaggle/working/\"\n    \"vit_b16_FULL_population_AOPC_MaxSensitivity_summary.csv\"\n)\n\npopulation_df.to_csv(\n    population_out,\n    index=False\n)\n\nsummary_df = (\n    population_df\n    .groupby([\"class\", \"group\"], observed=False)\n    .agg(\n        aopc_mean=(\"aopc\", \"mean\"),\n        aopc_std=(\"aopc\", \"std\"),\n        sensitivity_mean=(\"max_sensitivity\", \"mean\"),\n        sensitivity_std=(\"max_sensitivity\", \"std\"),\n        n=(\"aopc\", \"count\")\n    )\n    .reset_index()\n)\n\nsummary_df.to_csv(\n    summary_out,\n    index=False\n)\n\nprint(\"=\" * 80)\nprint(\"FULL POPULATION OUTPUTS SAVED\")\nprint(\"=\" * 80)\nprint(population_out)\nprint(summary_out)\nprint(\"=\" * 80)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T05:08:32.30301Z","iopub.execute_input":"2026-09-14T05:08:32.303649Z","iopub.status.idle":"2026-09-14T05:08:33.724005Z","shell.execute_reply.started":"2026-09-14T05:08:32.303613Z","shell.execute_reply":"2026-09-14T05:08:33.7233Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#cell30\n# ================================================================\n# CELL 30 — PARQUET + ViT-B/16 METADATA\n# ================================================================\n\nimport json\n\nassert \"population_df\" in globals()\n\nparquet_path = (\n    \"/kaggle/working/\"\n    \"vit_b16_FULL_population_AOPC_MaxSensitivity.parquet\"\n)\n\nmetadata_path = (\n    \"/kaggle/working/\"\n    \"vit_b16_FULL_population_metadata.json\"\n)\n\npopulation_df.to_parquet(\n    parquet_path,\n    index=False\n)\n\nmetadata = {\n    \"model\": \"ViT-B/16\",\n    \"input_resolution\": \"224x224\",\n    \"patch_size\": \"16x16\",\n    \"patch_grid\": \"14x14\",\n    \"spatial_patch_tokens\": 196,\n    \"cls_token\": \"removed before Grad-CAM spatial reshape\",\n    \"gradcam_target_layer\": \"model.encoder.layers[-1].ln_1\",\n    \"aopc\": \"Mean probability drop after progressively perturbing the highest-ranked native ViT patches at 10%, 20%, ..., 100%.\",\n    \"aopc_baseline\": \"Per-channel mean fill computed from the original image.\",\n    \"max_sensitivity\": \"Maximum L2 difference between the original Grad-CAM and Grad-CAMs recomputed after small random input-noise perturbations.\",\n    \"max_sensitivity_repeats\": N_SENSITIVITY_REPEATS,\n    \"noise_std_fraction\": NOISE_STD_FRACTION,\n    \"population\": \"Every eligible test image x every hemorrhage class, assigned to TN, FP, FN, or TP using validation-derived thresholds.\",\n    \"threshold_source\": \"Validation split F1-maximizing threshold per class.\",\n    \"random_seed\": RANDOM_SEED,\n    \"checkpoint\": POPULATION_CHECKPOINT\n}\n\nwith open(\n    metadata_path,\n    \"w\",\n    encoding=\"utf-8\"\n) as f:\n    json.dump(metadata, f, indent=2)\n\nprint(\"=\" * 80)\nprint(\"PARQUET + METADATA SAVED\")\nprint(\"=\" * 80)\nprint(parquet_path)\nprint(metadata_path)\nprint(\"=\" * 80)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T05:08:33.724894Z","iopub.execute_input":"2026-09-14T05:08:33.725176Z","iopub.status.idle":"2026-09-14T05:08:34.022342Z","shell.execute_reply.started":"2026-09-14T05:08:33.725146Z","shell.execute_reply":"2026-09-14T05:08:34.02165Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#cell31\n# ================================================================\n# CELL 31 — FULL POPULATION INTEGRITY + STATISTICS\n# ================================================================\n\nimport pandas as pd\n\nassert \"population_df\" in globals()\n\nstats_df = population_df.copy()\n\nstats_df[\"class\"] = pd.Categorical(\n    stats_df[\"class\"],\n    categories=LABEL_COLS,\n    ordered=True\n)\n\nstats_df[\"group\"] = pd.Categorical(\n    stats_df[\"group\"],\n    categories=[\"TN\", \"FP\", \"FN\", \"TP\"],\n    ordered=True\n)\n\nexpected_rows = len(test_df) * len(LABEL_COLS)\n\nprint(\"=\" * 80)\nprint(\"FULL POPULATION INTEGRITY CHECK\")\nprint(\"=\" * 80)\n\nprint(\"Observed rows:\", f\"{len(stats_df):,}\")\nprint(\"Expected rows:\", f\"{expected_rows:,}\")\n\nassert len(stats_df) == expected_rows\n\nmissing = stats_df.isna().sum()\nprint(\"\\nMissing values:\")\nprint(\n    missing[missing > 0].to_string()\n    if (missing > 0).any()\n    else \"None\"\n)\n\nduplicates = stats_df.duplicated(\n    subset=[\"img_id\", \"class\", \"group\"]\n).sum()\n\nprint(\n    \"\\nDuplicate (img_id, class, group) rows:\",\n    duplicates\n)\n\nassert duplicates == 0\n\nimage_class_counts = (\n    stats_df\n    .groupby([\"img_id\", \"class\"], observed=False)\n    .size()\n)\n\nassert (image_class_counts == 1).all()\n\nprint(\n    \"Image/class combinations:\",\n    len(image_class_counts)\n)\n\nprint(\"\\nObservations per class/group:\")\n\ncounts = (\n    stats_df\n    .groupby([\"class\", \"group\"], observed=False)\n    .size()\n    .unstack(fill_value=0)\n    .reindex(\n        columns=[\"TN\", \"FP\", \"FN\", \"TP\"],\n        fill_value=0\n    )\n)\n\nprint(counts.to_string())\n\nprint(\"\\nAOPC statistics:\")\n\nprint(\n    stats_df\n    .groupby([\"class\", \"group\"], observed=False)[\"aopc\"]\n    .agg([\"count\", \"mean\", \"std\", \"median\"])\n    .round(4)\n    .to_string()\n)\n\nprint(\"\\nMax-Sensitivity statistics:\")\n\nprint(\n    stats_df\n    .groupby([\"class\", \"group\"], observed=False)[\"max_sensitivity\"]\n    .agg([\"count\", \"mean\", \"std\", \"median\"])\n    .round(4)\n    .to_string()\n)\n\nprint(\"=\" * 80)\nprint(\"INTEGRITY CHECK COMPLETE\")\nprint(\"=\" * 80)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T05:08:34.025277Z","iopub.execute_input":"2026-09-14T05:08:34.025727Z","iopub.status.idle":"2026-09-14T05:08:34.308444Z","shell.execute_reply.started":"2026-09-14T05:08:34.025701Z","shell.execute_reply":"2026-09-14T05:08:34.307828Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 26 — STATISTICAL ANALYSIS DATASET\n# ================================================================\n\nstats_df = full_explainability_df.copy()\n\n# Explicit categorical ordering\nstats_df[\"class\"] = pd.Categorical(\n    stats_df[\"class\"],\n    categories=LABEL_COLS,\n    ordered=True\n)\n\nstats_df[\"group\"] = pd.Categorical(\n    stats_df[\"group\"],\n    categories=[\"TN\", \"FP\", \"FN\", \"TP\"],\n    ordered=True\n)\n\n\nprint(\"=\" * 80)\nprint(\"STATISTICAL ANALYSIS DATASET\")\nprint(\"=\" * 80)\n\nprint(\n    stats_df[\n        [\n            \"img_id\",\n            \"class\",\n            \"group\",\n            \"probability\",\n            \"aopc\",\n            \"max_sensitivity\"\n        ]\n    ].head()\n)\n\n\nprint(\"\\nMissing values:\")\nprint(\n    stats_df[\n        [\n            \"probability\",\n            \"aopc\",\n            \"max_sensitivity\"\n        ]\n    ]\n    .isna()\n    .sum()\n)\n\n\nprint(\"\\nObservations per class/group:\")\nprint(\n    stats_df\n    .groupby(\n        [\"class\", \"group\"],\n        observed=True\n    )\n    .size()\n    .unstack(fill_value=0)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T05:08:34.309314Z","iopub.execute_input":"2026-09-14T05:08:34.309558Z","iopub.status.idle":"2026-09-14T05:08:34.317597Z","shell.execute_reply.started":"2026-09-14T05:08:34.309536Z","shell.execute_reply":"2026-09-14T05:08:34.316705Z"}},"outputs":[],"execution_count":null}]}