{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.13"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"6fde78f3","cell_type":"markdown","source":"# RSNA Lumbar Spine — Attention Ensemble (EfficientNet+CBAM & VGG16+SE)\n\nA PyTorch pipeline mirroring the IQ-OTH/NCCD attention-ensemble pattern, adapted for the RSNA 2024 Lumbar Spine Degenerative Classification challenge.\n\n**What this notebook does**\n1. Loads RSNA DICOMs with VOI LUT preprocessing.\n2. Crops anatomical ROIs around `train_label_coordinates` for sharper supervision.\n3. Trains two attention-augmented CNNs (full fine-tune) on **all series combined** for 3-class severity (`normal_mild`/`moderate`/`severe`):\n   - **EfficientNetV2-S + CBAM** (channel + spatial attention)\n   - **VGG16 + Squeeze-and-Excitation** (channel attention)\n4. Builds two ensembles: **weighted-average** and **unified feature-fusion (attention)**.\n5. Reports classification report, confusion matrix, ROC/PR curves, accuracy/loss curves, and a final model-comparison chart.","metadata":{}},{"id":"5356c471","cell_type":"markdown","source":"## 1. Setup","metadata":{}},{"id":"aef9a0e9","cell_type":"code","source":"import os\nimport random\nimport warnings\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm.auto import tqdm\nfrom copy import deepcopy\n\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.models as tvm\nimport torchvision.transforms as T\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (classification_report, confusion_matrix,\n                             roc_curve, auc, precision_recall_curve)\nfrom sklearn.utils.class_weight import compute_class_weight\n\nwarnings.filterwarnings('ignore')\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nNUM_GPUS = torch.cuda.device_count() if torch.cuda.is_available() else 1\nprint('Device:', device, f'| GPUs: {NUM_GPUS}')\n","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:44:51.719671Z","iopub.execute_input":"2026-06-10T07:44:51.720105Z","iopub.status.idle":"2026-06-10T07:44:51.731251Z","shell.execute_reply.started":"2026-06-10T07:44:51.720078Z","shell.execute_reply":"2026-06-10T07:44:51.730386Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"6fa3f7ce-625b-4ba9-b7db-557a78e49959","cell_type":"code","source":"import torch\n\nprint(\"CUDA available:\", torch.cuda.is_available())\nprint(\"GPU count:\", torch.cuda.device_count())\n\nfor i in range(torch.cuda.device_count()):\n    print(f\"GPU {i}: {torch.cuda.get_device_name(i)}\")","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:44:51.732299Z","iopub.execute_input":"2026-06-10T07:44:51.732478Z","iopub.status.idle":"2026-06-10T07:44:51.742973Z","shell.execute_reply.started":"2026-06-10T07:44:51.73246Z","shell.execute_reply":"2026-06-10T07:44:51.742267Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"21169612","cell_type":"code","source":"# Configuration\nDATA_DIR   = '/kaggle/input/competitions/rsna-2024-lumbar-spine-degenerative-classification'\nTRAIN_IMG_DIR = os.path.join(DATA_DIR, 'train_images')\n\nIMG_SIZE     = 224           # input size to the CNNs\nROI_SIZE     = 128           # pixel-radius around the labeled (x, y) before resize\nBATCH_SIZE   = 32\nNUM_EPOCHS   = 10            # raise if you have GPU time\nLR_BACKBONE  = 1e-4\nLR_HEAD      = 1e-3\nLABEL_SMOOTH = 0.1           # label smoothing in CE loss\nUSE_AMP      = True          # mixed precision (faster, same accuracy)\nUSE_TTA      = True          # horizontal-flip test-time augmentation at eval\nNUM_CLASSES  = 3\nNUM_WORKERS  = 2\nCLASS_NAMES  = ['normal_mild', 'moderate', 'severe']\nLABEL_MAP    = {c: i for i, c in enumerate(CLASS_NAMES)}","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:44:51.743832Z","iopub.execute_input":"2026-06-10T07:44:51.744241Z","iopub.status.idle":"2026-06-10T07:44:51.754472Z","shell.execute_reply.started":"2026-06-10T07:44:51.744213Z","shell.execute_reply":"2026-06-10T07:44:51.753761Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"58e4b013","cell_type":"markdown","source":"## 2. Load and shape the labels\n\nWe reshape `train.csv` from wide to long (one row per `study_id × condition × level`), then join with `train_label_coordinates.csv` to get the (x, y) point for ROI cropping and with `train_series_descriptions.csv` to know which MRI sequence each series is.","metadata":{}},{"id":"f42ed3d4","cell_type":"code","source":"train       = pd.read_csv(f'{DATA_DIR}/train.csv')\ncoords      = pd.read_csv(f'{DATA_DIR}/train_label_coordinates.csv')\ntrain_desc  = pd.read_csv(f'{DATA_DIR}/train_series_descriptions.csv')\nprint('train:', train.shape, 'coords:', coords.shape, 'series:', train_desc.shape)","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:44:51.755962Z","iopub.execute_input":"2026-06-10T07:44:51.756235Z","iopub.status.idle":"2026-06-10T07:44:51.909417Z","shell.execute_reply.started":"2026-06-10T07:44:51.756208Z","shell.execute_reply":"2026-06-10T07:44:51.908759Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"84607c9d","cell_type":"code","source":"def reshape_train(df):\n    rows = []\n    for _, r in df.iterrows():\n        for col, val in r.items():\n            if col == 'study_id':\n                continue\n            parts = col.split('_')\n            condition = ' '.join(p.capitalize() for p in parts[:-2])\n            level = parts[-2].capitalize() + '/' + parts[-1].capitalize()\n            rows.append({'study_id': r['study_id'], 'condition': condition,\n                          'level': level, 'severity': val})\n    return pd.DataFrame(rows)\n\nlong_train = reshape_train(train)\nprint(long_train.head())\nprint('Severity counts:\\n', long_train['severity'].value_counts(dropna=False))","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:44:51.910433Z","iopub.execute_input":"2026-06-10T07:44:51.910743Z","iopub.status.idle":"2026-06-10T07:44:52.244021Z","shell.execute_reply.started":"2026-06-10T07:44:51.91071Z","shell.execute_reply":"2026-06-10T07:44:52.243332Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"29d73b00","cell_type":"code","source":"# Join long labels with ROI coordinates and series descriptions\nmerged = long_train.merge(coords, on=['study_id', 'condition', 'level'], how='inner')\nmerged = merged.merge(train_desc, on=['study_id', 'series_id'], how='inner')\n\n# Normalize severity strings\nmerged['severity'] = merged['severity'].map(\n    {'Normal/Mild': 'normal_mild', 'Moderate': 'moderate', 'Severe': 'severe'}\n)\nmerged = merged.dropna(subset=['severity'])\n\n# Build per-image path and a row_id\nmerged['image_path'] = (TRAIN_IMG_DIR + '/' +\n                        merged['study_id'].astype(str) + '/' +\n                        merged['series_id'].astype(str) + '/' +\n                        merged['instance_number'].astype(str) + '.dcm')\nmerged['label'] = merged['severity'].map(LABEL_MAP).astype(int)\n\n# Keep only existing files\nmerged = merged[merged['image_path'].map(os.path.exists)].reset_index(drop=True)\nprint('Final rows:', len(merged))\nprint(merged[['series_description', 'severity']].value_counts())","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:44:52.244954Z","iopub.execute_input":"2026-06-10T07:44:52.245238Z","iopub.status.idle":"2026-06-10T07:46:41.351009Z","shell.execute_reply.started":"2026-06-10T07:44:52.245216Z","shell.execute_reply":"2026-06-10T07:46:41.35027Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"554ca063","cell_type":"markdown","source":"## 3. DICOM loading (with VOI LUT) and ROI cropping","metadata":{}},{"id":"aeab4bfe","cell_type":"code","source":"def load_dicom(path):\n    \"\"\"Load DICOM, apply VOI LUT if present, return uint8 grayscale array.\"\"\"\n    ds = pydicom.dcmread(path)\n    arr = ds.pixel_array.astype(np.float32)\n    try:\n        arr = apply_voi_lut(ds.pixel_array, ds).astype(np.float32)\n    except Exception:\n        pass\n    if getattr(ds, 'PhotometricInterpretation', '') == 'MONOCHROME1':\n        arr = arr.max() - arr\n    arr -= arr.min()\n    if arr.max() > 0:\n        arr = arr / arr.max()\n    return (arr * 255).astype(np.uint8)\n\ndef crop_roi(img, x, y, half=ROI_SIZE // 2):\n    \"\"\"Crop a square ROI around (x, y); pad with zeros if it spills out.\"\"\"\n    h, w = img.shape\n    x, y = int(round(x)), int(round(y))\n    x0, x1 = x - half, x + half\n    y0, y1 = y - half, y + half\n    pad_l = max(0, -x0); pad_r = max(0, x1 - w)\n    pad_t = max(0, -y0); pad_b = max(0, y1 - h)\n    x0c, x1c = max(0, x0), min(w, x1)\n    y0c, y1c = max(0, y0), min(h, y1)\n    crop = img[y0c:y1c, x0c:x1c]\n    if pad_l or pad_r or pad_t or pad_b:\n        crop = np.pad(crop, ((pad_t, pad_b), (pad_l, pad_r)), mode='constant')\n    return crop","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:46:41.351976Z","iopub.execute_input":"2026-06-10T07:46:41.352413Z","iopub.status.idle":"2026-06-10T07:46:41.359672Z","shell.execute_reply.started":"2026-06-10T07:46:41.352389Z","shell.execute_reply":"2026-06-10T07:46:41.358832Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"fe06f955","cell_type":"code","source":"# Visual sanity-check: a few ROI crops per class\nfig, axes = plt.subplots(3, 4, figsize=(12, 9))\nfor r, cls in enumerate(CLASS_NAMES):\n    samples = merged[merged['severity'] == cls].sample(min(4, len(merged)), random_state=SEED)\n    for c, (_, row) in enumerate(samples.iterrows()):\n        img = load_dicom(row['image_path'])\n        roi = crop_roi(img, row['x'], row['y'])\n        axes[r, c].imshow(roi, cmap='gray')\n        axes[r, c].set_title(f\"{cls}\\n{row['condition'][:18]}\", fontsize=8)\n        axes[r, c].axis('off')\nplt.tight_layout(); plt.show()","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:46:41.360607Z","iopub.execute_input":"2026-06-10T07:46:41.360887Z","iopub.status.idle":"2026-06-10T07:46:42.482819Z","shell.execute_reply.started":"2026-06-10T07:46:41.360859Z","shell.execute_reply":"2026-06-10T07:46:42.481957Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"d225c381","cell_type":"markdown","source":"## 4. Dataset & DataLoaders\n\nSingle combined dataset across all series. Train transforms include flips, small rotations, and brightness jitter; validation only resizes.","metadata":{}},{"id":"09c4e390","cell_type":"code","source":"IMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.ToPILImage(),\n    T.Grayscale(num_output_channels=3),\n    T.Resize((IMG_SIZE, IMG_SIZE)),\n    T.RandomHorizontalFlip(p=0.5),\n    T.RandomAffine(degrees=15, translate=(0.05, 0.05), scale=(0.9, 1.1)),\n    T.ColorJitter(brightness=0.15, contrast=0.15),\n    T.ToTensor(),\n    T.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n    T.RandomErasing(p=0.25, scale=(0.02, 0.1)),\n])\n\nval_tf = T.Compose([\n    T.ToPILImage(),\n    T.Grayscale(num_output_channels=3),\n    T.Resize((IMG_SIZE, IMG_SIZE)),\n    T.ToTensor(),\n    T.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n])\n\nclass SpineROIDataset(Dataset):\n    def __init__(self, df, transform):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i):\n        row = self.df.iloc[i]\n        img = load_dicom(row['image_path'])\n        roi = crop_roi(img, row['x'], row['y'])\n        return self.transform(roi), int(row['label'])","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:46:42.483764Z","iopub.execute_input":"2026-06-10T07:46:42.484094Z","iopub.status.idle":"2026-06-10T07:46:42.492169Z","shell.execute_reply.started":"2026-06-10T07:46:42.484072Z","shell.execute_reply":"2026-06-10T07:46:42.491427Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"f3bdd7dc","cell_type":"code","source":"# Stratified split by severity\ntrain_df, val_df = train_test_split(\n    merged, test_size=0.2, random_state=SEED, stratify=merged['label']\n)\nprint('train:', len(train_df), 'val:', len(val_df))\nprint('train severity:\\n', train_df['severity'].value_counts())\nprint('val severity:\\n', val_df['severity'].value_counts())\n\ntrain_ds = SpineROIDataset(train_df, train_tf)\nval_ds   = SpineROIDataset(val_df,   val_tf)\n\n# Use class-weighted loss only (no oversampler) to avoid double-correction\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True,\n                          num_workers=NUM_WORKERS, pin_memory=True, drop_last=True)\nval_loader   = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=NUM_WORKERS, pin_memory=True)\n\n# Class weights for the loss function\nclass_weights = compute_class_weight('balanced',\n                                     classes=np.arange(NUM_CLASSES),\n                                     y=train_df['label'].values)\nclass_weights = torch.tensor(class_weights, dtype=torch.float32, device=device)\nprint('Class weights:', class_weights)","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:46:42.494327Z","iopub.execute_input":"2026-06-10T07:46:42.494616Z","iopub.status.idle":"2026-06-10T07:46:43.066039Z","shell.execute_reply.started":"2026-06-10T07:46:42.494597Z","shell.execute_reply":"2026-06-10T07:46:43.065295Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"fd120bfa","cell_type":"markdown","source":"## 5. Attention modules (PyTorch ports of the reference notebook)","metadata":{}},{"id":"7ded1598","cell_type":"code","source":"class ChannelAttention(nn.Module):\n    def __init__(self, in_ch, reduction=16):\n        super().__init__()\n        hidden = max(in_ch // reduction, 4)\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.max_pool = nn.AdaptiveMaxPool2d(1)\n        self.mlp = nn.Sequential(\n            nn.Conv2d(in_ch, hidden, 1, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(hidden, in_ch, 1, bias=False),\n        )\n\n    def forward(self, x):\n        a = self.mlp(self.avg_pool(x))\n        m = self.mlp(self.max_pool(x))\n        return torch.sigmoid(a + m)\n\nclass SpatialAttention(nn.Module):\n    def __init__(self, kernel_size=7):\n        super().__init__()\n        self.conv = nn.Conv2d(2, 1, kernel_size, padding=kernel_size // 2, bias=False)\n\n    def forward(self, x):\n        avg = torch.mean(x, dim=1, keepdim=True)\n        mx, _ = torch.max(x, dim=1, keepdim=True)\n        return torch.sigmoid(self.conv(torch.cat([avg, mx], dim=1)))\n\nclass CBAM(nn.Module):\n    \"\"\"Convolutional Block Attention Module (Woo et al., 2018).\"\"\"\n    def __init__(self, in_ch, reduction=16):\n        super().__init__()\n        self.ca = ChannelAttention(in_ch, reduction)\n        self.sa = SpatialAttention()\n\n    def forward(self, x):\n        x = x * self.ca(x)\n        x = x * self.sa(x)\n        return x\n\nclass SEBlock(nn.Module):\n    \"\"\"Squeeze-and-Excitation (Hu et al., 2018).\"\"\"\n    def __init__(self, in_ch, reduction=16):\n        super().__init__()\n        hidden = max(in_ch // reduction, 4)\n        self.fc = nn.Sequential(\n            nn.Linear(in_ch, hidden, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Linear(hidden, in_ch, bias=False),\n            nn.Sigmoid(),\n        )\n\n    def forward(self, x):\n        b, c, _, _ = x.shape\n        y = F.adaptive_avg_pool2d(x, 1).view(b, c)\n        y = self.fc(y).view(b, c, 1, 1)\n        return x * y","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:46:43.066988Z","iopub.execute_input":"2026-06-10T07:46:43.067336Z","iopub.status.idle":"2026-06-10T07:46:43.077848Z","shell.execute_reply.started":"2026-06-10T07:46:43.067314Z","shell.execute_reply":"2026-06-10T07:46:43.076817Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"c3bb7b69","cell_type":"markdown","source":"## 6. Models","metadata":{}},{"id":"fb9a0417","cell_type":"code","source":"class EffNetCBAM(nn.Module):\n    \"\"\"EfficientNetV2-S backbone + CBAM after last conv block.\"\"\"\n    FEAT_DIM = 1280\n\n    def __init__(self, num_classes=NUM_CLASSES, pretrained=True):\n        super().__init__()\n        weights = tvm.EfficientNet_V2_S_Weights.IMAGENET1K_V1 if pretrained else None\n        backbone = tvm.efficientnet_v2_s(weights=weights)\n        self.features = backbone.features\n        self.cbam = CBAM(self.FEAT_DIM)\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.head = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(self.FEAT_DIM, 512),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.5),\n            nn.Linear(512, num_classes),\n        )\n\n    def forward_features(self, x):\n        x = self.features(x)\n        x = self.cbam(x)\n        return self.pool(x).flatten(1)\n\n    def forward(self, x):\n        return self.head(self.forward_features(x))\n\n\nclass VGGSE(nn.Module):\n    \"\"\"VGG16 backbone + Squeeze-and-Excitation on the final feature map.\"\"\"\n    FEAT_DIM = 512\n\n    def __init__(self, num_classes=NUM_CLASSES, pretrained=True):\n        super().__init__()\n        weights = tvm.VGG16_Weights.IMAGENET1K_V1 if pretrained else None\n        backbone = tvm.vgg16(weights=weights)\n        self.features = backbone.features\n        self.se = SEBlock(self.FEAT_DIM)\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.head = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(self.FEAT_DIM, 256),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.5),\n            nn.Linear(256, num_classes),\n        )\n\n    def forward_features(self, x):\n        x = self.features(x)\n        x = self.se(x)\n        return self.pool(x).flatten(1)\n\n    def forward(self, x):\n        return self.head(self.forward_features(x))\n\n\neff_model = EffNetCBAM().to(device)\nvgg_model = VGGSE().to(device)\n\nn_eff = sum(p.numel() for p in eff_model.parameters() if p.requires_grad)\nn_vgg = sum(p.numel() for p in vgg_model.parameters() if p.requires_grad)\nprint(f'EffNet+CBAM trainable params: {n_eff:,}')\nprint(f'VGG16+SE    trainable params: {n_vgg:,}')\n\n# Wrap with DataParallel when multiple GPUs are available\nif NUM_GPUS > 1:\n    eff_model = nn.DataParallel(eff_model)\n    vgg_model = nn.DataParallel(vgg_model)\n    print(f'Using DataParallel across {NUM_GPUS} GPUs')\n","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:46:43.078789Z","iopub.execute_input":"2026-06-10T07:46:43.079337Z","iopub.status.idle":"2026-06-10T07:46:48.175142Z","shell.execute_reply.started":"2026-06-10T07:46:43.079296Z","shell.execute_reply":"2026-06-10T07:46:48.17401Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"22aec003","cell_type":"markdown","source":"## 7. Training loop\n\nFull fine-tune with a discriminative learning rate (lower for backbone, higher for the attention + classifier head), class-weighted cross-entropy, cosine LR schedule, and early stopping on best validation accuracy.","metadata":{}},{"id":"8f105c79","cell_type":"code","source":"def _unwrap(model):\n    \"\"\"Return the underlying module, stripping DataParallel if present.\"\"\"\n    return model.module if isinstance(model, nn.DataParallel) else model\n\ndef make_optimizer(model):\n    backbone_params, head_params = [], []\n    for name, p in model.named_parameters():\n        if not p.requires_grad:\n            continue\n        # With DataParallel names are 'module.features.*'; without: 'features.*'\n        if 'features.' in name:\n            backbone_params.append(p)\n        else:\n            head_params.append(p)\n    return torch.optim.AdamW(\n        [{'params': backbone_params, 'lr': LR_BACKBONE},\n         {'params': head_params,     'lr': LR_HEAD}],\n        weight_decay=1e-4,\n    )\n\ndef train_model(model, name, num_epochs=NUM_EPOCHS, patience=6):\n    optimizer = make_optimizer(model)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)\n    criterion = nn.CrossEntropyLoss(weight=class_weights, label_smoothing=LABEL_SMOOTH)\n    scaler = torch.cuda.amp.GradScaler(enabled=USE_AMP and device.type == 'cuda')\n\n    history = {'train_loss': [], 'val_loss': [], 'train_acc': [], 'val_acc': []}\n    best_acc, best_state, stale = 0.0, deepcopy(_unwrap(model).state_dict()), 0\n\n    for epoch in range(num_epochs):\n        # ---- train ----\n        model.train()\n        tl, tc, tn = 0.0, 0, 0\n        for x, y in tqdm(train_loader, desc=f'{name} ep {epoch+1}/{num_epochs}', leave=False):\n            x, y = x.to(device), y.to(device)\n            optimizer.zero_grad()\n            with torch.cuda.amp.autocast(enabled=USE_AMP and device.type == 'cuda'):\n                out = model(x)\n                loss = criterion(out, y)\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            tl += loss.item() * x.size(0)\n            tc += (out.argmax(1) == y).sum().item()\n            tn += x.size(0)\n        scheduler.step()\n        tr_loss, tr_acc = tl / tn, tc / tn\n\n        # ---- validate ----\n        model.eval()\n        vl, vc, vn = 0.0, 0, 0\n        with torch.no_grad():\n            for x, y in val_loader:\n                x, y = x.to(device), y.to(device)\n                out = model(x)\n                loss = criterion(out, y)\n                vl += loss.item() * x.size(0)\n                vc += (out.argmax(1) == y).sum().item()\n                vn += x.size(0)\n        va_loss, va_acc = vl / vn, vc / vn\n\n        history['train_loss'].append(tr_loss); history['val_loss'].append(va_loss)\n        history['train_acc'].append(tr_acc);   history['val_acc'].append(va_acc)\n        print(f'[{name}] ep {epoch+1:02d} | train loss {tr_loss:.4f} acc {tr_acc:.4f} | '\n              f'val loss {va_loss:.4f} acc {va_acc:.4f}')\n\n        if va_acc > best_acc:\n            best_acc, best_state, stale = va_acc, deepcopy(_unwrap(model).state_dict()), 0\n            torch.save(best_state, f'best_{name}.pth')\n        else:\n            stale += 1\n            if stale >= patience:\n                print(f'Early stop at epoch {epoch+1}')\n                break\n\n    _unwrap(model).load_state_dict(best_state)\n    return history, best_acc\n","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:46:48.176182Z","iopub.execute_input":"2026-06-10T07:46:48.176473Z","iopub.status.idle":"2026-06-10T07:46:48.188418Z","shell.execute_reply.started":"2026-06-10T07:46:48.176451Z","shell.execute_reply":"2026-06-10T07:46:48.187781Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"ed8cb029","cell_type":"code","source":"print('=== Training EfficientNetV2-S + CBAM ===')\neff_history, eff_best = train_model(eff_model, 'effnet_cbam')\nprint(f'Best val acc: {eff_best:.4f}')","metadata":{"execution":{"iopub.status.busy":"2026-06-10T07:46:48.189187Z","iopub.execute_input":"2026-06-10T07:46:48.189525Z","execution_failed":"2026-06-10T08:14:44.702Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"058d9c15","cell_type":"code","source":"print('=== Training VGG16 + SE ===')\nvgg_history, vgg_best = train_model(vgg_model, 'vgg_se')\nprint(f'Best val acc: {vgg_best:.4f}')","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.703Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"d0a88c15","cell_type":"markdown","source":"## 8. Training curves","metadata":{}},{"id":"7767756c","cell_type":"code","source":"def plot_history(history, title):\n    fig, ax = plt.subplots(1, 2, figsize=(12, 4))\n    ax[0].plot(history['train_acc'], color='blue',    label='Train')\n    ax[0].plot(history['val_acc'],   color='magenta', label='Validation')\n    ax[0].set_title(f'{title} — Accuracy'); ax[0].set_xlabel('Epoch'); ax[0].set_ylabel('Accuracy'); ax[0].legend()\n    ax[1].plot(history['train_loss'], color='blue',    label='Train')\n    ax[1].plot(history['val_loss'],   color='magenta', label='Validation')\n    ax[1].set_title(f'{title} — Loss'); ax[1].set_xlabel('Epoch'); ax[1].set_ylabel('Loss'); ax[1].legend()\n    plt.tight_layout(); plt.show()\n\nplot_history(eff_history, 'EfficientNetV2-S + CBAM')\nplot_history(vgg_history, 'VGG16 + SE')","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.703Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"f5038665","cell_type":"markdown","source":"## 9. Per-model evaluation","metadata":{}},{"id":"aba31f65","cell_type":"code","source":"@torch.no_grad()\ndef predict(model, loader, tta=USE_TTA):\n    \"\"\"Forward pass with optional horizontal-flip TTA averaging.\"\"\"\n    model.eval()\n    all_probs, all_y = [], []\n    for x, y in loader:\n        x = x.to(device)\n        probs = F.softmax(model(x), dim=1)\n        if tta:\n            probs_flip = F.softmax(model(torch.flip(x, dims=[3])), dim=1)\n            probs = (probs + probs_flip) / 2\n        all_probs.append(probs.cpu().numpy()); all_y.append(y.numpy())\n    return np.concatenate(all_probs), np.concatenate(all_y)\n\ndef plot_confusion(y_true, y_pred, title):\n    cm = confusion_matrix(y_true, y_pred)\n    cm_df = pd.DataFrame(cm, index=CLASS_NAMES, columns=CLASS_NAMES)\n    plt.figure(figsize=(6, 5))\n    sns.heatmap(cm_df, cmap='Blues', annot=True, fmt='d', linewidths=0.5, linecolor='black')\n    plt.title(f'{title} — Confusion Matrix'); plt.ylabel('True'); plt.xlabel('Predicted')\n    plt.tight_layout(); plt.show()\n\ndef plot_roc_pr(y_true, probs, title):\n    fig, ax = plt.subplots(1, 2, figsize=(13, 5))\n    for i, cls in enumerate(CLASS_NAMES):\n        fpr, tpr, _ = roc_curve(y_true == i, probs[:, i])\n        ax[0].plot(fpr, tpr, lw=2, label=f'{cls} (AUC={auc(fpr, tpr):.3f})')\n        pr, rc, _ = precision_recall_curve(y_true == i, probs[:, i])\n        ax[1].plot(rc, pr, lw=2, label=f'{cls} (AP={auc(rc, pr):.3f})')\n    ax[0].plot([0, 1], [0, 1], 'k--', lw=1)\n    ax[0].set_xlabel('FPR'); ax[0].set_ylabel('TPR'); ax[0].set_title(f'{title} — ROC'); ax[0].legend(loc='lower right')\n    ax[1].set_xlabel('Recall'); ax[1].set_ylabel('Precision'); ax[1].set_title(f'{title} — PR'); ax[1].legend(loc='lower left')\n    plt.tight_layout(); plt.show()\n\ndef evaluate(model, name):\n    probs, y_true = predict(model, val_loader)\n    y_pred = probs.argmax(1)\n    acc = (y_pred == y_true).mean()\n    print(f'\\n=== {name} | val accuracy: {acc:.4f} ===')\n    print(classification_report(y_true, y_pred, target_names=CLASS_NAMES, digits=3, zero_division=0))\n    plot_confusion(y_true, y_pred, name)\n    plot_roc_pr(y_true, probs, name)\n    return probs, y_true, acc","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.703Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"83371008","cell_type":"code","source":"eff_probs, y_val, eff_acc = evaluate(eff_model, 'EfficientNetV2-S + CBAM')","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.703Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"3f5be3ef","cell_type":"code","source":"vgg_probs, _, vgg_acc = evaluate(vgg_model, 'VGG16 + SE')","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.704Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"7cc415b7","cell_type":"markdown","source":"## 10. Weighted-average ensemble\n\nWeights chosen proportionally to each model's validation accuracy.","metadata":{}},{"id":"bc291f86","cell_type":"code","source":"w = np.array([eff_acc, vgg_acc], dtype=np.float32)\nw = w / w.sum()\nprint(f'Ensemble weights — EffNet+CBAM: {w[0]:.3f}, VGG16+SE: {w[1]:.3f}')\nweighted_probs = w[0] * eff_probs + w[1] * vgg_probs\nweighted_pred  = weighted_probs.argmax(1)\nweighted_acc   = (weighted_pred == y_val).mean()\nprint(f'\\n=== Weighted Ensemble | val accuracy: {weighted_acc:.4f} ===')\nprint(classification_report(y_val, weighted_pred, target_names=CLASS_NAMES, digits=3, zero_division=0))\nplot_confusion(y_val, weighted_pred, 'Weighted Ensemble')\nplot_roc_pr(y_val, weighted_probs, 'Weighted Ensemble')","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.704Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"e5734964","cell_type":"markdown","source":"## 11. Unified feature-fusion ensemble\n\nFreezes both backbones, concatenates their attention-refined feature vectors, applies a learned attention gate (matching the IQ-OTH/NCCD reference), then trains a small classification head.","metadata":{}},{"id":"aeda740f","cell_type":"code","source":"class UnifiedEnsemble(nn.Module):\n    def __init__(self, eff, vgg, num_classes=NUM_CLASSES):\n        super().__init__()\n        # Unwrap DataParallel so forward_features is directly accessible\n        eff = eff.module if isinstance(eff, nn.DataParallel) else eff\n        vgg = vgg.module if isinstance(vgg, nn.DataParallel) else vgg\n        # Freeze base models — we only train the fusion head\n        for p in eff.parameters(): p.requires_grad = False\n        for p in vgg.parameters(): p.requires_grad = False\n        self.eff = eff\n        self.vgg = vgg\n        combined = EffNetCBAM.FEAT_DIM + VGGSE.FEAT_DIM  # 1280 + 512\n        self.attention = nn.Sequential(\n            nn.Linear(combined, combined),\n            nn.Sigmoid(),\n        )\n        self.classifier = nn.Sequential(\n            nn.Linear(combined, 256),\n            nn.BatchNorm1d(256),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.5),\n            nn.Linear(256, num_classes),\n        )\n\n    def forward(self, x):\n        with torch.no_grad():\n            fe = self.eff.forward_features(x)\n            fv = self.vgg.forward_features(x)\n        c = torch.cat([fe, fv], dim=1)\n        c = c * self.attention(c)\n        return self.classifier(c)\n\nunified = UnifiedEnsemble(eff_model, vgg_model).to(device)\n\nif NUM_GPUS > 1:\n    unified = nn.DataParallel(unified)\n\nn_uni = sum(p.numel() for p in unified.parameters() if p.requires_grad)\nprint(f'Unified ensemble trainable params: {n_uni:,}')\n","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.704Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"d02319bf","cell_type":"code","source":"def train_unified(model, num_epochs=12):\n    optimizer = torch.optim.AdamW(\n        [p for p in model.parameters() if p.requires_grad], lr=1e-3, weight_decay=1e-4\n    )\n    criterion = nn.CrossEntropyLoss(weight=class_weights, label_smoothing=LABEL_SMOOTH)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)\n    history = {'train_loss': [], 'val_loss': [], 'train_acc': [], 'val_acc': []}\n    best_acc, best_state = 0.0, deepcopy(_unwrap(model).state_dict())\n    for epoch in range(num_epochs):\n        model.train()\n        # keep frozen backbones in eval mode (BatchNorm running stats)\n        _base = _unwrap(model)\n        _base.eff.eval(); _base.vgg.eval()\n        tl, tc, tn = 0.0, 0, 0\n        for x, y in tqdm(train_loader, desc=f'unified ep {epoch+1}/{num_epochs}', leave=False):\n            x, y = x.to(device), y.to(device)\n            optimizer.zero_grad()\n            out = model(x)\n            loss = criterion(out, y)\n            loss.backward()\n            optimizer.step()\n            tl += loss.item() * x.size(0)\n            tc += (out.argmax(1) == y).sum().item()\n            tn += x.size(0)\n        scheduler.step()\n        model.eval()\n        vl, vc, vn = 0.0, 0, 0\n        with torch.no_grad():\n            for x, y in val_loader:\n                x, y = x.to(device), y.to(device)\n                out = model(x)\n                loss = criterion(out, y)\n                vl += loss.item() * x.size(0)\n                vc += (out.argmax(1) == y).sum().item()\n                vn += x.size(0)\n        tr_loss, tr_acc = tl / tn, tc / tn\n        va_loss, va_acc = vl / vn, vc / vn\n        history['train_loss'].append(tr_loss); history['val_loss'].append(va_loss)\n        history['train_acc'].append(tr_acc);   history['val_acc'].append(va_acc)\n        print(f'[unified] ep {epoch+1:02d} | train loss {tr_loss:.4f} acc {tr_acc:.4f} | '\n              f'val loss {va_loss:.4f} acc {va_acc:.4f}')\n        if va_acc > best_acc:\n            best_acc, best_state = va_acc, deepcopy(_unwrap(model).state_dict())\n            torch.save(best_state, 'best_unified.pth')\n    _unwrap(model).load_state_dict(best_state)\n    return history, best_acc\n\n\nuni_history, uni_best = train_unified(unified, num_epochs=12)\nprint(f'Unified ensemble best val acc: {uni_best:.4f}')\n","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.704Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"b77477f3","cell_type":"code","source":"plot_history(uni_history, 'Unified Ensemble')\nuni_probs, _, uni_acc = evaluate(unified, 'Unified Ensemble')","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.704Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"09acea87","cell_type":"markdown","source":"## 12. Final comparison","metadata":{}},{"id":"5c9ac8bb","cell_type":"code","source":"results = {\n    'EffNet+CBAM':       eff_acc,\n    'VGG16+SE':          vgg_acc,\n    'Weighted Ensemble': weighted_acc,\n    'Unified Ensemble':  uni_acc,\n}\n\nplt.figure(figsize=(10, 5))\nbars = plt.bar(results.keys(), results.values(),\n               color=['#2b7a78', '#3aafa9', '#feffff', '#17252a'],\n               edgecolor='black')\nplt.ylabel('Validation Accuracy'); plt.title('Model Comparison')\nplt.ylim([0, max(results.values()) * 1.15])\nfor b, v in zip(bars, results.values()):\n    plt.text(b.get_x() + b.get_width() / 2, v + 0.005, f'{v:.4f}', ha='center', fontsize=11)\nplt.tight_layout(); plt.show()\n\nsummary = pd.DataFrame({'Model': list(results.keys()), 'Val Accuracy': list(results.values())})\nsummary.sort_values('Val Accuracy', ascending=False, inplace=True)\nsummary.reset_index(drop=True, inplace=True)\nprint(summary.to_string(index=False))","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.704Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"8bb91bfd-fdc5-4f05-8845-e5c0d6888de3","cell_type":"markdown","source":"## 13. Exploratory data visualization\n\nDistributions of the merged training set across severity, condition, vertebral level, and MRI series type.","metadata":{}},{"id":"eaf24014-4d66-4f1f-9e71-4079d7523368","cell_type":"code","source":"fig, axes = plt.subplots(2, 2, figsize=(15, 11))\n\n# 1. Severity distribution (overall)\nsev_counts = merged['severity'].value_counts().reindex(CLASS_NAMES)\naxes[0, 0].bar(sev_counts.index, sev_counts.values,\n               color=['#2b7a78', '#3aafa9', '#17252a'], edgecolor='black')\naxes[0, 0].set_title('Severity distribution (all ROIs)')\naxes[0, 0].set_ylabel('Count')\nfor i, v in enumerate(sev_counts.values):\n    axes[0, 0].text(i, v + max(sev_counts.values) * 0.01, f'{v:,}',\n                    ha='center', fontsize=10)\n\n# 2. Severity by MRI series type\nct = pd.crosstab(merged['series_description'], merged['severity']).reindex(columns=CLASS_NAMES)\nct.plot(kind='bar', stacked=True, ax=axes[0, 1],\n        color=['#2b7a78', '#3aafa9', '#17252a'], edgecolor='black')\naxes[0, 1].set_title('Severity by MRI series')\naxes[0, 1].set_ylabel('Count'); axes[0, 1].set_xlabel('')\naxes[0, 1].tick_params(axis='x', rotation=20)\naxes[0, 1].legend(title='Severity', fontsize=8)\n\n# 3. Severity by condition\nct2 = pd.crosstab(merged['condition'], merged['severity']).reindex(columns=CLASS_NAMES)\nct2.plot(kind='barh', stacked=True, ax=axes[1, 0],\n         color=['#2b7a78', '#3aafa9', '#17252a'], edgecolor='black')\naxes[1, 0].set_title('Severity by condition')\naxes[1, 0].set_xlabel('Count'); axes[1, 0].set_ylabel('')\naxes[1, 0].legend(title='Severity', fontsize=8)\n\n# 4. Severity by vertebral level\nct3 = pd.crosstab(merged['level'], merged['severity']).reindex(columns=CLASS_NAMES)\nct3.plot(kind='bar', stacked=True, ax=axes[1, 1],\n         color=['#2b7a78', '#3aafa9', '#17252a'], edgecolor='black')\naxes[1, 1].set_title('Severity by vertebral level')\naxes[1, 1].set_ylabel('Count'); axes[1, 1].set_xlabel('')\naxes[1, 1].tick_params(axis='x', rotation=0)\naxes[1, 1].legend(title='Severity', fontsize=8)\n\nplt.tight_layout(); plt.show()\n\n# Pie chart of class proportions (train vs val)\nfig, axes = plt.subplots(1, 2, figsize=(11, 5))\ncolors = ['#2b7a78', '#3aafa9', '#17252a']\nfor ax, df_, title in [(axes[0], train_df, 'Train'), (axes[1], val_df, 'Validation')]:\n    counts = df_['severity'].value_counts().reindex(CLASS_NAMES)\n    ax.pie(counts.values, labels=counts.index, colors=colors,\n           autopct='%1.1f%%', startangle=90,\n           wedgeprops={'edgecolor': 'white', 'linewidth': 2})\n    ax.set_title(f'{title} severity proportions')\nplt.tight_layout(); plt.show()","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.705Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"6730ac4d-0be8-4fe8-95ef-b18134f81c37","cell_type":"markdown","source":"## 14. Feature-space visualization (t-SNE & PCA)\n\nExtract the 1280-d EffNet+CBAM features and the 1792-d unified-ensemble features for a stratified subset of the validation set, then project them to 2-D with both PCA and t-SNE to see how well the attention models separate the three severity classes.","metadata":{}},{"id":"e6cbf527-041e-4203-a386-db55494bb936","cell_type":"code","source":"from sklearn.manifold import TSNE\nfrom sklearn.decomposition import PCA\n\n# Stratified subset of validation rows (cap for t-SNE speed)\nPER_CLASS = 300\nsubset_parts = []\nfor lbl in range(NUM_CLASSES):\n    pool = val_df[val_df['label'] == lbl]\n    subset_parts.append(pool.sample(min(PER_CLASS, len(pool)), random_state=SEED))\nsubset_df = pd.concat(subset_parts).reset_index(drop=True)\nprint(f'Subset size: {len(subset_df)}  ({subset_df[\"severity\"].value_counts().to_dict()})')\n\nsubset_ds = SpineROIDataset(subset_df, val_tf)\nsubset_loader = DataLoader(subset_ds, batch_size=BATCH_SIZE, shuffle=False,\n                           num_workers=NUM_WORKERS, pin_memory=True)\n\n@torch.no_grad()\ndef extract_features(model, loader, kind='effnet'):\n    \"\"\"kind in {'effnet', 'vgg', 'unified'} — returns (N, D) numpy.\"\"\"\n    model.eval()\n    _base = _unwrap(model)\n    feats, labels = [], []\n    for x, y in tqdm(loader, desc=f'feats:{kind}', leave=False):\n        x = x.to(device)\n        if kind in ('effnet', 'vgg'):\n            f = _base.forward_features(x)\n        else:  # unified — concatenated 1792-d before classifier\n            fe = _base.eff.forward_features(x)\n            fv = _base.vgg.forward_features(x)\n            c = torch.cat([fe, fv], dim=1)\n            f = c * _base.attention(c)\n        feats.append(f.cpu().numpy()); labels.append(y.numpy())\n    return np.concatenate(feats), np.concatenate(labels)\n\neff_feats,  y_sub = extract_features(eff_model, subset_loader, 'effnet')\nuni_feats,  _      = extract_features(unified,  subset_loader, 'unified')\nprint('EffNet feats:', eff_feats.shape, ' Unified feats:', uni_feats.shape)\n","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.705Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"5c6c7e74-77c5-4127-9972-ae22bc444423","cell_type":"code","source":"def plot_embedding(emb, labels, title, ax):\n    colors = ['#2ca02c', '#d62728', '#1f77b4']  # green, red, blue]\n    for i, cls in enumerate(CLASS_NAMES):\n        m = labels == i\n        ax.scatter(emb[m, 0], emb[m, 1], s=14, alpha=0.7,\n                   color=colors[i], edgecolor='white', linewidth=0.3, label=cls)\n    ax.set_title(title); ax.set_xticks([]); ax.set_yticks([])\n    ax.legend(loc='best', fontsize=9)\n\n# PCA first (deterministic, fast) — useful sanity-check baseline\npca_eff = PCA(n_components=2, random_state=SEED).fit_transform(eff_feats)\npca_uni = PCA(n_components=2, random_state=SEED).fit_transform(uni_feats)\n\n# t-SNE — non-linear, slower\ntsne_kwargs = dict(n_components=2, perplexity=30, init='pca',\n                   learning_rate='auto', random_state=SEED)\ntsne_eff = TSNE(**tsne_kwargs).fit_transform(eff_feats)\ntsne_uni = TSNE(**tsne_kwargs).fit_transform(uni_feats)\n\nfig, axes = plt.subplots(2, 2, figsize=(13, 11))\nplot_embedding(pca_eff,  y_sub, 'PCA — EfficientNet+CBAM features (1280-d)',  axes[0, 0])\nplot_embedding(pca_uni,  y_sub, 'PCA — Unified ensemble features (1792-d)',   axes[0, 1])\nplot_embedding(tsne_eff, y_sub, 't-SNE — EfficientNet+CBAM features',         axes[1, 0])\nplot_embedding(tsne_uni, y_sub, 't-SNE — Unified ensemble features',          axes[1, 1])\nplt.tight_layout(); plt.show()","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.705Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"7cd1e0c0-3794-45e4-a921-48580cba3d1e","cell_type":"markdown","source":"## 15. Explainable AI — Grad-CAM\n\nGrad-CAM (Selvaraju et al., 2017) highlights which spatial regions of the ROI the model used to make its prediction. We hook the **last conv block of EfficientNetV2-S, *after* the CBAM spatial-attention module**, so the heat-map reflects the attention-refined features the classifier actually consumes.","metadata":{}},{"id":"3fd90ce5-8a0f-426e-8151-6270b2f27db4","cell_type":"code","source":"class GradCAM:\n    \"\"\"Generic Grad-CAM: hooks a target conv module and returns a heat-map.\"\"\"\n    def __init__(self, model, target_module):\n        self.model = model\n        self.target = target_module\n        self.activations = None\n        self.gradients = None\n        self.h1 = target_module.register_forward_hook(self._save_act)\n        self.h2 = target_module.register_full_backward_hook(self._save_grad)\n\n    def _save_act(self, _m, _i, out):     self.activations = out.detach()\n    def _save_grad(self, _m, _gi, go):    self.gradients   = go[0].detach()\n\n    def __call__(self, x, class_idx=None):\n        self.model.eval()\n        x = x.clone().detach().requires_grad_(True)\n        logits = self.model(x)\n        if class_idx is None:\n            class_idx = logits.argmax(dim=1)\n        score = logits.gather(1, class_idx.view(-1, 1)).sum()\n        self.model.zero_grad(); score.backward(retain_graph=False)\n\n        # weights = global-avg-pool of gradients\n        w = self.gradients.mean(dim=(2, 3), keepdim=True)          # (B, C, 1, 1)\n        cam = F.relu((w * self.activations).sum(dim=1, keepdim=True))  # (B, 1, h, w)\n        cam = F.interpolate(cam, size=x.shape[-2:], mode='bilinear', align_corners=False)\n        # normalise per-image to [0, 1]\n        cam = cam.squeeze(1)\n        cam -= cam.amin(dim=(1, 2), keepdim=True)\n        cam /= (cam.amax(dim=(1, 2), keepdim=True) + 1e-8)\n        return cam.cpu().numpy(), F.softmax(logits, dim=1).detach().cpu().numpy()\n\n    def close(self):\n        self.h1.remove(); self.h2.remove()\n\n\ndef denormalize(t):\n    \"\"\"Undo ImageNet normalization for display. t: (C, H, W) tensor.\"\"\"\n    mean = torch.tensor(IMAGENET_MEAN).view(3, 1, 1)\n    std  = torch.tensor(IMAGENET_STD).view(3, 1, 1)\n    return (t.cpu() * std + mean).clamp(0, 1).permute(1, 2, 0).numpy()\n\n\ndef show_gradcam(model, target_module, df, n_per_class=2, title='Grad-CAM'):\n    cam_tool = GradCAM(model, target_module)\n    rows = []\n    for lbl in range(NUM_CLASSES):\n        pool = df[df['label'] == lbl]\n        rows.append(pool.sample(min(n_per_class, len(pool)), random_state=SEED))\n    samples = pd.concat(rows).reset_index(drop=True)\n\n    ds = SpineROIDataset(samples, val_tf)\n    imgs  = torch.stack([ds[i][0] for i in range(len(ds))]).to(device)\n    ys    = torch.tensor([ds[i][1] for i in range(len(ds))]).to(device)\n    cams, probs = cam_tool(imgs)\n    cam_tool.close()\n\n    n = len(samples)\n    fig, axes = plt.subplots(2, n, figsize=(2.4 * n, 5.2))\n    for j in range(n):\n        img = denormalize(imgs[j])\n        pred = int(probs[j].argmax())\n        conf = probs[j][pred]\n        axes[0, j].imshow(img); axes[0, j].axis('off')\n        axes[0, j].set_title(f'T:{CLASS_NAMES[ys[j].item()]}\\nP:{CLASS_NAMES[pred]} ({conf:.2f})',\n                             fontsize=8,\n                             color='green' if pred == ys[j].item() else 'red')\n        axes[1, j].imshow(img)\n        axes[1, j].imshow(cams[j], cmap='jet', alpha=0.45)\n        axes[1, j].axis('off')\n    fig.suptitle(title, fontsize=13)\n    plt.tight_layout(); plt.show()\n\n\n# Access underlying module for hooks (DataParallel wraps the module)\neff_base = _unwrap(eff_model)\nvgg_base = _unwrap(vgg_model)\n\n# EffNet+CBAM Grad-CAM: hook the CBAM output (post attention, pre pooling)\nshow_gradcam(eff_model, eff_base.cbam, val_df, n_per_class=3,\n             title='Grad-CAM — EfficientNetV2-S + CBAM')\n\n# VGG16+SE Grad-CAM: hook the SE output\nshow_gradcam(vgg_model, vgg_base.se, val_df, n_per_class=3,\n             title='Grad-CAM — VGG16 + SE')\n","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.705Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"ddd5e5ec","cell_type":"markdown","source":"## 16. Save unified model weights and metadata","metadata":{}},{"id":"4caf24d4","cell_type":"code","source":"import json\n\nos.makedirs('unified_model', exist_ok=True)\n\n# Save the module's state dict (without DataParallel 'module.' prefix) for portability\ntorch.save(_unwrap(unified).state_dict(), 'unified_model/unified_ensemble.pth')\n\nmetadata = {\n    'model_type': 'UnifiedEnsemble',\n    'best_val_accuracy': float(uni_acc),\n    'class_names': CLASS_NAMES,\n    'num_classes': NUM_CLASSES,\n    'img_size': IMG_SIZE,\n    'roi_size': ROI_SIZE,\n    'imagenet_mean': IMAGENET_MEAN,\n    'imagenet_std': IMAGENET_STD,\n    'eff_feat_dim': 1280,\n    'vgg_feat_dim': 512,\n    'description': 'Unified attention ensemble: EfficientNetV2-S + VGG16 with feature fusion'\n}\n\nwith open('unified_model/metadata.json', 'w') as f:\n    json.dump(metadata, f, indent=2)\n\nprint('✓ Model saved: unified_model/unified_ensemble.pth')\nprint('✓ Metadata saved: unified_model/metadata.json')\nprint(f'\\nModel accuracy: {uni_acc:.4f}')\nprint(f'Model size: {os.path.getsize(\"unified_model/unified_ensemble.pth\") / 1e6:.1f} MB')\n","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.705Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"937788f7","cell_type":"markdown","source":"## 17. Load saved model and predict on new images","metadata":{}},{"id":"7a89d19e","cell_type":"code","source":"def load_model(checkpoint_path='unified_model/unified_ensemble.pth', meta_path='unified_model/metadata.json'):\n    \"\"\"Load the unified ensemble model with metadata.\"\"\"\n    with open(meta_path, 'r') as f:\n        meta = json.load(f)\n    \n    model = UnifiedEnsemble(eff_model, vgg_model).to(device)\n    model.load_state_dict(torch.load(checkpoint_path, map_location=device))\n    model.eval()\n    return model, meta\n\n@torch.no_grad()\ndef predict_single(model, dicom_path, roi_x, roi_y, use_tta=True):\n    \"\"\"Predict on a single DICOM ROI. Returns class label, probabilities, confidence.\"\"\"\n    img = load_dicom(dicom_path)\n    roi = crop_roi(img, roi_x, roi_y)\n    img_tensor = val_tf(roi).unsqueeze(0).to(device)\n    \n    logits = model(img_tensor)\n    probs = F.softmax(logits, dim=1)\n    \n    if use_tta:\n        logits_flip = model(torch.flip(img_tensor, dims=[3]))\n        probs_flip = F.softmax(logits_flip, dim=1)\n        probs = (probs + probs_flip) / 2\n    \n    pred_idx = probs.argmax(1).item()\n    confidence = probs[0, pred_idx].item()\n    probs_dict = {CLASS_NAMES[i]: probs[0, i].item() for i in range(len(CLASS_NAMES))}\n    \n    return CLASS_NAMES[pred_idx], probs_dict, confidence\n\n# Example usage: predict on a validation sample\nprint('Loading model...')\nmodel_loaded, meta = load_model()\nprint(f'✓ Loaded: {meta[\"description\"]}')\nprint(f'  Classes: {meta[\"class_names\"]}')\n\n# Test on 3 random validation samples\nprint('\\n' + '='*60)\nprint('PREDICTIONS ON VALIDATION SET SAMPLES')\nprint('='*60)\nfor i in range(3):\n    sample = val_df.sample(1, random_state=SEED+i).iloc[0]\n    pred_label, probs, conf = predict_single(model_loaded, sample['image_path'], sample['x'], sample['y'])\n    true_label = sample['severity']\n    \n    print(f'\\nSample {i+1}:')\n    print(f'  True label:  {true_label}')\n    print(f'  Predicted:   {pred_label} (confidence: {conf:.4f})')\n    print(f'  Probabilities:')\n    for cls, prob in probs.items():\n        print(f'    {cls:15s}: {prob:.4f}')\n    print(f'  ✓ CORRECT' if pred_label == true_label else f'  ✗ INCORRECT')","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.705Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"5c5bf3b0","cell_type":"markdown","source":"## 18. Package and download model files","metadata":{}},{"id":"d68258a1","cell_type":"code","source":"import zipfile\nimport shutil\n\ndef create_download_package(output_name='rsna_spine_ensemble_model.zip'):\n    \"\"\"Create a zip file with model weights, metadata, and inference code.\"\"\"\n    \n    # Create inference helper script\n    inference_code = '''#!/usr/bin/env python3\n\"\"\"\nRSNA Lumbar Spine Unified Ensemble Model — Inference Script\nUnified attention ensemble: EfficientNetV2-S + VGG16 with feature fusion\n\"\"\"\n\nimport os\nimport json\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as tvm\nimport torchvision.transforms as T\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nimport pydicom\n\ndef load_dicom(path):\n    \"\"\"Load DICOM, apply VOI LUT if present, return uint8 grayscale.\"\"\"\n    ds = pydicom.dcmread(path)\n    arr = ds.pixel_array.astype(np.float32)\n    try:\n        arr = apply_voi_lut(ds.pixel_array, ds).astype(np.float32)\n    except:\n        pass\n    if getattr(ds, 'PhotometricInterpretation', '') == 'MONOCHROME1':\n        arr = arr.max() - arr\n    arr -= arr.min()\n    if arr.max() > 0:\n        arr = arr / arr.max()\n    return (arr * 255).astype(np.uint8)\n\ndef crop_roi(img, x, y, roi_size=128):\n    \"\"\"Crop square ROI around (x, y); pad with zeros if spills out.\"\"\"\n    h, w = img.shape\n    half = roi_size // 2\n    x, y = int(round(x)), int(round(y))\n    x0, x1 = x - half, x + half\n    y0, y1 = y - half, y + half\n    pad_l, pad_r = max(0, -x0), max(0, x1 - w)\n    pad_t, pad_b = max(0, -y0), max(0, y1 - h)\n    x0c, x1c = max(0, x0), min(w, x1)\n    y0c, y1c = max(0, y0), min(h, y1)\n    crop = img[y0c:y1c, x0c:x1c]\n    if pad_l or pad_r or pad_t or pad_b:\n        crop = np.pad(crop, ((pad_t, pad_b), (pad_l, pad_r)), mode='constant')\n    return crop\n\n# Load model, classes, transforms from package\nwith open('metadata.json') as f:\n    meta = json.load(f)\n\nCLASS_NAMES = meta['class_names']\nIMG_SIZE = meta['img_size']\nMEAN = meta['imagenet_mean']\nSTD = meta['imagenet_std']\n\nval_tf = T.Compose([\n    T.ToPILImage(),\n    T.Grayscale(num_output_channels=3),\n    T.Resize((IMG_SIZE, IMG_SIZE)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\n# If you have the full model definition, load it:\n# from model_def import EffNetCBAM, VGGSE, UnifiedEnsemble\n# device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n# eff_model = EffNetCBAM().to(device)\n# vgg_model = VGGSE().to(device)\n# unified = UnifiedEnsemble(eff_model, vgg_model).to(device)\n# unified.load_state_dict(torch.load('unified_ensemble.pth', map_location=device))\n\nprint(\"✓ Model package loaded. Define model architecture and load weights to make predictions.\")\nprint(f\"  Classes: {CLASS_NAMES}\")\n'''\n    \n    with open('unified_model/inference.py', 'w') as f:\n        f.write(inference_code)\n    \n    # Create README\n    readme = f'''# RSNA Lumbar Spine Unified Ensemble Model\n\n## Model Details\n- **Type**: Unified Attention Ensemble\n- **Architecture**: EfficientNetV2-S + VGG16 with feature fusion\n- **Best Validation Accuracy**: {uni_acc:.4f}\n- **Classes**: {\", \".join(CLASS_NAMES)}\n- **Input Size**: 224×224 (grayscale MRI ROI)\n- **Model Size**: {os.path.getsize(\"unified_model/unified_ensemble.pth\") / 1e6:.1f} MB\n\n## Files\n- `unified_ensemble.pth`: Model weights\n- `metadata.json`: Configuration and hyperparameters\n- `inference.py`: Inference helper script template\n- `README.md`: This file\n\n## Usage\n\n### Python (with training notebook environment)\n```python\nimport torch\nimport json\n\n# Load metadata\nwith open('metadata.json') as f:\n    meta = json.load(f)\n\n# Load model (requires model architecture definition)\nmodel = UnifiedEnsemble(eff_model, vgg_model)\nmodel.load_state_dict(torch.load('unified_ensemble.pth'))\nmodel.eval()\n\n# Predict on new image\nfrom predict import predict_single\npred_label, probs, confidence = predict_single(\n    model, \n    'path/to/dicom.dcm', \n    roi_x=512, \n    roi_y=384\n)\nprint(f\"Prediction: {{pred_label}} ({{confidence:.4f}})\")\n```\n\n## Training Details\n- **Dataset**: RSNA 2024 Lumbar Spine (48,657 ROI crops)\n- **Train/Val Split**: 80/20 stratified\n- **Optimization**: AdamW with discriminative learning rates\n- **Loss**: Class-weighted cross-entropy with label smoothing\n- **Regularization**: Mixed precision, dropout, random erasing\n- **Ensemble Strategy**: Frozen backbone features + learned attention gate\n\n## Model Architecture\n\n### Individual Models\n1. **EfficientNetV2-S + CBAM**: 21.0M params\n   - Channel + spatial attention after last conv block\n   \n2. **VGG16 + SE**: 14.9M params\n   - Squeeze-and-Excitation attention on final features\n\n### Fusion Head\n- Concatenates frozen features: 1280 + 512 = 1792 dims\n- Learned attention gate: 1792 → 1792\n- Classifier: 1792 → 256 (BN + ReLU) → 3 classes\n- Trainable params: 3.7M\n\n## Performance\n\n| Model | Validation Accuracy |\n|-------|-------------------|\n| EffNet+CBAM | 81.87% |\n| VGG16+SE | 80.78% |\n| Weighted Ensemble | 82.54% |\n| **Unified Ensemble** | **83.13%** |\n\n## Citation\nIf you use this model, please cite:\n- Woo et al., 2018. CBAM: Convolutional Block Attention Module\n- Hu et al., 2018. Squeeze-and-Excitation Networks\n- RSNA 2024 Lumbar Spine Degenerative Classification Challenge\n'''\n    \n    with open('unified_model/README.md', 'w') as f:\n        f.write(readme)\n    \n    # Create zip\n    if os.path.exists(output_name):\n        os.remove(output_name)\n    \n    with zipfile.ZipFile(output_name, 'w', zipfile.ZIP_DEFLATED) as zf:\n        for root, dirs, files in os.walk('unified_model'):\n            for file in files:\n                file_path = os.path.join(root, file)\n                arcname = os.path.relpath(file_path, 'unified_model')\n                zf.write(file_path, arcname)\n    \n    zip_size = os.path.getsize(output_name) / 1e6\n    print(f'✓ Package created: {output_name}')\n    print(f'  Size: {zip_size:.1f} MB')\n    print(f'\\nContents:')\n    with zipfile.ZipFile(output_name, 'r') as zf:\n        for info in zf.filelist:\n            print(f'  • {info.filename} ({info.file_size / 1e6:.1f} MB)')\n    \n    return output_name\n\n# Create the package\npackage_name = create_download_package()\n\n# For Kaggle/Colab notebook downloads\nprint('\\n' + '='*60)\nprint('DOWNLOAD INSTRUCTIONS')\nprint('='*60)\nprint(f'\\nFile to download: {package_name}')\nprint('\\nIf in Kaggle notebook:')\nprint(f'  from IPython.display import FileLink')\nprint(f'  FileLink(r\"{package_name}\")')\nprint('\\nIf in Google Colab:')\nprint(f'  from google.colab import files')\nprint(f'  files.download(\"{package_name}\")')\nprint('\\nOtherwise: Right-click file in file explorer and download.')","metadata":{"execution":{"execution_failed":"2026-06-10T08:14:44.705Z"},"trusted":true},"outputs":[],"execution_count":null}]}