{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":10338,"databundleVersionId":862042,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# =============================================================================\n# MS-MHAA: Multi-Scale Multi-Head Attention Autoencoder\n# Ablation Study — RSNA Pneumonia Detection Challenge\n# Self-Supervised Anomaly Detection (Zero Labels During Training)\n#\n# 8 Model Variants:\n#   M0 — Baseline         (no innovations)\n#   M1 — Multi-Scale only (Innovation 1)\n#   M2 — Attention only   (Innovation 2)\n#   M3 — Multi-Head only  (Innovation 3)\n#   M4 — Scale + Attn     (1 + 2)\n#   M5 — Scale + MultiHd  (1 + 3)\n#   M6 — Attn + MultiHd   (2 + 3)\n#   M7 — Full MS-MHAA     (1 + 2 + 3)  ← proposed method\n# =============================================================================\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 1 — Imports & Configuration\n# ─────────────────────────────────────────────────────────────────────────────\n\nimport os\nimport random\nimport warnings\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom PIL import Image\nfrom sklearn.metrics import (\n    roc_auc_score, recall_score, precision_score,\n    f1_score, accuracy_score, confusion_matrix, roc_curve\n)\nimport matplotlib\nmatplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nimport matplotlib.gridspec as gridspec\nfrom collections import defaultdict\n\nwarnings.filterwarnings('ignore')\n\n# ── Reproducibility ──────────────────────────────────────────────────────────\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Device: {DEVICE}\")\n\n# ── Dataset paths ─────────────────────────────────────────────────────────────\nBASE_PATH = '/kaggle/input/competitions/rsna-pneumonia-detection-challenge'\nTRAIN_IMG_DIR = '/kaggle/input/competitions/rsna-pneumonia-detection-challenge/stage_2_train_images'\nCLASS_CSV = '/kaggle/input/competitions/rsna-pneumonia-detection-challenge/stage_2_detailed_class_info.csv'\nLABEL_CSV = '/kaggle/input/competitions/rsna-pneumonia-detection-challenge/stage_2_train_labels.csv'\n\n# ── Training config ───────────────────────────────────────────────────────────\nCONFIG = {\n    'target_size'         : (128, 128),\n    'batch_size'          : 32,\n    'learning_rate'       : 1e-3,\n    'weight_decay'        : 1e-5,\n    'epochs'              : 60,\n    'scheduler_patience'  : 5,\n    'scheduler_factor'    : 0.5,\n    'early_stop_patience' : 12,\n    'val_split'           : 0.20,\n    'num_heads'           : 4,\n    'ssim_weight'         : 0.7,\n    'mse_weight'          : 0.3,\n    'score_w_pixel'       : 1.0,\n    'score_w_var'         : 3.0,   # head disagreement — most sensitive to OOD\n    'score_w_feat'        : 1.0,\n    'min_recall_threshold': 0.80,  # medical priority: miss fewer pneumonia cases\n    'random_seed'         : SEED,\n}\n\nprint(\"Config loaded.\")\nprint(f\"  Image size : {CONFIG['target_size']}\")\nprint(f\"  Batch size : {CONFIG['batch_size']}\")\nprint(f\"  Epochs     : {CONFIG['epochs']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T01:26:16.610280Z","iopub.execute_input":"2026-03-25T01:26:16.610624Z","iopub.status.idle":"2026-03-25T01:26:26.175118Z","shell.execute_reply.started":"2026-03-25T01:26:16.610599Z","shell.execute_reply":"2026-03-25T01:26:26.174230Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nbase = '/kaggle/input/competitions/rsna-pneumonia-detection-challenge'\nprint(\"Dataset exists:\", os.path.exists(base))\nprint(\"Files:\", os.listdir(base))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T01:26:31.478017Z","iopub.execute_input":"2026-03-25T01:26:31.478930Z","iopub.status.idle":"2026-03-25T01:26:31.486229Z","shell.execute_reply.started":"2026-03-25T01:26:31.478876Z","shell.execute_reply":"2026-03-25T01:26:31.485242Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 2 — DICOM Dataset & Data Splits\n# ─────────────────────────────────────────────────────────────────────────────\n\nclass RSNADicomDataset(Dataset):\n    \"\"\"\n    Reads DICOM chest X-rays from RSNA Pneumonia Challenge.\n\n    label=0 → Normal  (used for training + test negatives)\n    label=1 → Lung Opacity / Pneumonia  (test positives only)\n    'No Lung Opacity / Not Normal' is excluded entirely.\n    \"\"\"\n\n    def __init__(self, patient_ids, img_dir, transform=None, label=0):\n        self.patient_ids = patient_ids\n        self.img_dir     = img_dir\n        self.transform   = transform\n        self.label       = label   # fixed label for this split\n\n    def __len__(self):\n        return len(self.patient_ids)\n\n    def _read_dicom(self, pid):\n        path   = os.path.join(self.img_dir, f'{pid}.dcm')\n        ds     = pydicom.dcmread(path)\n        pixels = ds.pixel_array.astype(np.float32)\n\n        # Apply RescaleSlope / RescaleIntercept if present\n        if hasattr(ds, 'RescaleSlope'):\n            pixels = pixels * float(ds.RescaleSlope)\n        if hasattr(ds, 'RescaleIntercept'):\n            pixels = pixels + float(ds.RescaleIntercept)\n\n        # Clip and normalise to [0, 255] uint8\n        pixels = np.clip(pixels, pixels.min(), pixels.max())\n        pmin, pmax = pixels.min(), pixels.max()\n        if pmax > pmin:\n            pixels = (pixels - pmin) / (pmax - pmin) * 255.0\n        else:\n            pixels = np.zeros_like(pixels)\n\n        img = Image.fromarray(pixels.astype(np.uint8)).convert('L')\n        return img\n\n    def __getitem__(self, idx):\n        pid = self.patient_ids[idx]\n        img = self._read_dicom(pid)\n        if self.transform:\n            img = self.transform(img)\n        return img, self.label\n\n\ndef build_data_splits(class_csv, label_csv, img_dir, config):\n    \"\"\"\n    Returns:\n        train_loader  — Normal images only (no labels used)\n        val_loader    — Held-out Normal images (monitors recon loss)\n        test_loader   — Normal + Lung Opacity (final evaluation)\n        test_labels   — Ground-truth binary labels for test set\n    \"\"\"\n    class_df = pd.read_csv(class_csv)\n    label_df = pd.read_csv(label_csv)\n\n    # Merge to get unique patientId → class mapping\n    merged = class_df[['patientId', 'class']].drop_duplicates('patientId')\n\n    normal_ids  = merged[merged['class'] == 'Normal']['patientId'].tolist()\n    opacity_ids = merged[merged['class'] == 'Lung Opacity']['patientId'].tolist()\n    # 'No Lung Opacity / Not Normal' → excluded\n\n    print(f\"  Normal images      : {len(normal_ids)}\")\n    print(f\"  Lung Opacity images: {len(opacity_ids)}\")\n    print(f\"  Ambiguous (skipped): \"\n          f\"{(merged['class'] == 'No Lung Opacity / Not Normal').sum()}\")\n\n    # Shuffle normal split for reproducibility\n    rng = np.random.default_rng(config['random_seed'])\n    normal_ids = list(rng.permutation(normal_ids))\n\n    val_n    = int(len(normal_ids) * config['val_split'])\n    val_ids  = normal_ids[:val_n]\n    train_ids= normal_ids[val_n:]\n\n    # Test = held-out normal + all opacity\n    test_normal_ids  = val_ids           # reuse val for simplicity (or separate)\n    test_opacity_ids = opacity_ids\n\n    H, W = config['target_size']\n\n    # Training augmentation (mild — preserve normal anatomy)\n    train_tf = transforms.Compose([\n        transforms.Resize((H, W)),\n        transforms.RandomHorizontalFlip(p=0.5),\n        transforms.RandomRotation(degrees=5),\n        transforms.ToTensor(),           # → [0,1] float32\n    ])\n\n    eval_tf = transforms.Compose([\n        transforms.Resize((H, W)),\n        transforms.ToTensor(),\n    ])\n\n    train_ds = RSNADicomDataset(train_ids,        img_dir, train_tf, label=0)\n    val_ds   = RSNADicomDataset(val_ids,          img_dir, eval_tf,  label=0)\n    test_normal_ds  = RSNADicomDataset(test_normal_ids,  img_dir, eval_tf, label=0)\n    test_opacity_ds = RSNADicomDataset(test_opacity_ids, img_dir, eval_tf, label=1)\n\n    from torch.utils.data import ConcatDataset\n    test_ds = ConcatDataset([test_normal_ds, test_opacity_ds])\n\n    test_labels = ([0] * len(test_normal_ids) +\n                   [1] * len(test_opacity_ids))\n\n    BS = config['batch_size']\n    train_loader = DataLoader(train_ds, batch_size=BS, shuffle=True,\n                              num_workers=0, pin_memory=True)\n    val_loader   = DataLoader(val_ds,   batch_size=BS, shuffle=False,\n                              num_workers=0, pin_memory=True)\n    test_loader  = DataLoader(test_ds,  batch_size=BS, shuffle=False,\n                              num_workers=0, pin_memory=True)\n\n    return train_loader, val_loader, test_loader, test_labels\n\n\nprint(\"Loading dataset splits...\")\ntrain_loader, val_loader, test_loader, test_labels = build_data_splits(\n    CLASS_CSV, LABEL_CSV, TRAIN_IMG_DIR, CONFIG\n)\nprint(f\"  Train batches : {len(train_loader)}\")\nprint(f\"  Val batches   : {len(val_loader)}\")\nprint(f\"  Test samples  : {len(test_labels)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T01:26:34.310961Z","iopub.execute_input":"2026-03-25T01:26:34.311756Z","iopub.status.idle":"2026-03-25T01:26:34.500000Z","shell.execute_reply.started":"2026-03-25T01:26:34.311726Z","shell.execute_reply":"2026-03-25T01:26:34.499226Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 3 — Compute Dataset Normalisation Stats (DICOM pixel distribution)\n# ─────────────────────────────────────────────────────────────────────────────\n\ndef compute_mean_std(loader, max_batches=50):\n    \"\"\"Estimate channel mean and std from first N batches.\"\"\"\n    running_mean = 0.0\n    running_var  = 0.0\n    n_pixels     = 0\n\n    for i, (imgs, _) in enumerate(loader):\n        if i >= max_batches:\n            break\n        B, C, H, W = imgs.shape\n        n = B * H * W\n        running_mean += imgs.mean().item() * n\n        running_var  += imgs.var().item()  * n\n        n_pixels     += n\n\n    mean = running_mean / n_pixels\n    std  = (running_var / n_pixels) ** 0.5\n    return mean, std\n\n\nprint(\"Computing normalisation stats from training set (first 50 batches)...\")\nNORM_MEAN, NORM_STD = compute_mean_std(train_loader)\nprint(f\"  Mean : {NORM_MEAN:.4f}\")\nprint(f\"  Std  : {NORM_STD:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T01:26:40.408885Z","iopub.execute_input":"2026-03-25T01:26:40.409788Z","iopub.status.idle":"2026-03-25T01:27:29.688002Z","shell.execute_reply.started":"2026-03-25T01:26:40.409757Z","shell.execute_reply":"2026-03-25T01:27:29.687103Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 4 — SSIM Loss Helper\n# ─────────────────────────────────────────────────────────────────────────────\n\ndef ssim_loss(x, y, window_size=11, eps=1e-8):\n    \"\"\"\n    Structural Similarity loss: 1 - SSIM(x, y).\n    x, y : (B, C, H, W) in [0, 1]\n    \"\"\"\n    C1 = (0.01) ** 2\n    C2 = (0.03) ** 2\n\n    # Gaussian window\n    coords = torch.arange(window_size, dtype=torch.float32,\n                          device=x.device) - window_size // 2\n    g = torch.exp(-(coords ** 2) / (2 * 1.5 ** 2))\n    g /= g.sum()\n    kernel = g.unsqueeze(0) * g.unsqueeze(1)    # (win, win)\n    kernel = kernel.unsqueeze(0).unsqueeze(0)   # (1, 1, win, win)\n    kernel = kernel.expand(x.shape[1], 1, window_size, window_size)\n    pad = window_size // 2\n\n    mu_x  = F.conv2d(x, kernel, padding=pad, groups=x.shape[1])\n    mu_y  = F.conv2d(y, kernel, padding=pad, groups=x.shape[1])\n    mu_x2 = mu_x * mu_x\n    mu_y2 = mu_y * mu_y\n    mu_xy = mu_x * mu_y\n\n    sig_x2  = F.conv2d(x * x, kernel, padding=pad, groups=x.shape[1]) - mu_x2\n    sig_y2  = F.conv2d(y * y, kernel, padding=pad, groups=x.shape[1]) - mu_y2\n    sig_xy  = F.conv2d(x * y, kernel, padding=pad, groups=x.shape[1]) - mu_xy\n\n    num = (2 * mu_xy + C1) * (2 * sig_xy + C2)\n    den = (mu_x2 + mu_y2 + C1) * (sig_x2 + sig_y2 + C2) + eps\n\n    ssim_map = num / den\n    return 1.0 - ssim_map.mean()\n\n\ndef combined_loss(recon, target, ssim_w=0.7, mse_w=0.3):\n    \"\"\"\n    Normalise both to [0,1], then compute SSIM + MSE loss.\n    Normalise per-batch for stable SSIM.\n    \"\"\"\n    # Clip to valid range\n    r = torch.clamp(recon,  0.0, 1.0)\n    t = torch.clamp(target, 0.0, 1.0)\n    return ssim_w * ssim_loss(r, t) + mse_w * F.mse_loss(r, t)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T01:27:35.843055Z","iopub.execute_input":"2026-03-25T01:27:35.843385Z","iopub.status.idle":"2026-03-25T01:27:35.853380Z","shell.execute_reply.started":"2026-03-25T01:27:35.843358Z","shell.execute_reply":"2026-03-25T01:27:35.852655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 5 — Shared Building Blocks\n# ─────────────────────────────────────────────────────────────────────────────\n\ndef encoder_block(in_ch, out_ch, kernel, stride=2):\n    \"\"\"Single encoder stage: Conv → BN → ReLU (spatially downsamples 2×).\"\"\"\n    pad = kernel // 2\n    return nn.Sequential(\n        nn.Conv2d(in_ch, out_ch, kernel, stride=stride, padding=pad,\n                  bias=False),\n        nn.BatchNorm2d(out_ch),\n        nn.ReLU(inplace=True),\n    )\n\n\ndef make_encoder(kernel):\n    \"\"\"3-stage encoder: 1→32→64→128, spatial 128→64→32→16.\"\"\"\n    return nn.Sequential(\n        encoder_block(1,   32,  kernel),\n        encoder_block(32,  64,  kernel),\n        encoder_block(64,  128, kernel),\n    )\n\n\nclass SharedDecoderTrunk(nn.Module):\n    \"\"\"Upsamples bottleneck 16→32→64 (shared across all heads).\"\"\"\n    def __init__(self, in_ch=128):\n        super().__init__()\n        self.up = nn.Sequential(\n            nn.ConvTranspose2d(in_ch, 64, 4, stride=2, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(inplace=True),\n            nn.ConvTranspose2d(64, 32, 4, stride=2, padding=1, bias=False),\n            nn.BatchNorm2d(32),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        return self.up(x)   # (B, 32, 64, 64)\n\n\nclass DecoderHead(nn.Module):\n    \"\"\"Single reconstruction head: 32ch, 64×64 → 1ch, 128×128.\"\"\"\n    def __init__(self):\n        super().__init__()\n        self.head = nn.Sequential(\n            nn.ConvTranspose2d(32, 16, 4, stride=2, padding=1, bias=False),\n            nn.BatchNorm2d(16),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(16, 1, 3, padding=1),\n            nn.Sigmoid(),\n        )\n\n    def forward(self, x):\n        return self.head(x)  # (B, 1, 128, 128)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T01:27:42.044051Z","iopub.execute_input":"2026-03-25T01:27:42.044438Z","iopub.status.idle":"2026-03-25T01:27:42.054677Z","shell.execute_reply.started":"2026-03-25T01:27:42.044413Z","shell.execute_reply":"2026-03-25T01:27:42.053788Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 6 — All 8 Model Architectures\n# ─────────────────────────────────────────────────────────────────────────────\n\n# ── M0 — Baseline: single encoder (3×3), single decoder head ─────────────────\n\nclass M0_Baseline(nn.Module):\n    \"\"\"Standard single-path autoencoder. Zero innovations.\"\"\"\n    def __init__(self):\n        super().__init__()\n        self.encoder = make_encoder(3)\n        self.trunk   = SharedDecoderTrunk(128)\n        self.head    = DecoderHead()\n\n    def encode(self, x):\n        return self.encoder(x)\n\n    def forward(self, x):\n        z    = self.encode(x)\n        feat = self.trunk(z)\n        recon= self.head(feat)\n        return [recon], z   # list for uniform interface\n\n\n# ── M1 — Multi-Scale only (avg fusion, single head) ──────────────────────────\n\nclass M1_MultiScale(nn.Module):\n    \"\"\"Innovation 1: 3 parallel encoders, simple average fusion.\"\"\"\n    def __init__(self):\n        super().__init__()\n        self.enc_fine   = make_encoder(3)\n        self.enc_mid    = make_encoder(5)\n        self.enc_coarse = make_encoder(7)\n        # simple 1×1 projection after average\n        self.project = nn.Sequential(\n            nn.Conv2d(128, 128, 1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(inplace=True),\n        )\n        self.trunk = SharedDecoderTrunk(128)\n        self.head  = DecoderHead()\n\n    def encode(self, x):\n        z = (self.enc_fine(x) + self.enc_mid(x) + self.enc_coarse(x)) / 3.0\n        return self.project(z)\n\n    def forward(self, x):\n        z    = self.encode(x)\n        feat = self.trunk(z)\n        recon= self.head(feat)\n        return [recon], z\n\n\n# ── M2 — Attention only (single encoder + channel self-attention) ─────────────\n\nclass ChannelAttention(nn.Module):\n    \"\"\"Squeeze-and-Excitation style channel attention.\"\"\"\n    def __init__(self, channels=128, reduction=16):\n        super().__init__()\n        self.fc = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(),\n            nn.Linear(channels, channels // reduction),\n            nn.ReLU(inplace=True),\n            nn.Linear(channels // reduction, channels),\n            nn.Sigmoid(),\n        )\n\n    def forward(self, x):\n        w = self.fc(x).view(x.shape[0], x.shape[1], 1, 1)\n        return x * w\n\n\nclass M2_Attention(nn.Module):\n    \"\"\"Innovation 2: single encoder + channel attention refinement.\"\"\"\n    def __init__(self):\n        super().__init__()\n        self.encoder  = make_encoder(3)\n        self.attn     = ChannelAttention(128)\n        self.trunk    = SharedDecoderTrunk(128)\n        self.head     = DecoderHead()\n\n    def encode(self, x):\n        return self.attn(self.encoder(x))\n\n    def forward(self, x):\n        z    = self.encode(x)\n        feat = self.trunk(z)\n        recon= self.head(feat)\n        return [recon], z\n\n\n# ── M3 — Multi-Head only (single encoder, 4 decoder heads) ───────────────────\n\nclass M3_MultiHead(nn.Module):\n    \"\"\"Innovation 3: single encoder + 4 independent decoder heads.\"\"\"\n    def __init__(self, num_heads=4):\n        super().__init__()\n        self.encoder = make_encoder(3)\n        self.trunk   = SharedDecoderTrunk(128)\n        self.heads   = nn.ModuleList([DecoderHead() for _ in range(num_heads)])\n\n    def encode(self, x):\n        return self.encoder(x)\n\n    def forward(self, x):\n        z     = self.encode(x)\n        feat  = self.trunk(z)\n        recons= [h(feat) for h in self.heads]\n        return recons, z\n\n\n# ── M4 — Multi-Scale + Attention (no multi-head) ─────────────────────────────\n\nclass SpatialAttentionFusion(nn.Module):\n    \"\"\"\n    Spatial attention gate across 3 encoder scales.\n    Learns per-location weights: which scale contributes most where.\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        # Compress 3×128 = 384 channels into weights over 3 branches\n        self.gate = nn.Sequential(\n            nn.Conv2d(384, 128, 1, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(128, 3,  1, bias=False),\n        )\n        self.project = nn.Sequential(\n            nn.Conv2d(128, 128, 1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, fine, mid, coarse):\n        stacked = torch.cat([fine, mid, coarse], dim=1)  # (B,384,16,16)\n        w = self.gate(stacked)                            # (B,3,16,16)\n        w = F.softmax(w, dim=1)                           # per-location weights\n        fused = (w[:, 0:1] * fine +\n                 w[:, 1:2] * mid  +\n                 w[:, 2:3] * coarse)\n        return self.project(fused)\n\n\nclass M4_ScaleAttn(nn.Module):\n    \"\"\"Innovations 1+2: multi-scale parallel encoders + spatial attention.\"\"\"\n    def __init__(self):\n        super().__init__()\n        self.enc_fine   = make_encoder(3)\n        self.enc_mid    = make_encoder(5)\n        self.enc_coarse = make_encoder(7)\n        self.fusion     = SpatialAttentionFusion()\n        self.trunk      = SharedDecoderTrunk(128)\n        self.head       = DecoderHead()\n\n    def encode(self, x):\n        return self.fusion(\n            self.enc_fine(x),\n            self.enc_mid(x),\n            self.enc_coarse(x),\n        )\n\n    def forward(self, x):\n        z    = self.encode(x)\n        feat = self.trunk(z)\n        recon= self.head(feat)\n        return [recon], z\n\n\n# ── M5 — Multi-Scale + Multi-Head (avg fusion, no attention gate) ────────────\n\nclass M5_ScaleMultiHead(nn.Module):\n    \"\"\"Innovations 1+3: multi-scale avg fusion + 4 decoder heads.\"\"\"\n    def __init__(self, num_heads=4):\n        super().__init__()\n        self.enc_fine   = make_encoder(3)\n        self.enc_mid    = make_encoder(5)\n        self.enc_coarse = make_encoder(7)\n        self.project = nn.Sequential(\n            nn.Conv2d(128, 128, 1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(inplace=True),\n        )\n        self.trunk = SharedDecoderTrunk(128)\n        self.heads = nn.ModuleList([DecoderHead() for _ in range(num_heads)])\n\n    def encode(self, x):\n        z = (self.enc_fine(x) + self.enc_mid(x) + self.enc_coarse(x)) / 3.0\n        return self.project(z)\n\n    def forward(self, x):\n        z     = self.encode(x)\n        feat  = self.trunk(z)\n        recons= [h(feat) for h in self.heads]\n        return recons, z\n\n\n# ── M6 — Attention + Multi-Head (single encoder, no multi-scale) ─────────────\n\nclass M6_AttnMultiHead(nn.Module):\n    \"\"\"Innovations 2+3: channel attention + 4 decoder heads.\"\"\"\n    def __init__(self, num_heads=4):\n        super().__init__()\n        self.encoder = make_encoder(3)\n        self.attn    = ChannelAttention(128)\n        self.trunk   = SharedDecoderTrunk(128)\n        self.heads   = nn.ModuleList([DecoderHead() for _ in range(num_heads)])\n\n    def encode(self, x):\n        return self.attn(self.encoder(x))\n\n    def forward(self, x):\n        z     = self.encode(x)\n        feat  = self.trunk(z)\n        recons= [h(feat) for h in self.heads]\n        return recons, z\n\n\n# ── M7 — Full MS-MHAA (all 3 innovations combined) ───────────────────────────\n\nclass M7_MSMHAA(nn.Module):\n    \"\"\"\n    Full proposed method:\n    Innovation 1 — Multi-Scale parallel encoders (3×3, 5×5, 7×7)\n    Innovation 2 — Spatial Attention Fusion Gate\n    Innovation 3 — Multi-Head Anomaly Decoder (4 heads)\n    \"\"\"\n    def __init__(self, num_heads=4):\n        super().__init__()\n        self.enc_fine   = make_encoder(3)\n        self.enc_mid    = make_encoder(5)\n        self.enc_coarse = make_encoder(7)\n        self.fusion     = SpatialAttentionFusion()\n        self.trunk      = SharedDecoderTrunk(128)\n        self.heads      = nn.ModuleList([DecoderHead() for _ in range(num_heads)])\n\n    def encode(self, x):\n        return self.fusion(\n            self.enc_fine(x),\n            self.enc_mid(x),\n            self.enc_coarse(x),\n        )\n\n    def forward(self, x):\n        z     = self.encode(x)\n        feat  = self.trunk(z)\n        recons= [h(feat) for h in self.heads]\n        return recons, z\n\n\n# ── Model registry ────────────────────────────────────────────────────────────\nMODEL_REGISTRY = {\n    #'M0_Baseline'     : M0_Baseline,\n    #'M1_MultiScale'   : M1_MultiScale,\n    # 'M2_Attention'    : M2_Attention,\n    # 'M3_MultiHead'    : M3_MultiHead,\n    # 'M4_ScaleAttn'    : M4_ScaleAttn,\n    # 'M5_ScaleMultiHd' : M5_ScaleMultiHead,\n    'M6_AttnMultiHd'  : M6_AttnMultiHead,\n    'M7_MSMHAA'       : M7_MSMHAA,\n}\n\n# Count parameters for each model\nprint(\"\\nModel parameter counts:\")\nfor name, cls in MODEL_REGISTRY.items():\n    m = cls().to('cpu')\n    n = sum(p.numel() for p in m.parameters())\n    print(f\"  {name:<20}: {n:>10,} params\")\n    del m\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T01:28:51.057697Z","iopub.execute_input":"2026-03-25T01:28:51.058036Z","iopub.status.idle":"2026-03-25T01:28:51.109965Z","shell.execute_reply.started":"2026-03-25T01:28:51.058007Z","shell.execute_reply":"2026-03-25T01:28:51.109027Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 7 — Generic Training Loop (works for all 8 models)\n# ─────────────────────────────────────────────────────────────────────────────\n\ndef train_model(model, train_loader, val_loader, config, model_name, resume=True):\n    model = model.to(DEVICE)\n    optim = torch.optim.Adam(\n        model.parameters(),\n        lr=config['learning_rate'],\n        weight_decay=config['weight_decay'],\n    )\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n        optim,\n        patience=config['scheduler_patience'],\n        factor=config['scheduler_factor'],\n    )\n\n    ssim_w       = config['ssim_weight']\n    mse_w        = config['mse_weight']\n    best_val     = float('inf')\n    patience_ctr = 0\n    start_epoch  = 0\n    history      = {'train_loss': [], 'val_loss': []}\n\n    # ── Resume from checkpoint if exists ─────────────────────────\n    ckpt_path      = f'/kaggle/working/{model_name}_best.pth'\n    history_path   = f'/kaggle/working/{model_name}_history.npy'\n    meta_path      = f'/kaggle/working/{model_name}_meta.npy'\n\n    if resume and os.path.exists(ckpt_path):\n        print(f\"  ↳ Resuming {model_name} from checkpoint...\")\n        model.load_state_dict(torch.load(ckpt_path, map_location=DEVICE))\n\n        if os.path.exists(history_path):\n            saved = np.load(history_path, allow_pickle=True).item()\n            history     = saved\n            start_epoch = len(history['train_loss'])\n            print(f\"    Resumed from epoch {start_epoch}\")\n\n        if os.path.exists(meta_path):\n            meta         = np.load(meta_path, allow_pickle=True).item()\n            best_val     = meta['best_val']\n            patience_ctr = meta['patience_ctr']\n            # restore LR to last known value\n            for pg in optim.param_groups:\n                pg['lr'] = meta['lr']\n            print(f\"    best_val={best_val:.4f}  patience={patience_ctr}  lr={meta['lr']:.6f}\")\n    else:\n        print(f\"\\n  Training {model_name} from scratch | {config['epochs']} epochs max\")\n\n    print(f\"  {'Epoch':<8} {'Train':>8} {'Val':>8} {'LR':>10} {'Best Val':>10} {'Patience':>9}\")\n    print(f\"  {'-'*60}\")\n\n    for epoch in range(start_epoch, config['epochs']):\n        # ── Train ────────────────────────────────────────────────\n        model.train()\n        train_loss = 0.0\n        for imgs, _ in train_loader:\n            imgs = imgs.to(DEVICE)\n            recons, _ = model(imgs)\n            loss = sum(combined_loss(r, imgs, ssim_w, mse_w)\n                       for r in recons) / len(recons)\n            optim.zero_grad()\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            optim.step()\n            train_loss += loss.item()\n        train_loss /= len(train_loader)\n\n        # ── Validate ─────────────────────────────────────────────\n        model.eval()\n        val_loss = 0.0\n        with torch.no_grad():\n            for imgs, _ in val_loader:\n                imgs = imgs.to(DEVICE)\n                recons, _ = model(imgs)\n                loss = sum(combined_loss(r, imgs, ssim_w, mse_w)\n                           for r in recons) / len(recons)\n                val_loss += loss.item()\n        val_loss /= len(val_loader)\n\n        history['train_loss'].append(train_loss)\n        history['val_loss'].append(val_loss)\n        scheduler.step(val_loss)\n        lr_now = optim.param_groups[0]['lr']\n\n        print(f\"  {epoch+1:<8} {train_loss:>8.4f} {val_loss:>8.4f} \"\n              f\"{lr_now:>10.6f} {best_val:>10.4f} {patience_ctr:>9}\")\n\n        # ── Save history + meta every epoch ──────────────────────\n        np.save(history_path, history)\n        np.save(meta_path, {\n            'best_val'    : best_val,\n            'patience_ctr': patience_ctr,\n            'lr'          : lr_now,\n            'last_epoch'  : epoch + 1,\n        })\n\n        # ── Early stopping ────────────────────────────────────────\n        if val_loss < best_val - 1e-5:\n            best_val     = val_loss\n            patience_ctr = 0\n            torch.save(model.state_dict(), ckpt_path)\n            print(f\"           ↑ new best saved\")\n        else:\n            patience_ctr += 1\n            if patience_ctr >= config['early_stop_patience']:\n                print(f\"\\n  Early stop at epoch {epoch+1}\")\n                break\n\n    print(f\"\\n  ✓ {model_name} complete — best val loss: {best_val:.4f}\\n\")\n    model.load_state_dict(torch.load(ckpt_path, map_location=DEVICE))\n    return model, history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T01:28:04.630468Z","iopub.execute_input":"2026-03-25T01:28:04.631125Z","iopub.status.idle":"2026-03-25T01:28:04.645508Z","shell.execute_reply.started":"2026-03-25T01:28:04.631096Z","shell.execute_reply":"2026-03-25T01:28:04.644737Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 8 — Anomaly Scoring\n# ─────────────────────────────────────────────────────────────────────────────\n\ndef compute_anomaly_scores(model, loader, config):\n    \"\"\"\n    Composite anomaly score:\n        score = w_pixel * pixel_mse\n              + w_var   * head_var    (only if multi-head)\n              + w_feat  * feat_dist\n\n    head_var = variance across head outputs → high when OOD\n    feat_dist= bottleneck distance between input and its reconstruction\n    \"\"\"\n    model.eval()\n    scores  = []\n    labels  = []\n\n    w_pixel = config['score_w_pixel']\n    w_var   = config['score_w_var']\n    w_feat  = config['score_w_feat']\n\n    with torch.no_grad():\n        for imgs, lbls in loader:\n            imgs = imgs.to(DEVICE)\n\n            recons, z_orig = model(imgs)\n            mean_recon = torch.stack(recons, dim=0).mean(dim=0)  # (B,1,H,W)\n\n            # Pixel MSE\n            pixel_mse = ((imgs - mean_recon) ** 2).mean(dim=[1, 2, 3])\n\n            # Head disagreement (variance across heads)\n            if len(recons) > 1:\n                stacked  = torch.stack(recons, dim=0)  # (heads,B,1,H,W)\n                head_var = stacked.var(dim=0).mean(dim=[1, 2, 3])\n            else:\n                head_var = torch.zeros_like(pixel_mse)\n\n            # Feature-space distance\n            _, z_recon = model(mean_recon)\n            feat_dist  = ((z_orig - z_recon) ** 2).mean(dim=[1, 2, 3])\n\n            score = (w_pixel * pixel_mse +\n                     w_var   * head_var  +\n                     w_feat  * feat_dist)\n\n            scores.extend(score.cpu().numpy().tolist())\n            labels.extend(lbls.numpy().tolist())\n\n    scores = np.array(scores)\n    labels = np.array(labels)\n    return scores, labels\n\n\ndef auto_correct_scores(scores, labels):\n    \"\"\"\n    If normal mean > anomaly mean (inverted), negate scores.\n    Handles rare cases where reconstruction is better for anomalies.\n    \"\"\"\n    normal_mean  = scores[labels == 0].mean()\n    anomaly_mean = scores[labels == 1].mean()\n    if normal_mean > anomaly_mean:\n        print(\"    [Auto-correct] Score direction inverted — negating scores.\")\n        scores = -scores\n    return scores\n\n\ndef select_threshold(scores, labels, min_recall=0.80):\n    \"\"\"\n    Medical priority threshold: find all thresholds where recall >= min_recall,\n    then pick the one with lowest FPR (fewer false alarms).\n    Falls back to standard 0.5 percentile split if none qualify.\n    \"\"\"\n    thresholds = np.percentile(scores, np.arange(1, 100))\n    best_thresh = thresholds[len(thresholds) // 2]\n    best_fpr    = 1.0\n\n    for t in thresholds:\n        preds   = (scores >= t).astype(int)\n        if labels.sum() == 0:\n            break\n        recall  = recall_score(labels, preds, zero_division=0)\n        tn = ((preds == 0) & (labels == 0)).sum()\n        fp = ((preds == 1) & (labels == 0)).sum()\n        fpr = fp / (fp + tn + 1e-8)\n\n        if recall >= min_recall and fpr < best_fpr:\n            best_fpr    = fpr\n            best_thresh = t\n\n    return best_thresh","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T01:28:09.124480Z","iopub.execute_input":"2026-03-25T01:28:09.124814Z","iopub.status.idle":"2026-03-25T01:28:09.136411Z","shell.execute_reply.started":"2026-03-25T01:28:09.124787Z","shell.execute_reply":"2026-03-25T01:28:09.135707Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 9 — Evaluation\n# ─────────────────────────────────────────────────────────────────────────────\n\ndef evaluate(scores, labels, model_name, config):\n    \"\"\"Returns a dict of all metrics for this model.\"\"\"\n    scores = auto_correct_scores(scores, labels)\n    thresh = select_threshold(scores, labels,\n                              min_recall=config['min_recall_threshold'])\n    preds  = (scores >= thresh).astype(int)\n\n    auc       = roc_auc_score(labels, scores)\n    recall    = recall_score(labels, preds,    zero_division=0)\n    precision = precision_score(labels, preds, zero_division=0)\n    f1        = f1_score(labels, preds,        zero_division=0)\n    acc       = accuracy_score(labels, preds)\n    cm        = confusion_matrix(labels, preds)\n\n    tn, fp, fn, tp = cm.ravel()\n    specificity = tn / (tn + fp + 1e-8)\n\n    return {\n        'model'      : model_name,\n        'AUC'        : round(auc,       4),\n        'Recall'     : round(recall,    4),\n        'Precision'  : round(precision, 4),\n        'F1'         : round(f1,        4),\n        'Accuracy'   : round(acc,       4),\n        'Specificity': round(specificity,4),\n        'TP'         : int(tp),\n        'FP'         : int(fp),\n        'TN'         : int(tn),\n        'FN'         : int(fn),\n        'threshold'  : round(thresh,    6),\n        'scores'     : scores,\n        'labels'     : labels,\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T01:29:01.195707Z","iopub.execute_input":"2026-03-25T01:29:01.196037Z","iopub.status.idle":"2026-03-25T01:29:01.203460Z","shell.execute_reply.started":"2026-03-25T01:29:01.196000Z","shell.execute_reply":"2026-03-25T01:29:01.202575Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 10 — Full Ablation Runner\n# ─────────────────────────────────────────────────────────────────────────────\n# If M0 already finished, it skips straight to M1, M2, etc.\nall_results   = {}\nall_histories = {}\nfor model_name, ModelClass in MODEL_REGISTRY.items():\n    print(f\"\\n{'='*60}\")\n    print(f\"  Running: {model_name}\")\n    print(f\"{'='*60}\")\n\n    # Skip fully completed models\n    meta_path = f'/kaggle/working/{model_name}_meta.npy'\n    if os.path.exists(meta_path):\n        meta = np.load(meta_path, allow_pickle=True).item()\n        last = meta.get('last_epoch', 0)\n        if last >= CONFIG['epochs']:\n            print(f\"  ✓ Already complete ({last} epochs) — loading results...\")\n            # just load scores for already-trained model\n            model = ModelClass().to(DEVICE)\n            model.load_state_dict(torch.load(\n                f'/kaggle/working/{model_name}_best.pth', map_location=DEVICE))\n            scores, labels = compute_anomaly_scores(model, test_loader, CONFIG)\n            all_results[model_name]   = evaluate(scores, labels, model_name, CONFIG)\n            all_histories[model_name] = np.load(\n                f'/kaggle/working/{model_name}_history.npy',\n                allow_pickle=True).item()\n            del model\n            torch.cuda.empty_cache()\n            continue\n\n    model = ModelClass().to(DEVICE)\n    trained_model, history = train_model(\n        model, train_loader, val_loader, CONFIG, model_name, resume=True\n    )\n    all_histories[model_name] = history\n    scores, labels = compute_anomaly_scores(trained_model, test_loader, CONFIG)\n    all_results[model_name]   = evaluate(scores, labels, model_name, CONFIG)\n    del trained_model, model\n    torch.cuda.empty_cache()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T05:31:03.943934Z","iopub.execute_input":"2026-03-25T05:31:03.944301Z","iopub.status.idle":"2026-03-25T07:23:04.546819Z","shell.execute_reply.started":"2026-03-25T05:31:03.944273Z","shell.execute_reply":"2026-03-25T07:23:04.546202Z"},"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import zipfile\n\nzip_path = \"/kaggle/working/M6_files.zip\"\n\nwith zipfile.ZipFile(zip_path, 'w') as z:\n    z.write(\"/kaggle/working/M7_MSMHAA_best.pth\")\n    z.write(\"/kaggle/working/M7_MSMHAA_history.npy\")\n    z.write(\"/kaggle/working/M7_MSMHAA_meta.npy\")\n\nprint(\"Zip created!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython.display import FileLink\nFileLink(\"/kaggle/working/M6_files.zip\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T05:30:44.006944Z","iopub.status.idle":"2026-03-25T05:30:44.007330Z","shell.execute_reply.started":"2026-03-25T05:30:44.007116Z","shell.execute_reply":"2026-03-25T05:30:44.007139Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nprint(os.listdir(\"/kaggle/working\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 11 — Results Table\n# ─────────────────────────────────────────────────────────────────────────────\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"ABLATION STUDY RESULTS\")\nprint(\"=\"*80)\n\nrows = []\nfor name, r in all_results.items():\n    rows.append({\n        'Model'      : r['model'],\n        'AUC'        : r['AUC'],\n        'Recall'     : r['Recall'],\n        'Precision'  : r['Precision'],\n        'F1'         : r['F1'],\n        'Accuracy'   : r['Accuracy'],\n        'Specificity': r['Specificity'],\n        'TP'         : r['TP'],\n        'FP'         : r['FP'],\n        'TN'         : r['TN'],\n        'FN'         : r['FN'],\n    })\n\nresults_df = pd.DataFrame(rows)\nresults_df = results_df.sort_values('AUC', ascending=False).reset_index(drop=True)\nprint(results_df.to_string(index=False))\n\nresults_df.to_csv('/kaggle/working/ablation_results.csv', index=False)\nprint(\"\\nSaved: /kaggle/working/ablation_results.csv\")\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-03-24T05:04:03.089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 12 — Visualisations\n# ─────────────────────────────────────────────────────────────────────────────\n\nMODEL_COLORS = {\n    'M0_Baseline'     : '#888780',\n    'M1_MultiScale'   : '#7F77DD',\n    'M2_Attention'    : '#1D9E75',\n    'M3_MultiHead'    : '#D85A30',\n    'M4_ScaleAttn'    : '#378ADD',\n    'M5_ScaleMultiHd' : '#BA7517',\n    'M6_AttnMultiHd'  : '#D4537E',\n    'M7_MSMHAA'       : '#E24B4A',\n}\n\nfig = plt.figure(figsize=(20, 24))\nfig.patch.set_facecolor('#FFFFFF')\ngs  = gridspec.GridSpec(4, 2, figure=fig, hspace=0.45, wspace=0.35)\n\n# ── 1. AUC bar chart ──────────────────────────────────────────────────────────\nax1 = fig.add_subplot(gs[0, 0])\nnames = [r['model'] for r in all_results.values()]\naucs  = [r['AUC']   for r in all_results.values()]\ncolors= [MODEL_COLORS.get(n, '#888') for n in names]\nshort = [n.split('_')[0] for n in names]\n\nbars = ax1.barh(short, aucs, color=colors, height=0.6)\nax1.set_xlabel('ROC-AUC', fontsize=11)\nax1.set_title('AUC comparison across ablation variants', fontsize=12, fontweight='500')\nax1.axvline(x=0.7, color='gray', linestyle='--', alpha=0.5, linewidth=0.8)\nax1.set_xlim(0.4, 1.0)\nfor bar, val in zip(bars, aucs):\n    ax1.text(bar.get_width() + 0.005, bar.get_y() + bar.get_height()/2,\n             f'{val:.3f}', va='center', fontsize=9)\n\n# ── 2. Recall bar chart ───────────────────────────────────────────────────────\nax2 = fig.add_subplot(gs[0, 1])\nrecalls = [r['Recall'] for r in all_results.values()]\nbars2 = ax2.barh(short, recalls, color=colors, height=0.6)\nax2.set_xlabel('Recall (sensitivity)', fontsize=11)\nax2.set_title('Recall — minimising missed pneumonia (priority)', fontsize=12, fontweight='500')\nax2.axvline(x=0.80, color='red', linestyle='--', alpha=0.7, linewidth=0.8,\n            label='Target recall 0.80')\nax2.set_xlim(0.0, 1.1)\nax2.legend(fontsize=9)\nfor bar, val in zip(bars2, recalls):\n    ax2.text(bar.get_width() + 0.005, bar.get_y() + bar.get_height()/2,\n             f'{val:.3f}', va='center', fontsize=9)\n\n# ── 3. ROC curves ─────────────────────────────────────────────────────────────\nax3 = fig.add_subplot(gs[1, :])\nax3.plot([0, 1], [0, 1], 'k--', alpha=0.4, linewidth=0.8)\nfor name, r in all_results.items():\n    fpr_, tpr_, _ = roc_curve(r['labels'], r['scores'])\n    lw = 2.5 if name == 'M7_MSMHAA' else 1.2\n    ax3.plot(fpr_, tpr_,\n             color=MODEL_COLORS.get(name, '#888'),\n             linewidth=lw,\n             label=f\"{name.split('_')[0]} (AUC={r['AUC']:.3f})\")\nax3.set_xlabel('False Positive Rate', fontsize=11)\nax3.set_ylabel('True Positive Rate', fontsize=11)\nax3.set_title('ROC curves — all ablation variants', fontsize=12, fontweight='500')\nax3.legend(fontsize=9, loc='lower right', ncol=2)\nax3.grid(True, alpha=0.3)\n\n# ── 4. Score distribution for M7 (proposed method) ────────────────────────────\nax4 = fig.add_subplot(gs[2, 0])\nr7      = all_results.get('M7_MSMHAA', list(all_results.values())[-1])\nscores7 = r7['scores']\nlabels7 = r7['labels']\nax4.hist(scores7[labels7 == 0], bins=50, alpha=0.65, color='#378ADD',\n         label='Normal', density=True)\nax4.hist(scores7[labels7 == 1], bins=50, alpha=0.65, color='#E24B4A',\n         label='Pneumonia', density=True)\nax4.axvline(r7['threshold'], color='black', linestyle='--',\n            linewidth=1.2, label=f\"Threshold={r7['threshold']:.4f}\")\nax4.set_xlabel('Anomaly score', fontsize=11)\nax4.set_ylabel('Density', fontsize=11)\nax4.set_title('M7 MS-MHAA — anomaly score distribution', fontsize=12, fontweight='500')\nax4.legend(fontsize=9)\n\n# ── 5. Score distribution for M0 (baseline) ───────────────────────────────────\nax5 = fig.add_subplot(gs[2, 1])\nr0      = all_results.get('M0_Baseline', list(all_results.values())[0])\nscores0 = r0['scores']\nlabels0 = r0['labels']\nax5.hist(scores0[labels0 == 0], bins=50, alpha=0.65, color='#378ADD',\n         label='Normal', density=True)\nax5.hist(scores0[labels0 == 1], bins=50, alpha=0.65, color='#E24B4A',\n         label='Pneumonia', density=True)\nax5.axvline(r0['threshold'], color='black', linestyle='--',\n            linewidth=1.2, label=f\"Threshold={r0['threshold']:.4f}\")\nax5.set_xlabel('Anomaly score', fontsize=11)\nax5.set_ylabel('Density', fontsize=11)\nax5.set_title('M0 Baseline — anomaly score distribution', fontsize=12, fontweight='500')\nax5.legend(fontsize=9)\n\n# ── 6. Training loss curves ───────────────────────────────────────────────────\nax6 = fig.add_subplot(gs[3, :])\nfor name, h in all_histories.items():\n    lw = 2.5 if name == 'M7_MSMHAA' else 1.0\n    ax6.plot(h['val_loss'],\n             color=MODEL_COLORS.get(name, '#888'),\n             linewidth=lw,\n             label=f\"{name.split('_')[0]}\")\nax6.set_xlabel('Epoch', fontsize=11)\nax6.set_ylabel('Validation loss (SSIM + MSE)', fontsize=11)\nax6.set_title('Validation loss curves — all models', fontsize=12, fontweight='500')\nax6.legend(fontsize=9, ncol=4)\nax6.grid(True, alpha=0.3)\n\nplt.suptitle(\n    'MS-MHAA Ablation Study — RSNA Pneumonia Detection\\n'\n    'Self-supervised anomaly detection (zero labels during training)',\n    fontsize=14, fontweight='500', y=1.01\n)\n\nplt.savefig('/kaggle/working/ablation_results.png',\n            dpi=150, bbox_inches='tight', facecolor='white')\nplt.show()\nprint(\"Saved: /kaggle/working/ablation_results.png\")\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-03-24T05:04:03.089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 13 — Confusion Matrix Grid (all 8 models)\n# ─────────────────────────────────────────────────────────────────────────────\n\nfig2, axes = plt.subplots(2, 4, figsize=(20, 10))\nfig2.patch.set_facecolor('#FFFFFF')\n\nfor ax, (name, r) in zip(axes.flat, all_results.items()):\n    cm_vals = np.array([[r['TN'], r['FP']],\n                         [r['FN'], r['TP']]])\n    im = ax.imshow(cm_vals, interpolation='nearest',\n                   cmap='Blues', vmin=0)\n    ax.set_title(f\"{name.split('_')[0]}\\nAUC={r['AUC']:.3f}  Recall={r['Recall']:.3f}\",\n                 fontsize=10)\n    ax.set_xlabel('Predicted', fontsize=9)\n    ax.set_ylabel('Actual', fontsize=9)\n    ax.set_xticks([0, 1]); ax.set_xticklabels(['Normal', 'Pneumonia'], fontsize=8)\n    ax.set_yticks([0, 1]); ax.set_yticklabels(['Normal', 'Pneumonia'], fontsize=8)\n    thresh_cm = cm_vals.max() / 2.0\n    for i in range(2):\n        for j in range(2):\n            color = 'white' if cm_vals[i, j] > thresh_cm else 'black'\n            ax.text(j, i, str(cm_vals[i, j]), ha='center', va='center',\n                    color=color, fontsize=11, fontweight='500')\n\nplt.suptitle('Confusion matrices — ablation variants\\n(test set: Normal + Pneumonia)',\n             fontsize=13, fontweight='500')\nplt.tight_layout()\nplt.savefig('/kaggle/working/confusion_matrices.png',\n            dpi=150, bbox_inches='tight', facecolor='white')\nplt.show()\nprint(\"Saved: /kaggle/working/confusion_matrices.png\")\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-03-24T05:04:03.089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 14 — Sample Reconstructions (M7 on test images)\n# ─────────────────────────────────────────────────────────────────────────────\n\ndef show_reconstructions(model_class, model_name, test_loader,\n                         config, n_samples=4):\n    \"\"\"\n    Loads best weights and shows:\n        Row 1 — Normal images\n        Row 2 — Their reconstructions\n        Row 3 — Error heatmap |input - recon|\n\n        Row 4 — Pneumonia images\n        Row 5 — Their reconstructions\n        Row 6 — Error heatmap\n    \"\"\"\n    model = model_class().to(DEVICE)\n    ckpt  = f'/kaggle/working/{model_name}_best.pth'\n    if not os.path.exists(ckpt):\n        print(f\"Checkpoint not found: {ckpt}\")\n        return\n    model.load_state_dict(torch.load(ckpt, map_location=DEVICE))\n    model.eval()\n\n    normal_imgs    = []\n    pneumonia_imgs = []\n\n    for imgs, lbls in test_loader:\n        for img, lbl in zip(imgs, lbls):\n            if lbl == 0 and len(normal_imgs)    < n_samples:\n                normal_imgs.append(img)\n            if lbl == 1 and len(pneumonia_imgs) < n_samples:\n                pneumonia_imgs.append(img)\n        if (len(normal_imgs) >= n_samples and\n                len(pneumonia_imgs) >= n_samples):\n            break\n\n    fig3, axes = plt.subplots(6, n_samples, figsize=(n_samples * 4, 22))\n    fig3.patch.set_facecolor('#FFFFFF')\n    row_titles = ['Normal input', 'Reconstruction', 'Error map',\n                  'Pneumonia input', 'Reconstruction', 'Error map']\n    cmaps = ['gray', 'gray', 'hot', 'gray', 'gray', 'hot']\n\n    for col in range(n_samples):\n        for row_grp, imgs_grp in enumerate([normal_imgs, pneumonia_imgs]):\n            img = imgs_grp[col].unsqueeze(0).to(DEVICE)\n            with torch.no_grad():\n                recons, _ = model(img)\n                recon = torch.stack(recons, 0).mean(0)\n            img_np   = img.squeeze().cpu().numpy()\n            recon_np = recon.squeeze().cpu().numpy()\n            error_np = np.abs(img_np - recon_np)\n\n            offset = row_grp * 3\n            axes[offset + 0, col].imshow(img_np,   cmap='gray')\n            axes[offset + 1, col].imshow(recon_np, cmap='gray')\n            im = axes[offset + 2, col].imshow(error_np, cmap='hot')\n            plt.colorbar(im, ax=axes[offset + 2, col], fraction=0.046, pad=0.04)\n\n            for r in range(3):\n                axes[offset + r, col].axis('off')\n                if col == 0:\n                    axes[offset + r, col].set_title(\n                        row_titles[offset + r], fontsize=10, fontweight='500', loc='left'\n                    )\n\n    plt.suptitle(f'{model_name} — Reconstruction quality on test set',\n                 fontsize=13, fontweight='500')\n    plt.tight_layout()\n    plt.savefig(f'/kaggle/working/{model_name}_reconstructions.png',\n                dpi=120, bbox_inches='tight', facecolor='white')\n    plt.show()\n    del model\n    torch.cuda.empty_cache()\n\n\n# Run for proposed method (M7) and baseline (M0)\nshow_reconstructions(M7_MSMHAA,   'M7_MSMHAA',   test_loader, CONFIG)\nshow_reconstructions(M0_Baseline, 'M0_Baseline',  test_loader, CONFIG)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-03-24T05:04:03.089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 15 — Ablation Summary Table (publication-ready)\n# ─────────────────────────────────────────────────────────────────────────────\n\n# Innovation flags per model (for the summary)\nINNO_FLAGS = {\n    'M0_Baseline'     : ('✗', '✗', '✗'),\n    'M1_MultiScale'   : ('✓', '✗', '✗'),\n    'M2_Attention'    : ('✗', '✓', '✗'),\n    'M3_MultiHead'    : ('✗', '✗', '✓'),\n    'M4_ScaleAttn'    : ('✓', '✓', '✗'),\n    'M5_ScaleMultiHd' : ('✓', '✗', '✓'),\n    'M6_AttnMultiHd'  : ('✗', '✓', '✓'),\n    'M7_MSMHAA'       : ('✓', '✓', '✓'),\n}\n\nprint(\"\\n\" + \"=\"*100)\nprint(\"PUBLICATION-READY ABLATION TABLE\")\nprint(\"=\"*100)\nheader = (f\"{'Model':<20} | {'Multi-Scale':^12} | {'Attention':^9} | \"\n          f\"{'Multi-Head':^11} | {'AUC':^7} | {'Recall':^7} | \"\n          f\"{'Precision':^10} | {'F1':^7} | {'Acc':^7}\")\nprint(header)\nprint(\"-\" * 100)\n\nfor name in MODEL_REGISTRY.keys():\n    r  = all_results[name]\n    ms, at, mh = INNO_FLAGS[name]\n    mark = \" ← proposed\" if name == 'M7_MSMHAA' else \"\"\n    print(f\"{name:<20} | {ms:^12} | {at:^9} | {mh:^11} | \"\n          f\"{r['AUC']:^7.4f} | {r['Recall']:^7.4f} | \"\n          f\"{r['Precision']:^10.4f} | {r['F1']:^7.4f} | \"\n          f\"{r['Accuracy']:^7.4f}{mark}\")\n\nprint(\"=\"*100)\nprint(\"\\nKey: Multi-Scale=Innovation1, Attention=Innovation2, Multi-Head=Innovation3\")\nprint(\"Training set: NORMAL images only — no disease labels used during training.\")\nprint(\"Test set    : Normal + Lung Opacity (Pneumonia) from RSNA Pneumonia Challenge.\")\n\nprint(\"\\n\\nAll outputs saved to /kaggle/working/:\")\nprint(\"  ablation_results.csv\")\nprint(\"  ablation_results.png\")\nprint(\"  confusion_matrices.png\")\nprint(\"  M7_MSMHAA_reconstructions.png\")\nprint(\"  M0_Baseline_reconstructions.png\")\nprint(\"  [model]_best.pth  (one per model)\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"execution_failed":"2026-03-24T05:04:03.089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}