{"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.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":71549,"databundleVersionId":8561470},{"sourceType":"modelInstanceVersion","sourceId":854617,"databundleVersionId":17008644,"modelInstanceId":649528,"modelId":661539}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":38936.752009,"end_time":"2026-04-29T12:03:15.973730","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-04-29T01:14:19.221721","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport gc\nimport warnings\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport pydicom as dcm\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedShuffleSplit\nfrom joblib import Parallel, delayed\nimport json\nfrom PIL import Image\nfrom datetime import datetime\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\n\nprint(f\"  PyTorch version : {torch.__version__}\")\nprint(f\"  CUDA available  : {torch.cuda.is_available()}\")\nwarnings.filterwarnings('ignore')\nsns.set_style('whitegrid')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:02.375764Z","iopub.execute_input":"2026-05-01T17:25:02.376915Z","iopub.status.idle":"2026-05-01T17:25:02.382783Z","shell.execute_reply.started":"2026-05-01T17:25:02.376883Z","shell.execute_reply":"2026-05-01T17:25:02.382030Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BASE_PATH = '/kaggle/input/competitions/rsna-2024-lumbar-spine-degenerative-classification'\nTRAIN_CSV = f'{BASE_PATH}/train.csv'\nSERIES_DESC_CSV = f'{BASE_PATH}/train_series_descriptions.csv'\nTRAIN_IMAGES = f'{BASE_PATH}/train_images'\nOUTPUT_DIR = '/kaggle/working'\n\nTRAIN_RATIO = 0.70\nVAL_RATIO   = 0.15\nTEST_RATIO  = 0.15\nRANDOM_SEED = 42\n\nSEVERITY_MAP = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\n\nprint(\"✓ Configuration loaded\")\nprint(f\"\\nData split: {TRAIN_RATIO*100:.0f}% / {VAL_RATIO*100:.0f}% / {TEST_RATIO*100:.0f}%\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:02.384589Z","iopub.execute_input":"2026-05-01T17:25:02.384956Z","iopub.status.idle":"2026-05-01T17:25:02.402712Z","shell.execute_reply.started":"2026-05-01T17:25:02.384935Z","shell.execute_reply":"2026-05-01T17:25:02.402117Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 1: Load Dataset","metadata":{}},{"cell_type":"code","source":"train_df  = pd.read_csv(TRAIN_CSV)\nseries_df = pd.read_csv(SERIES_DESC_CSV)\n\nprint(f\"Training labels: {train_df.shape}\")\nprint(f\"Series description: {series_df.shape}\")\nprint(f\"\\nTotal studies: {train_df['study_id'].nunique()}\")\nprint(f\"Total series:  {len(series_df)}\")\nprint(\"\\nFirst 3 rows:\")\ntrain_df.head(3)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:02.404205Z","iopub.execute_input":"2026-05-01T17:25:02.404434Z","iopub.status.idle":"2026-05-01T17:25:02.443438Z","shell.execute_reply.started":"2026-05-01T17:25:02.404415Z","shell.execute_reply":"2026-05-01T17:25:02.442926Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 2: Analyze Class Distribution","metadata":{}},{"cell_type":"code","source":"stenosis_cols = [col for col in train_df.columns\n                 if ('stenosis' in col.lower() or 'narrowing' in col.lower())]\nprint(f\"Found {len(stenosis_cols)} stenosis classification columns\\n\")\n\ncondition_types = {\n    'Spinal Canal':           [c for c in stenosis_cols if 'spinal_canal'      in c],\n    'Left Neural Foraminal':  [c for c in stenosis_cols if 'left_neural'       in c],\n    'Right Neural Foraminal': [c for c in stenosis_cols if 'right_neural'      in c],\n    'Left Subarticular':      [c for c in stenosis_cols if 'left_subarticular'  in c],\n    'Right Subarticular':     [c for c in stenosis_cols if 'right_subarticular' in c],\n}\n\nfor condition_name, cols in condition_types.items():\n    print(f\"\\n{'='*60}\")\n    print(f\"{condition_name} ({len(cols)} labels)\")\n    print(f\"{'='*60}\")\n    all_values = pd.concat([train_df[col] for col in cols])\n    counts     = all_values.value_counts()\n    print(\"\\nCounts:\")\n    print(counts)\n    print(\"\\nPercentages:\")\n    for severity, count in counts.items():\n        pct = (count / len(all_values)) * 100\n        print(f\"  {severity}: {pct:.1f}%\")\n    if len(counts) > 1:\n        imbalance = counts.max() / counts.min()\n        print(f\"\\n⚠️  Imbalance Ratio: {imbalance:.2f}:1\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:02.444242Z","iopub.execute_input":"2026-05-01T17:25:02.444534Z","iopub.status.idle":"2026-05-01T17:25:02.461613Z","shell.execute_reply.started":"2026-05-01T17:25:02.444515Z","shell.execute_reply":"2026-05-01T17:25:02.460860Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 3: Prepare Multi-Task Labels\n\n**Why Multi-Task?**\n\nInstead of collapsing 25 conditions into one worst-case label, we keep all 25 targets.\nEach condition gets its own classification head, so the attention modules receive\ncondition-specific gradient signals.\n\n**Benefits:**\n- Granular per-condition feedback to CA + SAM attention layers\n- Weighted loss fights the ~18:1 Normal/Mild imbalance per condition\n- Clinically actionable: outputs severity for *each* disc level and condition type\n","metadata":{}},{"cell_type":"code","source":"for col in stenosis_cols:\n    train_df[col + '_label'] = train_df[col].map(SEVERITY_MAP).fillna(-100).astype(int)\n\nLABEL_COLS = [col + '_label' for col in stenosis_cols]\n\nprint(f\"Multi-task label columns: {len(LABEL_COLS)}\")\nprint(\"Sample (first 3 rows, first 5 labels):\")\nprint(train_df[LABEL_COLS[:5]].head(3))\n\ntrain_df['overall_severity'] = train_df[LABEL_COLS].max(axis=1)\nreverse_map = {v: k for k, v in SEVERITY_MAP.items()}\ntrain_df['overall_severity_label'] = train_df['overall_severity'].map(reverse_map)\n\nprint(\"\\nStratification proxy (worst-case) distribution:\")\nprint(train_df['overall_severity_label'].value_counts())\ncounts    = train_df['overall_severity'].value_counts()\nimbalance = counts.max() / counts.min()\nprint(f\"\\nImbalance Ratio: {imbalance:.2f}:1  (handled by weighted loss)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:02.462678Z","iopub.execute_input":"2026-05-01T17:25:02.463018Z","iopub.status.idle":"2026-05-01T17:25:02.501046Z","shell.execute_reply.started":"2026-05-01T17:25:02.462999Z","shell.execute_reply":"2026-05-01T17:25:02.500441Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 4: Visualize Class Distribution","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 2, figsize=(15, 5))\ncolors = ['#2ecc71', '#f39c12', '#e74c3c']\n\nax1 = axes[0]\ncounts = train_df['overall_severity_label'].value_counts()\nprint(counts)\ncounts.plot(kind='bar', ax=ax1, color=colors)\nax1.set_title(\"Overall Severity Distribution\", fontsize=14, fontweight='bold')\nax1.set_xlabel('Severity Level', fontsize=12)\nax1.set_ylabel(\"Number of Studies\", fontsize=12)\nax1.tick_params(axis='x', rotation=45)\n\nax2 = axes[1]\ncounts.plot(kind='pie', ax=ax2, colors=colors, autopct='%1.1f%%', startangle=90)\nax2.set_title('Percentage Distribution', fontsize=14, fontweight='bold')\nax2.set_ylabel(\"\")\n\nplt.tight_layout()\nplt.savefig(f'{OUTPUT_DIR}/class_distribution.png', dpi=300, bbox_inches='tight')\nplt.show()\nprint(\"✓ Saved: class_distribution.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:02.502826Z","iopub.execute_input":"2026-05-01T17:25:02.503360Z","iopub.status.idle":"2026-05-01T17:25:03.186641Z","shell.execute_reply.started":"2026-05-01T17:25:02.503339Z","shell.execute_reply":"2026-05-01T17:25:03.185879Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 5: Create Stratified Splits\n\nSplit Ratios:\n* Training: 70% | Validation: 15% | Test: 15%\n","metadata":{}},{"cell_type":"code","source":"# First split: train vs (val + test)\nsss1 = StratifiedShuffleSplit(n_splits=1, test_size=(VAL_RATIO + TEST_RATIO), random_state=RANDOM_SEED)\nfor train_idx, temp_idx in sss1.split(train_df, train_df['overall_severity']):\n    train_split = train_df.iloc[train_idx].copy()\n    temp_split  = train_df.iloc[temp_idx].copy()\n\n# FIX: second split used wrong variable (sss1 reused) — now correctly sss2\nsss2 = StratifiedShuffleSplit(n_splits=1, test_size=0.5, random_state=RANDOM_SEED)\nfor val_idx, test_idx in sss2.split(temp_split, temp_split['overall_severity']):\n    val_split  = temp_split.iloc[val_idx].copy()\n    test_split = temp_split.iloc[test_idx].copy()\n\ntotal = len(train_df)\nprint(\"Split Results:\")\nprint('=' * 60)\nprint(f\"Training  : {len(train_split):4d} studies ({len(train_split)/total*100:5.1f}%)\")\nprint(f\"Validation: {len(val_split):4d}   studies ({len(val_split)/total*100:5.1f}%)\")\nprint(f\"Test      : {len(test_split):4d}   studies ({len(test_split)/total*100:5.1f}%)\")\nprint(\"=\" * 60)\nprint(f\"Total     : {total:4d}   studies (100.0%)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:03.187496Z","iopub.execute_input":"2026-05-01T17:25:03.187975Z","iopub.status.idle":"2026-05-01T17:25:03.206237Z","shell.execute_reply.started":"2026-05-01T17:25:03.187951Z","shell.execute_reply":"2026-05-01T17:25:03.205374Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 6: Verify No Data Leakage","metadata":{}},{"cell_type":"code","source":"train_ids = set(train_split['study_id'])\nval_ids   = set(val_split['study_id'])\ntest_ids  = set(test_split['study_id'])\n\ntrain_val  = train_ids.intersection(val_ids)\ntrain_test = train_ids.intersection(test_ids)\nval_test   = val_ids.intersection(test_ids)\n\nprint(\"Data Leakage Check:\")\nprint(\"=\" * 60)\n\nall_clear = True\nfor name, overlap in [(\"train ∩ validation\", train_val), (\"train ∩ test\", train_test), (\"validation ∩ test\", val_test)]:\n    if len(overlap) > 0:\n        print(f\"❌ {len(overlap)} studies in both {name}!\")\n        all_clear = False\n    else:\n        print(f\"✓ No overlap: {name}\")\n\ntotal_unique = len(train_ids) + len(val_ids) + len(test_ids)\nif total_unique == total and all_clear:\n    print(\"\\n✅ ALL CHECKS PASSED — No data leakage detected!\")\nelse:\n    print(f\"\\n❌ ERROR: Expected {total}, found {total_unique} unique studies\")\n    raise ValueError(\"Data leakage detected!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:03.207155Z","iopub.execute_input":"2026-05-01T17:25:03.208110Z","iopub.status.idle":"2026-05-01T17:25:03.216441Z","shell.execute_reply.started":"2026-05-01T17:25:03.208088Z","shell.execute_reply":"2026-05-01T17:25:03.215688Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 7: Verify Stratification Quality","metadata":{}},{"cell_type":"code","source":"splits = {'Training': train_split, 'Validation': val_split, 'Test': test_split}\n\nprint(\"Severity Distribution by Split:\")\nprint(\"=\" * 80)\nfor name, split_df in splits.items():\n    print(f\"\\n{name} Set (n={len(split_df)}):\")\n    dist = split_df['overall_severity'].value_counts(normalize=True).sort_index() * 100\n    for sev_code, pct in dist.items():\n        label = reverse_map[sev_code]\n        count = (split_df['overall_severity'] == sev_code).sum()\n        print(f\"  {label:15s}: {count:4d} ({pct:5.1f}%)\")\nprint(\"\\n\" + \"=\" * 80)\nprint(\"✓ Distributions are similar (differences <2% are acceptable)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:03.218315Z","iopub.execute_input":"2026-05-01T17:25:03.218588Z","iopub.status.idle":"2026-05-01T17:25:03.234214Z","shell.execute_reply.started":"2026-05-01T17:25:03.218569Z","shell.execute_reply":"2026-05-01T17:25:03.233593Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 8: Visualize Data Splits","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 3, figsize=(15, 5))\nfor idx, (name, split_df) in enumerate(splits.items()):\n    ax     = axes[idx]\n    counts = split_df['overall_severity'].value_counts().sort_index()\n    labels = [reverse_map[i] for i in counts.index]\n    ax.bar(labels, counts.values, color=['#2ecc71', '#f39c12', '#e74c3c'])\n    ax.set_title(f'{name}\\n(n={len(split_df)})', fontsize=12, fontweight='bold')\n    ax.set_xlabel('Severity')\n    ax.set_ylabel('Count')\n    ax.tick_params(axis='x', rotation=45)\n    for i, v in enumerate(counts.values):\n        pct = (v / len(split_df)) * 100\n        ax.text(i, v, f'{v}\\n({pct:.1f}%)', ha='center', va='bottom', fontsize=9)\nplt.suptitle('Data Split Distribution (Stratified)', fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig(f'{OUTPUT_DIR}/data_splits.png', dpi=300, bbox_inches='tight')\nplt.show()\nprint(\"✓ Saved: data_splits.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:03.234989Z","iopub.execute_input":"2026-05-01T17:25:03.235229Z","iopub.status.idle":"2026-05-01T17:25:04.195285Z","shell.execute_reply.started":"2026-05-01T17:25:03.235183Z","shell.execute_reply":"2026-05-01T17:25:04.194490Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 9: Save Splits","metadata":{}},{"cell_type":"code","source":"train_split.to_csv(f'{OUTPUT_DIR}/train_split.csv', index=False)\nval_split.to_csv(f'{OUTPUT_DIR}/val_split.csv',   index=False)\ntest_split.to_csv(f'{OUTPUT_DIR}/test_split.csv',  index=False)\nprint(\"✓ Saved: train_split.csv\")\nprint(\"✓ Saved: val_split.csv\")\nprint(\"✓ Saved: test_split.csv\")\nprint(\"\\nThese will be used in Phase 2 (Model Training)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:04.196210Z","iopub.execute_input":"2026-05-01T17:25:04.196521Z","iopub.status.idle":"2026-05-01T17:25:04.235121Z","shell.execute_reply.started":"2026-05-01T17:25:04.196492Z","shell.execute_reply":"2026-05-01T17:25:04.234542Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 10: Visualize Sample DICOM Images","metadata":{}},{"cell_type":"code","source":"sample_id    = train_split['study_id'].iloc[0]\nstudy_series = series_df[series_df['study_id'] == sample_id]\n\nprint(f\"Sample Study: {sample_id}\")\nprint(f\"Series available: {len(study_series)}\\n\")\n\ntarget_series = ['Sagittal T2/STIR', 'Axial T2']\nfig, axes = plt.subplots(len(target_series), 5, figsize=(20, 8))\n\nfor row, series_desc in enumerate(target_series):\n    series_match = study_series[study_series['series_description'] == series_desc]\n    if len(series_match) == 0:\n        print(f\"  {series_desc} not found\")\n        continue\n    series_id   = series_match['series_id'].iloc[0]\n    series_path = f'{TRAIN_IMAGES}/{sample_id}/{series_id}'\n    print(f\"  {series_desc}: Series {series_id}\")\n    if os.path.exists(series_path):\n        dcm_files = sorted([f for f in os.listdir(series_path) if f.endswith('.dcm')])\n        print(f\"  Slices: {len(dcm_files)}\")\n        indices = np.linspace(0, len(dcm_files) - 1, 5, dtype=int)\n        for col, idx in enumerate(indices):\n            dcm_data = dcm.dcmread(f'{series_path}/{dcm_files[idx]}')\n            img      = dcm_data.pixel_array\n            ax = axes[row, col]\n            ax.imshow(img, cmap='gray')\n            # FIX: was f'Slice {idx+1/len(dcm_files)}' — operator precedence bug\n            ax.set_title(f'Slice {idx + 1}/{len(dcm_files)}', fontsize=10)\n            ax.axis('off')\n    print()\n\nplt.suptitle(f\"Sample MRI Scans — Study {sample_id}\", fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig(f'{OUTPUT_DIR}/sample_images.png', dpi=300, bbox_inches='tight')\nplt.show()\nprint(\"✓ Saved: sample_images.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:04.236019Z","iopub.execute_input":"2026-05-01T17:25:04.236308Z","iopub.status.idle":"2026-05-01T17:25:10.705241Z","shell.execute_reply.started":"2026-05-01T17:25:04.236278Z","shell.execute_reply":"2026-05-01T17:25:10.704303Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 11 — Phase 2: Data Loading Pipeline\n\nThis phase covers:\n1. DICOM loading and pixel extraction\n2. Image normalization and resizing\n3. Grayscale-to-3-channel conversion\n4. Data augmentation (training set only)\n5. Custom PyTorch Dataset class\n6. DataLoader creation for train / val / test splits\n","metadata":{}},{"cell_type":"code","source":"TRAIN_CSV   = f'{OUTPUT_DIR}/train_split.csv'\nVAL_CSV     = f'{OUTPUT_DIR}/val_split.csv'\nTEST_CSV    = f'{OUTPUT_DIR}/test_split.csv'\n\nIMG_SIZE    = 384\nIN_CHANNELS = 3\nBATCH_SIZE  = 16\nNUM_WORKERS = 4\nPIN_MEMORY  = torch.cuda.is_available()\n\nprint(\"✓ Configuration loaded\")\nprint(f\"  Image size  : {IMG_SIZE}×{IMG_SIZE}\")\nprint(f\"  Batch size  : {BATCH_SIZE}\")\nprint(f\"  Pin memory  : {PIN_MEMORY}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:10.706290Z","iopub.execute_input":"2026-05-01T17:25:10.706606Z","iopub.status.idle":"2026-05-01T17:25:10.712175Z","shell.execute_reply.started":"2026-05-01T17:25:10.706583Z","shell.execute_reply":"2026-05-01T17:25:10.711290Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 12 — DICOM Loading & Preprocessing Function","metadata":{}},{"cell_type":"code","source":"def load_and_preprocess_dicom(dicom_path):\n    \"\"\"\n    Load a single DICOM file and return a preprocessed RGB PIL Image.\n\n    Steps: read → pixel array → float32 → min-max normalise [0,1]\n           → uint8 [0,255] → resize IMG_SIZE × IMG_SIZE → convert RGB\n    \"\"\"\n    dicom       = dcm.dcmread(dicom_path)\n    pixel_array = dicom.pixel_array\n    image       = pixel_array.astype(np.float32)\n\n    min_val = image.min()\n    max_val = image.max()\n    if max_val - min_val > 0:\n        image = (image - min_val) / (max_val - min_val)\n    else:\n        image = np.zeros_like(image)\n\n    # FIX: these two lines were outside the function body due to broken indentation\n    image = (image * 255).astype(np.uint8)\n\n    pil_image = Image.fromarray(image)\n    pil_image = pil_image.resize((IMG_SIZE, IMG_SIZE), resample=Image.BILINEAR)\n    pil_image = pil_image.convert(\"RGB\")\n    return pil_image\n\nprint(\"✓ load_and_preprocess_dicom() defined\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:10.713323Z","iopub.execute_input":"2026-05-01T17:25:10.713835Z","iopub.status.idle":"2026-05-01T17:25:10.728139Z","shell.execute_reply.started":"2026-05-01T17:25:10.713810Z","shell.execute_reply":"2026-05-01T17:25:10.727260Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 13 — Data Augmentation Transforms","metadata":{}},{"cell_type":"code","source":"IMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD  = [0.229, 0.224, 0.225]\n\ntrain_transforms = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.RandomRotation(degrees=5, fill=0),        # was 15\n    transforms.RandomAffine(\n        degrees=0,\n        translate=(0.02, 0.02),   # was 0.05 — only 7px shift on 384px image\n        scale=(0.97, 1.03),       # was 0.90–1.10 — only ~11px zoom\n        fill=0\n    ),\n    transforms.ColorJitter(brightness=0.15, contrast=0.15),  # new, safe for MRI\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n])\n\nval_test_transforms = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n])\n\nprint(\"✓ Transforms defined\")\nprint(\"  train_transforms    : rotation(±5°) + affine(translate+zoom) + normalise\")\nprint(\"  val_test_transforms : resize + normalise only\")\nprint(\"  ✗ No horizontal flip (axial left/right matters clinically)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:10.729273Z","iopub.execute_input":"2026-05-01T17:25:10.729675Z","iopub.status.idle":"2026-05-01T17:25:10.744940Z","shell.execute_reply.started":"2026-05-01T17:25:10.729642Z","shell.execute_reply":"2026-05-01T17:25:10.744230Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 14 — Custom PyTorch Dataset Class","metadata":{}},{"cell_type":"code","source":"class LumbarSpineDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.df        = dataframe.reset_index(drop=True)\n        self.transform = transform\n        self._has_label_cols = all(c in self.df.columns for c in LABEL_COLS)\n        print(f\"  Dataset created: {len(self.df)} samples | label_cols_present={self._has_label_cols}\")\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row   = self.df.iloc[idx]\n        image = load_and_preprocess_dicom(row['dicom_path'])\n\n        if self.transform is not None:\n            image = self.transform(image)\n\n        if self._has_label_cols:\n            labels = row[LABEL_COLS].values.astype('int')\n        else:\n            labels = np.full(NUM_TASKS, int(row['overall_severity']), dtype='int')\n\n        # now returns study_id as third element\n        return image, torch.tensor(labels, dtype=torch.long), str(row['study_id'])\n\nprint(\"✓ LumbarSpineDataset updated — now returns (image, labels, study_id)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:10.745884Z","iopub.execute_input":"2026-05-01T17:25:10.746150Z","iopub.status.idle":"2026-05-01T17:25:10.758096Z","shell.execute_reply.started":"2026-05-01T17:25:10.746121Z","shell.execute_reply":"2026-05-01T17:25:10.757308Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 15 — Build DICOM Path Column (Series-Aware)","metadata":{}},{"cell_type":"code","source":"\n# ── Self-contained prerequisites for this cell ────────────────────────────────\nBASE_PATH    = '/kaggle/input/competitions/rsna-2024-lumbar-spine-degenerative-classification'\nTRAIN_IMAGES = f'{BASE_PATH}/train_images'\ntrain_df = pd.read_csv(TRAIN_CSV)\nval_df   = pd.read_csv(VAL_CSV)\ntest_df  = pd.read_csv(TEST_CSV)\nSEVERITY_MAP = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\n\nseries_df_ref = pd.read_csv(f'{BASE_PATH}/train_series_descriptions.csv')\n\nall_cols      = train_df.columns.tolist()\nstenosis_cols = [c.replace('_label', '') for c in all_cols\n                 if c.endswith('_label') and ('stenosis' in c or 'narrowing' in c)]\nLABEL_COLS    = [c + '_label' for c in stenosis_cols]\nNUM_TASKS     = len(LABEL_COLS)\n\nSERIES_CONDITION_MAP = {\n    'Sagittal T2/STIR': [c for c in stenosis_cols if 'spinal_canal'  in c],\n    'Sagittal T1':      [c for c in stenosis_cols if 'neural'        in c],\n    'Axial T2':         [c for c in stenosis_cols if 'subarticular'  in c],\n}\n\nSLICES_PER_SERIES = {\n    'Sagittal T2/STIR': 5,\n    'Sagittal T1':      7,\n    'Axial T2':         7,\n}\n\nprint(\"✓ All prerequisites defined\")\nprint(f\"  LABEL_COLS    : {NUM_TASKS} tasks\")\nprint(f\"  series_df_ref : {len(series_df_ref)} series\")\nfor k, v in SERIES_CONDITION_MAP.items():\n    print(f\"  {k:<20}: {len(v)} conditions\")\n\n\n# Series-specific slice counts\nSLICES_PER_SERIES = {\n    'Sagittal T2/STIR': 5,   # canal is central — 5 center-biased slices enough\n    'Sagittal T1':      7,   # foraminal pathology is lateral — need wider coverage\n    'Axial T2':         7,   # one slice per disc level — need more levels covered\n}\n\ndef select_slices(files, series_desc):\n    L = len(files)\n    n = SLICES_PER_SERIES.get(series_desc, 5)   # fallback 5 for anything unexpected\n\n    if series_desc == 'Sagittal T1':\n        indices = np.linspace(0, L-1, n, dtype=int)\n    elif series_desc == 'Axial T2':\n        indices = np.linspace(0, L-1, n, dtype=int)\n    else:  # Sagittal T2/STIR\n        mid     = L // 2\n        offsets = [-4, -2, 0, 2, 4]\n        indices = [max(0, min(L-1, mid + o)) for o in offsets]\n    return [files[i] for i in indices]\n\n\ndef add_dicom_paths(df, images_root, series_desc_df):  # removed slices_per_series arg\n    all_rows        = []\n    skipped_series  = 0\n    included_series = 0\n\n    for row in df.itertuples(index=False):\n        study_id   = str(row.study_id)\n        study_path = os.path.join(images_root, study_id)\n        if not os.path.exists(study_path):\n            continue\n\n        study_series = series_desc_df[series_desc_df['study_id'] == int(study_id)]\n\n        for series_row in study_series.itertuples(index=False):\n            series_id   = str(series_row.series_id)\n            series_desc = series_row.series_description\n            series_path = os.path.join(study_path, series_id)\n\n            if not os.path.isdir(series_path):\n                continue\n            if series_desc not in SERIES_CONDITION_MAP:\n                skipped_series += 1\n                continue\n\n            relevant_cols = SERIES_CONDITION_MAP[series_desc]\n            severity_vals = []\n            for col in relevant_cols:\n                val = getattr(row, col, None)\n                if val is not None and pd.notna(val):\n                    severity_vals.append(SEVERITY_MAP.get(val, 0))\n\n            if not severity_vals:\n                skipped_series += 1\n                continue\n\n            label = max(severity_vals)\n\n            files = sorted(\n                [f for f in os.listdir(series_path) if f.endswith('.dcm')],\n                key=lambda x: int(os.path.splitext(x)[0])\n            )\n            if not files:\n                continue\n\n            selected = select_slices(files, series_desc)  # no n arg needed anymore\n\n            for f in selected:\n                try:\n                    instance_number = int(os.path.splitext(f)[0])\n                except ValueError:\n                    instance_number = -1\n\n                row_dict = {\n                    'study_id':         study_id,\n                    'series_id':        series_id,\n                    'series_desc':      series_desc,\n                    'instance_number':  instance_number,\n                    'dicom_path':       os.path.join(series_path, f),\n                    'overall_severity': label,\n                }\n                for lc in LABEL_COLS:\n                    orig_col = lc.replace('_label', '')\n                    raw_val  = getattr(row, orig_col, None)\n                    row_dict[lc] = SEVERITY_MAP.get(raw_val, -100) if (raw_val is not None and pd.notna(raw_val)) else -100\n\n                all_rows.append(row_dict)\n            included_series += 1\n\n    slice_df = pd.DataFrame(all_rows)\n    print(f\"✓ Series-aware expansion complete\")\n    print(f\"  Studies: {len(df)} | Series included: {included_series} | Skipped: {skipped_series}\")\n    print(f\"  Total slices: {len(slice_df)}\")\n    if len(slice_df) > 0:\n        rev = {0: 'Normal/Mild', 1: 'Moderate', 2: 'Severe'}\n        counts = slice_df['overall_severity'].value_counts().sort_index()\n        for code, count in counts.items():\n            print(f\"    {rev[code]:<12}: {count:5d}  ({count/len(slice_df)*100:.1f}%)\")\n    return slice_df\n\n\n# Call without slices_per_series arg now\nprint(\"Building DICOM paths (series-aware, variable slice counts)...\")\ntrain_slice_df = add_dicom_paths(train_df, TRAIN_IMAGES, series_df_ref)\nval_slice_df   = add_dicom_paths(val_df,   TRAIN_IMAGES, series_df_ref)\ntest_slice_df  = add_dicom_paths(test_df,  TRAIN_IMAGES, series_df_ref)\n\nprint(f\"\\nExpected slice increase vs before:\")\nprint(f\"  Before: 3 series × 5 slices = ~15 slices/study\")\nprint(f\"  After : T2/STIR(5) + T1(7) + Axial(7) = ~19 slices/study\")\nprint(f\"  Train set: ~{len(train_df) * 19:,} slices (was ~{len(train_df) * 15:,})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:10.761166Z","iopub.execute_input":"2026-05-01T17:25:10.761413Z","iopub.status.idle":"2026-05-01T17:25:20.760414Z","shell.execute_reply.started":"2026-05-01T17:25:10.761394Z","shell.execute_reply":"2026-05-01T17:25:20.759764Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 16 — Load Split DataFrames and Build Datasets","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(TRAIN_CSV)\nval_df   = pd.read_csv(VAL_CSV)\ntest_df  = pd.read_csv(TEST_CSV)\n\nprint(f\"Splits loaded — Train: {len(train_df)} | Val: {len(val_df)} | Test: {len(test_df)}\")\n\n# Reconstruct stenosis_cols and LABEL_COLS from saved CSV columns\nall_cols      = train_df.columns.tolist()\nstenosis_cols = [c.replace('_label', '') for c in all_cols\n                 if c.endswith('_label') and ('stenosis' in c or 'narrowing' in c)]\nLABEL_COLS    = [c + '_label' for c in stenosis_cols]\nNUM_TASKS     = len(LABEL_COLS)\nprint(f\"  LABEL_COLS ({NUM_TASKS} tasks) reconstructed from CSV\")\n\n# series_df reference for add_dicom_paths\nseries_df_ref = pd.read_csv(\n    '/kaggle/input/competitions/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv'\n)\nTRAIN_IMAGES_PATH = '/kaggle/input/competitions/rsna-2024-lumbar-spine-degenerative-classification/train_images'\n\nprint(\"\\nBuilding DICOM paths (series-aware)...\")\ntrain_slice_df = add_dicom_paths(train_df, TRAIN_IMAGES, series_df_ref)\nval_slice_df   = add_dicom_paths(val_df,   TRAIN_IMAGES, series_df_ref)\ntest_slice_df  = add_dicom_paths(test_df,  TRAIN_IMAGES, series_df_ref)\n\nprint(\"\\nCreating Dataset objects...\")\ntrain_dataset = LumbarSpineDataset(train_slice_df, transform=train_transforms)\nval_dataset   = LumbarSpineDataset(val_slice_df,   transform=val_test_transforms)\ntest_dataset  = LumbarSpineDataset(test_slice_df,  transform=val_test_transforms)\nprint(\"\\n✓ All three datasets created\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:20.761272Z","iopub.execute_input":"2026-05-01T17:25:20.761542Z","iopub.status.idle":"2026-05-01T17:25:27.440414Z","shell.execute_reply.started":"2026-05-01T17:25:20.761507Z","shell.execute_reply":"2026-05-01T17:25:27.439488Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 17 — Create DataLoaders","metadata":{}},{"cell_type":"code","source":"train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True,\n    num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY, drop_last=True,\n    persistent_workers=(NUM_WORKERS > 0),\n    prefetch_factor=2 if NUM_WORKERS > 0 else None)\n\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False,\n    num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY, drop_last=False,\n    persistent_workers=(NUM_WORKERS > 0),\n    prefetch_factor=2 if NUM_WORKERS > 0 else None)\n\ntest_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False,\n    num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY, drop_last=False,\n    persistent_workers=(NUM_WORKERS > 0),\n    prefetch_factor=2 if NUM_WORKERS > 0 else None)\n\nprint(\"✓ DataLoaders created\")\nprint(f\"  Train batches : {len(train_loader)}\")\nprint(f\"  Val   batches : {len(val_loader)}\")\nprint(f\"  Test  batches : {len(test_loader)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:27.441543Z","iopub.execute_input":"2026-05-01T17:25:27.441891Z","iopub.status.idle":"2026-05-01T17:25:27.508381Z","shell.execute_reply.started":"2026-05-01T17:25:27.441867Z","shell.execute_reply":"2026-05-01T17:25:27.507783Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 18 — Sanity Check","metadata":{}},{"cell_type":"code","source":"# Step 18 — Sanity Check (updated for 3-value dataset)\nimages, labels, study_ids = next(iter(train_loader))\n\nprint(\"Batch sanity check:\")\nprint(f\"  Image tensor shape : {images.shape}\")       # expected: (16, 3, 384, 384)\nprint(f\"  Label tensor shape : {labels.shape}\")       # expected: (16, 25)\nprint(f\"  Image dtype        : {images.dtype}\")\nprint(f\"  Label dtype        : {labels.dtype}\")\nprint(f\"  Pixel value range  : [{images.min():.4f}, {images.max():.4f}]\")\nprint(f\"  Unique label values: {labels.unique().tolist()}\")\nprint(f\"  Study IDs (sample) : {study_ids[:3]}\")      # show first 3 study IDs\nprint(\"\\n✓ Batch structure is correct\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:27.509193Z","iopub.execute_input":"2026-05-01T17:25:27.509609Z","iopub.status.idle":"2026-05-01T17:25:28.325774Z","shell.execute_reply.started":"2026-05-01T17:25:27.509586Z","shell.execute_reply":"2026-05-01T17:25:28.324682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def denormalize(tensor, mean=IMAGENET_MEAN, std=IMAGENET_STD):\n    mean   = torch.tensor(mean).view(3, 1, 1)\n    std    = torch.tensor(std).view(3, 1, 1)\n    return (tensor * std + mean).clamp(0, 1)\n\nfig, axes = plt.subplots(2, 4, figsize=(16, 8))\nfig.suptitle(\"Sample Training Images (with Augmentation)\\nWorst-case label shown across 25 conditions\\n0=Normal/Mild | 1=Moderate | 2=Severe\", fontsize=11)\nlabel_names = {0: 'Normal/Mild', 1: 'Moderate', 2: 'Severe'}\n\nfor i, ax in enumerate(axes.flat):\n    if i < len(images):\n        img = denormalize(images[i]).permute(1, 2, 0).cpu().numpy()\n        # FIX: labels is (B,25); show worst-case for display\n        lbl = int(labels[i].max().item())\n        ax.imshow(img, cmap='gray')\n        ax.set_title(f\"Overall: {lbl} ({label_names[lbl]})\", fontsize=9)\n    else:\n        ax.axis('off')\n\nplt.tight_layout()\nplt.savefig(f'{OUTPUT_DIR}/phase2_augmented_samples.png', dpi=150, bbox_inches='tight')\nplt.show()\nprint(\"✓ Saved: phase2_augmented_samples.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:28.328656Z","iopub.execute_input":"2026-05-01T17:25:28.331354Z","iopub.status.idle":"2026-05-01T17:25:31.353250Z","shell.execute_reply.started":"2026-05-01T17:25:28.331307Z","shell.execute_reply":"2026-05-01T17:25:31.352390Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Updated visualization cell\n# images, labels, study_ids = next(iter(train_loader))\n\n# fig, axes = plt.subplots(2, 4, figsize=(16, 8))\n# fig.suptitle(\"Sample Training Images (with Augmentation)\\n0=Normal/Mild | 1=Moderate | 2=Severe\",\n#              fontsize=11)\n# label_names = {0: 'Normal/Mild', 1: 'Moderate', 2: 'Severe'}\n\n# for i, ax in enumerate(axes.flat):\n#     if i < len(images):\n#         img = denormalize(images[i]).permute(1, 2, 0).cpu().numpy()\n#         lbl = int(labels[i][labels[i] != -100].max().item())  # ignore -100\n#         ax.imshow(img, cmap='gray')\n#         ax.set_title(f\"Overall: {lbl} ({label_names[lbl]})\", fontsize=9)\n#     else:\n#         ax.axis('off')\n\n# plt.tight_layout()\n# plt.savefig(f'{OUTPUT_DIR}/phase2_augmented_samples.png', dpi=150, bbox_inches='tight')\n# plt.show()\n# print(\"✓ Saved: phase2_augmented_samples.png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:31.354349Z","iopub.execute_input":"2026-05-01T17:25:31.354716Z","iopub.status.idle":"2026-05-01T17:25:31.358648Z","shell.execute_reply.started":"2026-05-01T17:25:31.354681Z","shell.execute_reply":"2026-05-01T17:25:31.357860Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!nvidia-smi","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:31.359642Z","iopub.execute_input":"2026-05-01T17:25:31.359952Z","iopub.status.idle":"2026-05-01T17:25:31.824643Z","shell.execute_reply.started":"2026-05-01T17:25:31.359931Z","shell.execute_reply":"2026-05-01T17:25:31.823667Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"=\" * 55)\nprint(\"  Phase 2 Preprocessing Pipeline — Complete\")\nprint(\"=\" * 55)\nprint(f\"  Image resolution : {IMG_SIZE} × {IMG_SIZE}\")\nprint(f\"  Batch size       : {BATCH_SIZE}\")\nprint(f\"  Train samples    : {len(train_dataset)}\")\nprint(f\"  Val   samples    : {len(val_dataset)}\")\nprint(f\"  Test  samples    : {len(test_dataset)}\")\nprint(\"=\" * 55)\nprint(\"  Ready for Phase 3: Model Architecture\")\nprint(\"=\" * 55)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:31.826492Z","iopub.execute_input":"2026-05-01T17:25:31.826838Z","iopub.status.idle":"2026-05-01T17:25:31.835291Z","shell.execute_reply.started":"2026-05-01T17:25:31.826802Z","shell.execute_reply":"2026-05-01T17:25:31.834289Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Stage 2: Model Development\n\n### Core architecture components:\n1. EfficientNetB4 — CNN backbone\n2. Coordinate Attention (CA) — channel + spatial position\n3. Spatial Attention Module (SAM) — 2D clinical heatmap\n\nA plain baseline (no attention) is built for comparison.\n","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as models\nfrom torchvision.models import EfficientNet_B4_Weights\n\nNUM_CLASSES              = 3\nDROPOUT_RATE             = 0.4\nIN_CHANNELS              = 3\nIMG_SIZE                 = 384\nEFFICIENTNET_B4_FEATURES = 1792\n\nprint(\"=\" * 60)\nprint(\"  STAGE 2 — MODEL DEVELOPMENT\")\nprint(\"=\" * 60)\nprint(f\"  NUM_CLASSES              : {NUM_CLASSES}\")\nprint(f\"  DROPOUT_RATE             : {DROPOUT_RATE}\")\nprint(f\"  EfficientNetB4 features  : {EFFICIENTNET_B4_FEATURES}\")\nprint(f\"  Input size               : {IMG_SIZE}x{IMG_SIZE}x{IN_CHANNELS}\")\nprint(f\"✓ PyTorch version : {torch.__version__}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:31.836349Z","iopub.execute_input":"2026-05-01T17:25:31.836581Z","iopub.status.idle":"2026-05-01T17:25:31.851276Z","shell.execute_reply.started":"2026-05-01T17:25:31.836547Z","shell.execute_reply":"2026-05-01T17:25:31.850520Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CoordinateAttention(nn.Module):\n    \"\"\"\n    Coordinate Attention — captures WHAT (channels) and WHERE (H, W directions).\n    Decomposes 2D global pooling into two 1D poolings, preserving spatial location.\n    \"\"\"\n\n    def __init__(self, in_channels: int, reduction: int = 32):\n        super(CoordinateAttention, self).__init__()\n        mid_channels = max(8, in_channels // reduction)\n        self.conv1  = nn.Conv2d(in_channels, mid_channels, kernel_size=1, bias=False)\n        self.bn1    = nn.BatchNorm2d(mid_channels)\n        self.act    = nn.Hardswish(inplace=True)\n        self.conv_h = nn.Conv2d(mid_channels, in_channels, kernel_size=1, bias=False)\n        self.conv_w = nn.Conv2d(mid_channels, in_channels, kernel_size=1, bias=False)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        identity = x\n        B, C, H, W = x.shape\n        x_h = F.adaptive_avg_pool2d(x, (H, 1))\n        x_w = F.adaptive_avg_pool2d(x, (1, W)).permute(0, 1, 3, 2)\n        y   = torch.cat([x_h, x_w], dim=2)\n        y   = self.act(self.bn1(self.conv1(y)))\n        x_h_attn, x_w_attn = torch.split(y, [H, W], dim=2)\n        x_w_attn = x_w_attn.permute(0, 1, 3, 2)\n        a_h = torch.sigmoid(self.conv_h(x_h_attn))\n        a_w = torch.sigmoid(self.conv_w(x_w_attn))\n        return identity * a_h * a_w\n\n_ca  = CoordinateAttention(in_channels=1792)\n_inp = torch.zeros(2, 1792, 16, 16)\n_out = _ca(_inp)\nassert _out.shape == _inp.shape\nprint(f\"✓ CoordinateAttention | Input: {list(_inp.shape)} | Output: {list(_out.shape)} | Params: {sum(p.numel() for p in _ca.parameters()):,}\")\ndel _ca, _inp, _out\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:31.852333Z","iopub.execute_input":"2026-05-01T17:25:31.852602Z","iopub.status.idle":"2026-05-01T17:25:31.891234Z","shell.execute_reply.started":"2026-05-01T17:25:31.852574Z","shell.execute_reply":"2026-05-01T17:25:31.890496Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SpatialAttentionModule(nn.Module):\n    \"\"\"\n    Spatial Attention Module (SAM) — CBAM (Woo et al., 2018).\n    Produces a 2D heatmap highlighting WHERE the model focuses.\n    \"\"\"\n\n    def __init__(self, kernel_size: int = 7):\n        super(SpatialAttentionModule, self).__init__()\n        assert kernel_size in (3, 7)\n        self.conv    = nn.Conv2d(2, 1, kernel_size=kernel_size, padding=kernel_size // 2, bias=False)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x: torch.Tensor):\n        avg_out    = torch.mean(x, dim=1, keepdim=True)\n        max_out, _ = torch.max(x, dim=1, keepdim=True)\n        combined   = torch.cat([avg_out, max_out], dim=1)\n        attn_map   = self.sigmoid(self.conv(combined))\n        return x * attn_map, attn_map\n\n_sam = SpatialAttentionModule(kernel_size=7)\n_inp = torch.zeros(2, 1792, 16, 16)\n_out, _map = _sam(_inp)\nassert _out.shape == _inp.shape\nassert _map.shape == (2, 1, 16, 16)\nprint(f\"✓ SpatialAttentionModule | Output: {list(_out.shape)} | Attn map: {list(_map.shape)} | Params: {sum(p.numel() for p in _sam.parameters()):,}\")\ndel _sam, _inp, _out, _map\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:31.892244Z","iopub.execute_input":"2026-05-01T17:25:31.892583Z","iopub.status.idle":"2026-05-01T17:25:31.910025Z","shell.execute_reply.started":"2026-05-01T17:25:31.892562Z","shell.execute_reply":"2026-05-01T17:25:31.909269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class LSS_HybridModel(nn.Module):\n    def __init__(self, num_tasks=NUM_TASKS, num_classes=NUM_CLASSES,\n                 dropout_rate=DROPOUT_RATE, pretrained=True):\n        super(LSS_HybridModel, self).__init__()\n        weights       = EfficientNet_B4_Weights.IMAGENET1K_V1 if pretrained else None\n        self.backbone = models.efficientnet_b4(weights=weights).features\n        self.ca       = CoordinateAttention(in_channels=1792, reduction=32)\n        self.sam      = SpatialAttentionModule(kernel_size=7)\n        self.gap      = nn.AdaptiveAvgPool2d(1)\n        self.neck     = nn.Sequential(\n            nn.Linear(1792, 512),\n            nn.ReLU(),\n            nn.Dropout(0.3)\n        )\n        self.heads    = nn.ModuleList([nn.Linear(512, num_classes) for _ in range(num_tasks)])\n        self.last_attn_map = None\n\n    def forward(self, x):\n        x = self.backbone(x)\n        x = self.ca(x)\n        x, attn = self.sam(x)\n        self.last_attn_map = attn\n        x = self.gap(x).flatten(1)\n        x = self.neck(x)\n        return torch.stack([head(x) for head in self.heads], dim=1)  # [B, T, 3]\n\nprint(\"Building LSS_HybridModel...\")\nhybrid_model = LSS_HybridModel(num_tasks=NUM_TASKS, num_classes=NUM_CLASSES, dropout_rate=DROPOUT_RATE, pretrained=True)\n_x = torch.zeros(2, 3, IMG_SIZE, IMG_SIZE)\nhybrid_model.eval()\nwith torch.no_grad():\n    _logits = hybrid_model(_x)\nassert _logits.shape == (2, NUM_TASKS, NUM_CLASSES)\nprint(f\"✓ LSS_HybridModel | Params: {sum(p.numel() for p in hybrid_model.parameters()):,} | Output: {list(_logits.shape)}\")\ndel _x, _logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:31.910967Z","iopub.execute_input":"2026-05-01T17:25:31.911225Z","iopub.status.idle":"2026-05-01T17:25:32.972646Z","shell.execute_reply.started":"2026-05-01T17:25:31.911194Z","shell.execute_reply":"2026-05-01T17:25:32.971939Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class LSS_BaselineModel(nn.Module):\n    \"\"\"\n    Baseline: plain EfficientNetB4, no attention, single overall_severity output.\n    Trained identically to the hybrid — performance gap isolates CA+SAM contribution.\n    \"\"\"\n\n    def __init__(self, num_classes=3, dropout_rate=0.4, pretrained=True):\n        super(LSS_BaselineModel, self).__init__()\n        weights       = EfficientNet_B4_Weights.IMAGENET1K_V1 if pretrained else None\n        self.backbone = models.efficientnet_b4(weights=weights).features\n        self.gap      = nn.AdaptiveAvgPool2d(1)\n        self.dropout  = nn.Dropout(p=dropout_rate)\n        self.fc       = nn.Linear(EFFICIENTNET_B4_FEATURES, num_classes)\n\n    def forward(self, x):\n        out = self.gap(self.backbone(x)).flatten(1)\n        return self.fc(self.dropout(out))  # [B, 3]\n\nprint(\"Building LSS_BaselineModel...\")\nbaseline_model = LSS_BaselineModel(num_classes=NUM_CLASSES, dropout_rate=DROPOUT_RATE, pretrained=True)\n_x = torch.zeros(2, 3, IMG_SIZE, IMG_SIZE)\nbaseline_model.eval()\nwith torch.no_grad():\n    _logits = baseline_model(_x)\nassert _logits.shape == (2, NUM_CLASSES)\nprint(f\"✓ LSS_BaselineModel | Params: {sum(p.numel() for p in baseline_model.parameters()):,} | Output: {list(_logits.shape)}\")\ndel _x, _logits\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:32.973608Z","iopub.execute_input":"2026-05-01T17:25:32.973946Z","iopub.status.idle":"2026-05-01T17:25:33.881517Z","shell.execute_reply.started":"2026-05-01T17:25:32.973921Z","shell.execute_reply":"2026-05-01T17:25:33.880646Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"hybrid_total   = sum(p.numel() for p in hybrid_model.parameters())\nbaseline_total = sum(p.numel() for p in baseline_model.parameters())\nprint(\"=\" * 60)\nprint(\"  Stage 2 — Model Development Summary\")\nprint(\"=\" * 60)\nprint(f\"  {'Baseline (EfficientNetB4)':<30} {baseline_total:>15,}\")\nprint(f\"  {'Hybrid (+ CA + SAM)':<30} {hybrid_total:>15,}\")\nprint(f\"  {'Attention overhead':<30} {hybrid_total - baseline_total:>15,}\")\nprint(\"=\" * 60)\nprint(\"  Pipeline: Input → EfficientNetB4 → CA → SAM → GAP → Dropout → FC(×25)\")\nprint(\"  ✓ Stage 2 complete — Ready for Phase 3\")\nprint(\"=\" * 60)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:33.882596Z","iopub.execute_input":"2026-05-01T17:25:33.882928Z","iopub.status.idle":"2026-05-01T17:25:33.891877Z","shell.execute_reply.started":"2026-05-01T17:25:33.882900Z","shell.execute_reply":"2026-05-01T17:25:33.891081Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Phase 3: Training and Evaluation","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nfrom sklearn.metrics import accuracy_score, f1_score, confusion_matrix, classification_report, roc_auc_score\nimport numpy as np, matplotlib.pyplot as plt, seaborn as sns\nimport os, time\n\nEPOCHS        = 30\nLEARNING_RATE = 1e-4\nWEIGHT_DECAY  = 1e-5\nNUM_CLASSES   = 3\nRANDOM_SEED   = 42\nPATIENCE      = 8\nSAVE_DIR      = '/kaggle/working'\nCLASS_NAMES   = ['Normal/Mild', 'Moderate', 'Severe']\n\ntorch.manual_seed(RANDOM_SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed(RANDOM_SEED)\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(\"=\" * 60)\nprint(\"  PHASE 3 — TRAINING AND EVALUATION\")\nprint(\"=\" * 60)\nprint(f\"  Device   : {DEVICE}\")\nprint(f\"  Epochs   : {EPOCHS}\")\nprint(f\"  LR       : {LEARNING_RATE}\")\nprint(f\"  Patience : {PATIENCE}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:33.893020Z","iopub.execute_input":"2026-05-01T17:25:33.893360Z","iopub.status.idle":"2026-05-01T17:25:33.906868Z","shell.execute_reply.started":"2026-05-01T17:25:33.893333Z","shell.execute_reply":"2026-05-01T17:25:33.906131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── 1. Recompute weights from actual slice distribution ──────────────────────\ntrain_counts  = torch.tensor([8195, 7490, 6370], dtype=torch.float)\ntotal         = train_counts.sum()\nclass_weights = total / (3 * train_counts)\nclass_weights = class_weights / class_weights.min()\nclass_weights = class_weights.to(DEVICE)\n\nprint(f\"Class weights:\")\nprint(f\"  Normal/Mild : {class_weights[0]:.4f}\")\nprint(f\"  Moderate    : {class_weights[1]:.4f}\")\nprint(f\"  Severe      : {class_weights[2]:.4f}\")\n\n# ── 2. Updated criterion ─────────────────────────────────────────────────────\ncriterion = nn.CrossEntropyLoss(\n    weight        = class_weights,\n    label_smoothing = 0.1,\n    ignore_index  = -100\n)\n\n# ── 3. Loss functions (unchanged logic, works with new criterion) ─────────────\ndef multi_task_loss(outputs, targets):\n    T = outputs.shape[1]\n    losses = []\n    for i in range(T):\n        t = targets[:, i]\n        if (t != -100).any():\n            l = criterion(outputs[:, i, :], t)\n            if not torch.isnan(l):\n                losses.append(l)\n    return sum(losses) / max(len(losses), 1)\n\ndef single_task_loss(outputs, targets):\n    if targets.dim() == 2:\n        targets = targets.max(dim=1).values\n    return criterion(outputs, targets)\n\nprint(\"✓ Criterion updated with balanced weights and ignore_index=-100\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:33.907895Z","iopub.execute_input":"2026-05-01T17:25:33.908304Z","iopub.status.idle":"2026-05-01T17:25:33.922698Z","shell.execute_reply.started":"2026-05-01T17:25:33.908266Z","shell.execute_reply":"2026-05-01T17:25:33.921887Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.amp import autocast, GradScaler\n\nscaler = GradScaler()\n\ndef train_one_epoch(model, loader, optimizer, device, scaler, is_multitask=True):\n    model.train()\n    running_loss, correct, total, num_samples = 0.0, 0, 0, 0\n\n    for batch in loader:\n        images, targets = batch[0].to(device), batch[1].to(device)\n        optimizer.zero_grad()\n\n        with autocast(device_type='cuda' if device.type == 'cuda' else 'cpu'):\n            outputs = model(images)\n            if is_multitask:\n                loss  = multi_task_loss(outputs, targets)\n                preds = torch.argmax(outputs, dim=2)\n                mask  = targets != -100\n                correct += (preds[mask] == targets[mask]).sum().item()\n                total   += mask.sum().item()\n            else:\n                loss  = single_task_loss(outputs, targets)\n                preds = torch.argmax(outputs, dim=1)\n                tgt   = targets.max(dim=1).values if targets.dim() == 2 else targets\n                correct += (preds == tgt).sum().item()\n                total   += tgt.numel()\n\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += loss.item() * images.size(0)\n        num_samples  += images.size(0)\n\n    return running_loss / max(num_samples, 1), correct / max(total, 1)\n\n\ndef evaluate(model, loader, device, is_multitask=True):\n    model.eval()\n    running_loss, correct, total, num_samples = 0.0, 0, 0, 0\n    all_preds, all_labels, all_probs = [], [], []\n\n    with torch.no_grad():\n        for batch in loader:\n            images, targets = batch[0].to(device), batch[1].to(device)\n            outputs = model(images)\n\n            if is_multitask:\n                loss  = multi_task_loss(outputs, targets)\n                probs = torch.softmax(outputs, dim=2)\n                preds = probs.argmax(dim=2)\n                mask  = targets != -100\n                correct += (preds[mask] == targets[mask]).sum().item()\n                total   += mask.sum().item()\n                all_preds.extend(preds.cpu().numpy().flatten())\n                all_labels.extend(targets.cpu().numpy().flatten())\n                all_probs.extend(probs.cpu().numpy().reshape(-1, NUM_CLASSES))\n            else:\n                loss  = single_task_loss(outputs, targets)\n                probs = torch.softmax(outputs, dim=1)\n                preds = probs.argmax(dim=1)\n                tgt   = targets.max(dim=1).values if targets.dim() == 2 else targets\n                correct += (preds == tgt).sum().item()\n                total   += tgt.numel()\n                all_preds.extend(preds.cpu().numpy())\n                all_labels.extend(tgt.cpu().numpy())\n                all_probs.extend(probs.cpu().numpy())\n\n            running_loss += loss.item() * images.size(0)\n            num_samples  += images.size(0)\n\n    # filter out -100 from metrics\n    all_preds  = np.array(all_preds)\n    all_labels = np.array(all_labels)\n    all_probs  = np.array(all_probs)\n    valid      = all_labels != -100\n    return (running_loss / max(num_samples, 1), correct / max(total, 1),\n            all_preds[valid], all_labels[valid], all_probs[valid])\n\n\ndef train_model(model, train_loader, val_loader, optimizer, scheduler,\n                device, epochs, patience, model_name, save_dir, is_multitask=True):\n\n    history   = {'train_loss': [], 'val_loss': [], 'train_acc': [], 'val_acc': []}\n    best_val  = float('inf')\n    no_improv = 0\n    best_path = os.path.join(save_dir, f'{model_name}_best.pth')\n\n    print(f\"\\n  Training {model_name}\")\n    print(f\"  {'Epoch':<6} {'Train Loss':<12} {'Train Acc':<12} {'Val Loss':<12} {'Val Acc':<10}\")\n    print(f\"  {'-'*54}\")\n\n    for epoch in range(1, epochs + 1):\n        t0 = time.time()\n        tr_loss, tr_acc = train_one_epoch(model, train_loader, optimizer, device, scaler, is_multitask)\n        vl_loss, vl_acc, _, _, _ = evaluate(model, val_loader, device, is_multitask)\n\n        # FIX: compatible with both ReduceLROnPlateau and CosineAnnealingLR\n        try:\n            scheduler.step(vl_loss)\n        except TypeError:\n            scheduler.step()\n\n        history['train_loss'].append(tr_loss)\n        history['val_loss'].append(vl_loss)\n        history['train_acc'].append(tr_acc)\n        history['val_acc'].append(vl_acc)\n\n        marker = \" ✓ best\" if vl_loss < best_val else \"\"\n        print(f\"  {epoch:<6} {tr_loss:<12.4f} {tr_acc:<12.4f} {vl_loss:<12.4f} {vl_acc:<10.4f}  [{time.time()-t0:.0f}s]{marker}\")\n\n        if vl_loss < best_val:\n            best_val  = vl_loss\n            no_improv = 0\n            torch.save(model.state_dict(), best_path)\n        else:\n            no_improv += 1\n            if no_improv >= patience:\n                print(f\"\\n  Early stopping at epoch {epoch}\")\n                break\n\n    # Replace the last 3 lines of train_model with this:\n    if os.path.exists(best_path):\n        model.load_state_dict(torch.load(best_path, map_location=device, weights_only=True))\n        print(f\"\\n  Best val loss: {best_val:.4f} | Weights restored from: {best_path}\")\n    else:\n        print(f\"\\n  ⚠ No checkpoint saved (all losses were nan) — returning current weights\")\n    return model, history\n\nprint(\"✓ train_one_epoch(), evaluate(), train_model() defined and fixed\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:33.923716Z","iopub.execute_input":"2026-05-01T17:25:33.924253Z","iopub.status.idle":"2026-05-01T17:25:33.944168Z","shell.execute_reply.started":"2026-05-01T17:25:33.924232Z","shell.execute_reply":"2026-05-01T17:25:33.943275Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(train_dataset), len(val_dataset), len(test_dataset))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:33.945959Z","iopub.execute_input":"2026-05-01T17:25:33.946232Z","iopub.status.idle":"2026-05-01T17:25:33.957701Z","shell.execute_reply.started":"2026-05-01T17:25:33.946212Z","shell.execute_reply":"2026-05-01T17:25:33.957015Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ── Baseline Model — fair comparison ─────────────────────────────────────────\n# baseline_model = LSS_BaselineModel(\n#     num_classes=NUM_CLASSES, dropout_rate=DROPOUT_RATE, pretrained=True\n# ).to(DEVICE)\n\n# # Differential LR — same principle as hybrid, just no CA/SAM/neck\n# optimizer_b = torch.optim.AdamW([\n#     {'params': baseline_model.backbone.parameters(), 'lr': 1e-5},\n#     {'params': baseline_model.fc.parameters(),       'lr': 1e-4},\n# ], weight_decay=1e-5)\n\n# scheduler_b = torch.optim.lr_scheduler.CosineAnnealingLR(\n#     optimizer_b, T_max=EPOCHS, eta_min=1e-7\n# )\n\n# print(\"=\" * 60)\n# print(\"  Baseline: single-stage, differential LR, CosineAnnealingLR\")\n# print(\"=\" * 60)\n\n# baseline_model, baseline_history = train_model(\n#     baseline_model, train_loader, val_loader,\n#     optimizer_b, scheduler_b,\n#     DEVICE, EPOCHS, PATIENCE,\n#     'baseline', SAVE_DIR,\n#     is_multitask=False\n# )\n\n# torch.save({\n#     'model_state_dict': baseline_model.state_dict(),\n#     'history':          baseline_history,\n#     'best_val_loss':    min(baseline_history['val_loss']),\n#     'total_epochs':     len(baseline_history['val_loss']),\n# }, '/kaggle/working/baseline_best_final.pth')\n\n# print(\"✓ Baseline training complete.\")\n# print(\"✓ Saved: /kaggle/working/baseline_best_final.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:33.958806Z","iopub.execute_input":"2026-05-01T17:25:33.959146Z","iopub.status.idle":"2026-05-01T17:25:33.969185Z","shell.execute_reply.started":"2026-05-01T17:25:33.959116Z","shell.execute_reply":"2026-05-01T17:25:33.968543Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nbaseline_model = LSS_BaselineModel(\n    num_classes=NUM_CLASSES, dropout_rate=DROPOUT_RATE, pretrained=False\n).to(DEVICE)\n\nckpt_b = torch.load('/kaggle/input/models/adebayo22/lss-baseline/pytorch/default/1/baseline_best_final.pth', map_location=DEVICE, weights_only=True)\nbaseline_model.load_state_dict(ckpt_b['model_state_dict'])\nbaseline_model.eval()\nbaseline_history = ckpt_b['history']\n\nprint(f\"✓ Baseline loaded — best val loss: {ckpt_b['best_val_loss']:.4f}\")\nprint(f\"  Trained for {ckpt_b['total_epochs']} epochs\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:33.970102Z","iopub.execute_input":"2026-05-01T17:25:33.970374Z","iopub.status.idle":"2026-05-01T17:25:34.483098Z","shell.execute_reply.started":"2026-05-01T17:25:33.970345Z","shell.execute_reply":"2026-05-01T17:25:34.482359Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"hybrid_model = LSS_HybridModel(\n    num_tasks=NUM_TASKS, num_classes=NUM_CLASSES,\n    dropout_rate=DROPOUT_RATE, pretrained=True\n).to(DEVICE)\n\noptimizer = torch.optim.AdamW([\n    {'params': hybrid_model.backbone.parameters(), 'lr': 1e-5},\n    {'params': hybrid_model.ca.parameters(),       'lr': 5e-5},\n    {'params': hybrid_model.sam.parameters(),      'lr': 5e-5},\n    {'params': hybrid_model.neck.parameters(),     'lr': 5e-5},\n    {'params': hybrid_model.heads.parameters(),    'lr': 5e-5},\n], weight_decay=1e-5)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer, T_max=EPOCHS, eta_min=1e-7\n)\n\nprint(\"✓ Hybrid optimizer — reduced LR to 5e-5 to prevent NaN explosion\")\n\nprint(\"✓ Single-stage training — differential LR, no frozen backbone\")\nprint(\"  backbone lr : 1e-5  (fine-tune gently)\")\nprint(\"  CA/SAM/neck : 1e-4  (train freely)\")\n\nhybrid_model, hybrid_history = train_model(\n    hybrid_model, train_loader, val_loader,\n    optimizer, scheduler,\n    DEVICE, EPOCHS, PATIENCE,\n    'hybrid', SAVE_DIR,\n    is_multitask=True\n)\n\ntorch.save({\n    'model_state_dict': hybrid_model.state_dict(),\n    'history':          hybrid_history,\n    'best_val_loss':    min(hybrid_history['val_loss']),\n    'total_epochs':     len(hybrid_history['val_loss']),\n}, '/kaggle/working/hybrid_best_final.pth')\n\nprint(\"✓ Saved: hybrid_best_final.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T17:25:34.483998Z","iopub.execute_input":"2026-05-01T17:25:34.484272Z","execution_failed":"2026-05-01T17:40:23.449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Optional: Load from saved checkpoints (uncomment to use) ──────────────────\n# CHECKPOINT_DIR = '/kaggle/input/lss-model-checkpoints'\n# DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n# baseline_model = LSS_BaselineModel(num_classes=NUM_CLASSES, dropout_rate=DROPOUT_RATE, pretrained=False)\n# ckpt_b = torch.load(f'{CHECKPOINT_DIR}/baseline_best_final.pth', map_location=DEVICE, weights_only=True)\n# baseline_model.load_state_dict(ckpt_b['model_state_dict'])\n# baseline_model = baseline_model.to(DEVICE).eval()\n# baseline_history = ckpt_b['history']\n# print(f\"✓ Baseline loaded — best val loss: {ckpt_b['best_val_loss']:.4f}\")\n# hybrid_model = LSS_HybridModel(num_tasks=NUM_TASKS, num_classes=NUM_CLASSES, dropout_rate=DROPOUT_RATE, pretrained=False)\n# ckpt_h = torch.load(f'{CHECKPOINT_DIR}/hybrid_best_final.pth', map_location=DEVICE, weights_only=True)\n# hybrid_model.load_state_dict(ckpt_h['model_state_dict'])\n# hybrid_model = hybrid_model.to(DEVICE).eval()\n# hybrid_history = ckpt_h['history']\n# print(f\"✓ Hybrid loaded — best val loss: {ckpt_h['best_val_loss']:.4f}\")\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-01T17:40:23.455Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_training_curves(baseline_hist, hybrid_hist, save_dir):\n    fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n    fig.suptitle(\"Training Curves: Baseline vs Hybrid\", fontsize=14, fontweight='bold')\n    eb = range(1, len(baseline_hist['train_loss']) + 1)\n    eh = range(1, len(hybrid_hist['train_loss'])   + 1)\n    axes[0].plot(eb, baseline_hist['train_loss'], 'b--', label='Baseline Train', alpha=0.7)\n    axes[0].plot(eb, baseline_hist['val_loss'],   'b-',  label='Baseline Val',   linewidth=2)\n    axes[0].plot(eh, hybrid_hist['train_loss'],   'r--', label='Hybrid Train',   alpha=0.7)\n    axes[0].plot(eh, hybrid_hist['val_loss'],     'r-',  label='Hybrid Val',     linewidth=2)\n    axes[0].set(xlabel='Epoch', ylabel='Loss', title='Loss'); axes[0].legend(); axes[0].grid(True, alpha=0.3)\n    axes[1].plot(eb, baseline_hist['train_acc'], 'b--', alpha=0.7)\n    axes[1].plot(eb, baseline_hist['val_acc'],   'b-',  linewidth=2)\n    axes[1].plot(eh, hybrid_hist['train_acc'],   'r--', alpha=0.7)\n    axes[1].plot(eh, hybrid_hist['val_acc'],     'r-',  linewidth=2)\n    axes[1].set(xlabel='Epoch', ylabel='Accuracy', title='Accuracy'); axes[1].legend(['Baseline Train','Baseline Val','Hybrid Train','Hybrid Val']); axes[1].grid(True, alpha=0.3)\n    plt.tight_layout()\n    p = os.path.join(save_dir, 'phase3_training_curves.png')\n    plt.savefig(p, dpi=150, bbox_inches='tight'); plt.show()\n    print(f\"✓ Saved: {p}\")\n\nplot_training_curves(baseline_history, hybrid_history, SAVE_DIR)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-01T17:40:23.455Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from collections import defaultdict\n\ndef evaluate_study_level(model, loader, device, class_names):\n    model.eval()\n    study_probs  = defaultdict(list)\n    study_labels = {}\n\n    with torch.no_grad():\n        for batch in loader:\n            images, targets, study_ids = batch[0].to(device), batch[1], batch[2]\n            outputs = model(images)\n            probs   = torch.softmax(outputs, dim=2).cpu()  # [B, 25, 3]\n            for i, sid in enumerate(study_ids):\n                study_probs[sid].append(probs[i])\n                study_labels[sid] = targets[i]\n\n    all_preds, all_labels = [], []\n    for sid in study_probs:\n        avg_probs = torch.stack(study_probs[sid]).mean(0)   # [25, 3]\n        pred      = avg_probs.argmax(-1).max().item()        # worst-case across heads\n        label     = study_labels[sid].max().item()\n        if label == -100:\n            continue\n        all_preds.append(pred)\n        all_labels.append(label)\n\n    all_preds  = np.array(all_preds)\n    all_labels = np.array(all_labels)\n\n    acc        = (all_preds == all_labels).mean()\n    f1_macro   = f1_score(all_labels, all_preds, average='macro',    zero_division=0)\n    f1_weighted= f1_score(all_labels, all_preds, average='weighted', zero_division=0)\n    try:\n        # need probs for AUC — collect per-study averaged worst-case probs\n        study_worst_probs = []\n        for sid in study_probs:\n            avg = torch.stack(study_probs[sid]).mean(0)   # [25, 3]\n            study_worst_probs.append(avg.max(0).values.numpy())  # [3]\n        auc = roc_auc_score(all_labels,\n                            np.array(study_worst_probs)[:len(all_labels)],\n                            multi_class='ovr', average='macro')\n    except ValueError:\n        auc = float('nan')\n\n    print(f\"\\n  ── Study-level Test Results ──\")\n    print(f\"  Studies evaluated : {len(all_preds)}\")\n    print(f\"  Accuracy          : {acc*100:.2f}%\")\n    print(f\"  F1 macro          : {f1_macro:.4f}\")\n    print(f\"  F1 weighted       : {f1_weighted:.4f}\")\n    print(f\"  AUC-ROC           : {auc:.4f}\")\n    print(f\"\\n{classification_report(all_labels, all_preds, target_names=class_names, zero_division=0)}\")\n\n    return all_preds, all_labels\n\npreds, labels = evaluate_study_level(hybrid_model, test_loader, DEVICE, CLASS_NAMES)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-01T17:40:23.456Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def full_evaluation(model, loader, device, class_names, model_name, is_multitask=True):\n    \"\"\"Test-set evaluation. FIX: evaluate() no longer takes criterion as arg.\"\"\"\n    loss, acc, preds, labels, probs = evaluate(model, loader, device, is_multitask)\n\n    f1_macro     = f1_score(labels, preds, average='macro',    zero_division=0)\n    f1_weighted  = f1_score(labels, preds, average='weighted', zero_division=0)\n    f1_per_class = f1_score(labels, preds, average=None,       zero_division=0)\n    try:\n        auc = roc_auc_score(labels, probs, multi_class='ovr', average='macro')\n    except ValueError:\n        auc = float('nan')\n    cm = confusion_matrix(labels, preds)\n\n    print(f\"\\n  ── {model_name} — Test Results ──\")\n    print(f\"  Accuracy      : {acc:.4f}  ({acc*100:.2f}%)\")\n    print(f\"  F1 (macro)    : {f1_macro:.4f}\")\n    print(f\"  F1 (weighted) : {f1_weighted:.4f}\")\n    print(f\"  AUC-ROC (OvR) : {auc:.4f}\")\n    print(f\"\\n  Per-class F1:\")\n    for name, f1 in zip(class_names, f1_per_class):\n        print(f\"    [{name:<12}] F1 = {f1:.4f}\")\n    print(f\"\\n  Classification Report:\")\n    print(classification_report(labels, preds, target_names=class_names, zero_division=0))\n\n    return {'model': model_name, 'accuracy': acc, 'f1_macro': f1_macro,\n            'f1_weighted': f1_weighted, 'auc': auc, 'cm': cm,\n            'preds': preds, 'labels': labels, 'probs': probs, 'f1_per_class': f1_per_class}\n\nprint(\"Evaluating on test set...\")\nbaseline_results = full_evaluation(baseline_model, test_loader, DEVICE, CLASS_NAMES, 'Baseline', is_multitask=False)\nhybrid_results   = full_evaluation(hybrid_model,   test_loader, DEVICE, CLASS_NAMES, 'Hybrid (CA+SAM)', is_multitask=True)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-01T17:40:23.456Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_confusion_matrix(cm, class_names, model_name, ax):\n    cm_pct = cm.astype(float) / cm.sum(axis=1, keepdims=True) * 100\n    sns.heatmap(cm_pct, annot=True, fmt='.1f', cmap='Blues',\n                xticklabels=class_names, yticklabels=class_names, ax=ax, linewidths=0.5)\n    ax.set_xlabel('Predicted', fontsize=11)\n    ax.set_ylabel('Actual',    fontsize=11)\n    ax.set_title(f'{model_name}\\nConfusion Matrix (% of actual class)', fontsize=12)\n\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\nplot_confusion_matrix(baseline_results['cm'], CLASS_NAMES, 'Baseline',        axes[0])\nplot_confusion_matrix(hybrid_results['cm'],   CLASS_NAMES, 'Hybrid (CA+SAM)', axes[1])\nplt.tight_layout()\np = os.path.join(SAVE_DIR, 'phase3_confusion_matrices.png')\nplt.savefig(p, dpi=150, bbox_inches='tight'); plt.show()\nprint(f\"✓ Saved: {p}\")\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-01T17:40:23.461Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"=\" * 65)\nprint(\"  FINAL RESULTS SUMMARY\")\nprint(\"=\" * 65)\nprint(f\"  {'Model':<25} {'Accuracy':>10} {'F1 Macro':>10} {'AUC-ROC':>10}\")\nprint(f\"  {'-'*57}\")\nfor r in [baseline_results, hybrid_results]:\n    print(f\"  {r['model']:<25} {r['accuracy']*100:>9.2f}% {r['f1_macro']:>10.4f} {r['auc']:>10.4f}\")\nprint(\"=\" * 65)\ndelta_acc = (hybrid_results['accuracy'] - baseline_results['accuracy']) * 100\ndelta_f1  = hybrid_results['f1_macro']  - baseline_results['f1_macro']\nprint(f\"\\n  Hybrid improvement over Baseline:\")\nprint(f\"    Accuracy delta : {delta_acc:+.2f}%\")\nprint(f\"    F1 macro delta : {delta_f1:+.4f}\")\nprint(\"=\" * 65)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-01T17:40:23.462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class LSS_CAOnly(nn.Module):\n    \"\"\"EfficientNetB4 + Coordinate Attention only, multi-task.\"\"\"\n    def __init__(self, num_tasks=NUM_TASKS, num_classes=NUM_CLASSES,\n                 dropout_rate=DROPOUT_RATE, pretrained=True):\n        super().__init__()\n        weights       = EfficientNet_B4_Weights.IMAGENET1K_V1 if pretrained else None\n        self.backbone = models.efficientnet_b4(weights=weights).features\n        self.ca       = CoordinateAttention(in_channels=1792)\n        self.gap      = nn.AdaptiveAvgPool2d(1)\n        self.neck     = nn.Sequential(\n            nn.Linear(1792, 512), nn.ReLU(), nn.Dropout(0.3)\n        )\n        self.heads    = nn.ModuleList([nn.Linear(512, num_classes) for _ in range(num_tasks)])\n\n    def forward(self, x):\n        x = self.ca(self.backbone(x))\n        x = self.gap(x).flatten(1)\n        x = self.neck(x)\n        return torch.stack([head(x) for head in self.heads], dim=1)  # [B, T, 3]\n\n\nclass LSS_SAMOnly(nn.Module):\n    \"\"\"EfficientNetB4 + Spatial Attention only, multi-task.\"\"\"\n    def __init__(self, num_tasks=NUM_TASKS, num_classes=NUM_CLASSES,\n                 dropout_rate=DROPOUT_RATE, pretrained=True):\n        super().__init__()\n        weights            = EfficientNet_B4_Weights.IMAGENET1K_V1 if pretrained else None\n        self.backbone      = models.efficientnet_b4(weights=weights).features\n        self.sam           = SpatialAttentionModule(kernel_size=7)\n        self.gap           = nn.AdaptiveAvgPool2d(1)\n        self.neck          = nn.Sequential(\n            nn.Linear(1792, 512), nn.ReLU(), nn.Dropout(0.3)\n        )\n        self.heads         = nn.ModuleList([nn.Linear(512, num_classes) for _ in range(num_tasks)])\n        self.last_attn_map = None\n\n    def forward(self, x):\n        x, attn            = self.sam(self.backbone(x))\n        self.last_attn_map = attn.detach()\n        x = self.gap(x).flatten(1)\n        x = self.neck(x)\n        return torch.stack([head(x) for head in self.heads], dim=1)  # [B, T, 3]\n\nprint(\"✓ LSS_CAOnly and LSS_SAMOnly defined — multi-task, with neck\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-01T17:40:23.462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_ablation_variant(model_cls, name, train_loader, val_loader,\n                            test_loader, device, epochs, patience, save_dir):\n    model = model_cls(\n        num_tasks=NUM_TASKS, num_classes=NUM_CLASSES,\n        dropout_rate=DROPOUT_RATE, pretrained=True\n    ).to(device)\n\n    # Same differential LR as hybrid — fair comparison\n    optimizer = torch.optim.AdamW([\n        {'params': model.backbone.parameters(), 'lr': 1e-5},\n        {'params': model.neck.parameters(),     'lr': 1e-4},\n        {'params': model.heads.parameters(),    'lr': 1e-4},\n    ], weight_decay=1e-5)\n\n    # CA/SAM layers get added only if they exist\n    if hasattr(model, 'ca'):\n        optimizer.add_param_group({'params': model.ca.parameters(), 'lr': 1e-4})\n    if hasattr(model, 'sam'):\n        optimizer.add_param_group({'params': model.sam.parameters(), 'lr': 1e-4})\n\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer, T_max=epochs, eta_min=1e-7\n    )\n\n    print(f\"\\n  Training ablation variant: {name}\")\n    model, _ = train_model(\n        model, train_loader, val_loader,\n        optimizer, scheduler,\n        device, epochs, patience,\n        name, save_dir,\n        is_multitask=True      # both variants are multi-task\n    )\n\n    results = full_evaluation(model, test_loader, device, CLASS_NAMES, name,\n                               is_multitask=True)\n    del model\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    return results\n\nprint(\"✓ train_ablation_variant() fixed — cosine LR, multi-task, no criterion arg\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-01T17:40:23.462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"=\" * 60)\nprint(\"  Ablation Study — CA-only and SAM-only variants\")\nprint(\"=\" * 60)\n\nca_only_results  = train_ablation_variant(\n    LSS_CAOnly,  'CA-only',\n    train_loader, val_loader, test_loader,\n    DEVICE, EPOCHS, PATIENCE, SAVE_DIR\n)\n\nsam_only_results = train_ablation_variant(\n    LSS_SAMOnly, 'SAM-only',\n    train_loader, val_loader, test_loader,\n    DEVICE, EPOCHS, PATIENCE, SAVE_DIR\n)\n\nprint(\"\\n✓ Ablation variants trained.\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-01T17:40:23.468Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Collect all four results in order\nall_results  = [baseline_results, ca_only_results, sam_only_results, hybrid_results]\nmodel_labels = ['Baseline\\n(no attn)', 'CA-only', 'SAM-only', 'CA+SAM\\n(Hybrid)']\ncolors       = ['#4e79a7', '#59a14f', '#e15759', '#f28e2b']\n\n# ── Ablation table ────────────────────────────────────────────────────────────\nprint(\"=\" * 70)\nprint(\"  Ablation Study Results\")\nprint(\"=\" * 70)\nprint(f\"  {'Model':<22} {'Accuracy':>10} {'F1 Macro':>10} {'F1 Weighted':>12} {'AUC':>8}\")\nprint(f\"  {'-'*64}\")\nfor r in all_results:\n    print(f\"  {r['model']:<22} {r['accuracy']:>10.4f} {r['f1_macro']:>10.4f} \"\n          f\"{r['f1_weighted']:>12.4f} {r['auc']:>8.4f}\")\nprint(\"=\" * 70)\n\nprint(\"\\n  Per-class F1 Breakdown:\")\nprint(f\"  {'Model':<22} {'Normal/Mild':>13} {'Moderate':>10} {'Severe':>8}\")\nprint(f\"  {'-'*55}\")\nfor r in all_results:\n    f = r['f1_per_class']\n    print(f\"  {r['model']:<22} {f[0]:>13.4f} {f[1]:>10.4f} {f[2]:>8.4f}\")\n\n# ── Bar chart ─────────────────────────────────────────────────────────────────\nfig, axes = plt.subplots(1, 3, figsize=(16, 5))\nfig.suptitle(\"Ablation Study: Component Contribution\",\n             fontsize=14, fontweight='bold')\n\nfor ax, metric, title in zip(\n    axes,\n    ['accuracy', 'f1_macro', 'auc'],\n    ['Accuracy', 'F1 Score (Macro)', 'AUC-ROC (OvR Macro)']\n):\n    values = [r[metric] for r in all_results]\n    bars   = ax.bar(model_labels, values, color=colors,\n                    edgecolor='white', linewidth=1.2)\n    ax.set_ylim(max(0, min(values) - 0.05), min(1.0, max(values) + 0.05))\n    ax.set_title(title, fontweight='bold')\n    ax.grid(axis='y', alpha=0.3)\n    for bar, val in zip(bars, values):\n        ax.text(bar.get_x() + bar.get_width()/2,\n                bar.get_height() + 0.002,\n                f'{val:.3f}', ha='center', va='bottom', fontsize=9)\n\nplt.tight_layout()\np = os.path.join(SAVE_DIR, 'phase3_ablation_study.png')\nplt.savefig(p, dpi=150, bbox_inches='tight')\nplt.show()\nprint(f\"✓ Saved: {p}\")\n\n# ── Final summary ─────────────────────────────────────────────────────────────\nprint(\"=\" * 60)\nprint(\"  Phase 3 — Training & Evaluation Complete\")\nprint(\"=\" * 60)\nb = baseline_results\nh = hybrid_results\nprint(f\"  Baseline accuracy : {b['accuracy']*100:.2f}%\")\nprint(f\"  Hybrid   accuracy : {h['accuracy']*100:.2f}%\")\nprint(f\"  Improvement       : +{(h['accuracy']-b['accuracy'])*100:.2f} pp\")\nprint(f\"  F1 improvement    : +{h['f1_macro']-b['f1_macro']:.4f} (macro)\")\nprint()\nprint(\"  Saved artefacts:\")\nfor fname in ['baseline_best_final.pth', 'hybrid_best_final.pth',\n              'phase3_training_curves.png',\n              'phase3_confusion_matrices.png',\n              'phase3_ablation_study.png']:\n    print(f\"    ✓ {fname}\")\nprint(\"=\" * 60)\nprint(\"  ✓ Ready for Phase 4: Interpretability (SAM heatmaps)\")\nprint(\"=\" * 60)","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-01T17:40:23.468Z"}},"outputs":[],"execution_count":null}]}