{"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":"none","dataSources":[],"isGpuEnabled":false,"isInternetEnabled":true,"language":"python","sourceType":"notebook"},"papermill":{"default_parameters":{},"duration":3.782841,"end_time":"2026-03-16T12:41:40.190492","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-03-16T12:41:36.407651","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nn_files = sum(len(filenames) for _, _, filenames in os.walk('/kaggle/input'))\nprint(f'{n_files} files available under /kaggle/input (listing suppressed to keep the notebook responsive)')","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2026-08-11T01:12:45.386803Z","iopub.execute_input":"2026-08-11T01:12:45.387241Z","iopub.status.idle":"2026-08-11T01:14:55.505484Z","shell.execute_reply.started":"2026-08-11T01:12:45.387213Z","shell.execute_reply":"2026-08-11T01:14:55.504648Z"},"papermill":{"duration":0.828767,"end_time":"2026-03-16T12:41:39.772163","exception":false,"start_time":"2026-03-16T12:41:38.943396","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Spine-Specific Multimodal Foundation Models for Lumbar MRI\n\nThree-stage progressive pipeline for the RSNA 2024 Lumbar Spine Degenerative Classification (LumbarDISC) dataset: a text-only baseline, a 2D multi-view fusion model, and a 3D volumetric fusion model. Each stage predicts severity for all 25 condition-level targets and generates a pseudo radiology report; the ablation across stages measures the marginal value of imaging and, separately, of volumetric context.","metadata":{}},{"cell_type":"code","source":"# RSNA DICOM series are JPEG-lossless / JPEG2000 compressed; pydicom needs\n# these decoder plugins installed or pixel_array raises on read. sentencepiece\n# is required by T5Tokenizer and isn't preinstalled on this image.\n!pip install -q pydicom pylibjpeg pylibjpeg-libjpeg python-gdcm sentencepiece","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:14:55.507022Z","iopub.execute_input":"2026-08-11T01:14:55.507260Z","iopub.status.idle":"2026-08-11T01:15:02.392856Z","shell.execute_reply.started":"2026-08-11T01:14:55.507240Z","shell.execute_reply":"2026-08-11T01:15:02.392065Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport glob\nimport random\nimport warnings\nwarnings.filterwarnings('ignore')\nos.environ['TOKENIZERS_PARALLELISM'] = 'false'  # quiets HF tokenizer fork warnings with num_workers>0\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport cv2\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import f1_score, roc_auc_score, cohen_kappa_score\n\nimport timm\nfrom transformers import AutoTokenizer, AutoModel, BertTokenizer, BertModel, T5Tokenizer, T5ForConditionalGeneration\nfrom torchvision.models.video import mc3_18, MC3_18_Weights\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\nimport logging\nlogging.getLogger('huggingface_hub').setLevel(logging.ERROR)\nlogging.getLogger('transformers').setLevel(logging.ERROR)\nfrom transformers import logging as hf_logging\nhf_logging.set_verbosity_error()\nos.environ['HF_HUB_DISABLE_PROGRESS_BARS'] = '1'\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'device: {DEVICE}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:15:02.394091Z","iopub.execute_input":"2026-08-11T01:15:02.394798Z","iopub.status.idle":"2026-08-11T01:15:31.052940Z","shell.execute_reply.started":"2026-08-11T01:15:02.394764Z","shell.execute_reply":"2026-08-11T01:15:31.052110Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Configuration\n\nPaths, label vocabulary, split ratios, and model hyperparameters are collected in one place so that Stages 1-3 reference the same settings.","metadata":{}},{"cell_type":"code","source":"class Config:\n    data_dir = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\n    train_csv = os.path.join(data_dir, 'train.csv')\n    coords_csv = os.path.join(data_dir, 'train_label_coordinates.csv')\n    series_csv = os.path.join(data_dir, 'train_series_descriptions.csv')\n    image_dir = os.path.join(data_dir, 'train_images')\n    conditions = ['spinal_canal_stenosis', 'left_neural_foraminal_narrowing', 'right_neural_foraminal_narrowing', 'left_subarticular_stenosis', 'right_subarticular_stenosis']\n    levels = ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']\n    severity_map = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\n    class_weights = {0: 1.0, 1: 2.0, 2: 4.0}\n    train_frac, val_frac, test_frac = 0.70, 0.15, 0.15\n    biobert_name = 'dmis-lab/biobert-base-cased-v1.1'\n    t5_name = 't5-small'\n    max_text_len = 256\n    decoder_max_len = 128\n    decoder_loss_weight = 1.0\n    image_size = 224\n    slices_per_volume = 12\n    batch_size = 8\n    lr = 2e-5\n    epochs = 3\n    run_training = True\n\ncfg = Config()\nsearch_roots = ['/kaggle/input'] + [os.path.join('/kaggle/input', d) for d in os.listdir('/kaggle/input')]\ncandidates = []\nfor root in search_roots:\n    if os.path.isdir(root):\n        candidates += [os.path.join(root, d) for d in os.listdir(root) if 'lumbar' in d.lower() or 'rsna' in d.lower()]\ncfg.data_dir = candidates[0] if candidates else cfg.data_dir\ncfg.train_csv = os.path.join(cfg.data_dir, 'train.csv')\ncfg.coords_csv = os.path.join(cfg.data_dir, 'train_label_coordinates.csv')\ncfg.series_csv = os.path.join(cfg.data_dir, 'train_series_descriptions.csv')\ncfg.image_dir = os.path.join(cfg.data_dir, 'train_images')\nprint('cfg.data_dir resolved to:', cfg.data_dir)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:15:31.054124Z","iopub.execute_input":"2026-08-11T01:15:31.054702Z","iopub.status.idle":"2026-08-11T01:15:31.065660Z","shell.execute_reply.started":"2026-08-11T01:15:31.054674Z","shell.execute_reply":"2026-08-11T01:15:31.064706Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Loading\n\nLoading the three metadata files supplied by the RSNA 2024 LumbarDISC challenge: study-level severity labels, per-level anatomical coordinates, and series-to-sequence mappings.","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(cfg.train_csv)\ncoords_df = pd.read_csv(cfg.coords_csv)\nseries_df = pd.read_csv(cfg.series_csv)\n\nprint(f'studies: {train_df.shape[0]}, label columns: {train_df.shape[1] - 1}')\nprint(f'coordinate rows: {coords_df.shape[0]}')\nprint(f'series rows: {series_df.shape[0]}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:15:31.067724Z","iopub.execute_input":"2026-08-11T01:15:31.068179Z","iopub.status.idle":"2026-08-11T01:15:31.292578Z","shell.execute_reply.started":"2026-08-11T01:15:31.068152Z","shell.execute_reply":"2026-08-11T01:15:31.291627Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_label_coordinates.csv ships with duplicate rows (same study/series/\n# instance/condition/level combination repeated)\nbefore = len(coords_df)\ncoords_df = coords_df.drop_duplicates(\n    subset=['study_id', 'series_id', 'instance_number', 'condition', 'level']\n).reset_index(drop=True)\nprint(f'dropped {before - len(coords_df)} duplicate coordinate rows')\n\nlabel_cols = [c for c in train_df.columns if c != 'study_id']\nnull_counts = train_df[label_cols].isna().sum().sum()\nnull_frac = null_counts / (train_df.shape[0] * len(label_cols))\nprint(f'missing labels: {null_counts} ({null_frac:.4%})')\n\n# a handful of missing cells does not disqualify a study for the levels that\n# are labelled, so nulls are masked out target-by-target later on rather\n# than dropping the whole study","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:15:31.293820Z","iopub.execute_input":"2026-08-11T01:15:31.294251Z","iopub.status.idle":"2026-08-11T01:15:31.342844Z","shell.execute_reply.started":"2026-08-11T01:15:31.294200Z","shell.execute_reply":"2026-08-11T01:15:31.342104Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def wide_to_long(df, conditions, levels, severity_map):\n    records = []\n    for _, row in df.iterrows():\n        for cond in conditions:\n            for lvl in levels:\n                col = f'{cond}_{lvl}'\n                if col not in df.columns:\n                    continue\n                raw = row[col]\n                if pd.isna(raw):\n                    continue\n                records.append({\n                    'study_id': row['study_id'],\n                    'condition': cond,\n                    'level': lvl,\n                    'severity': severity_map[raw],\n                })\n    return pd.DataFrame(records)\n\nlong_df = wide_to_long(train_df, cfg.conditions, cfg.levels, cfg.severity_map)\nlong_df['target'] = long_df['condition'] + '_' + long_df['level']\nprint(f'long-format rows: {len(long_df)} across {long_df.study_id.nunique()} studies')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:15:31.344111Z","iopub.execute_input":"2026-08-11T01:15:31.344521Z","iopub.status.idle":"2026-08-11T01:15:31.748316Z","shell.execute_reply.started":"2026-08-11T01:15:31.344482Z","shell.execute_reply":"2026-08-11T01:15:31.747197Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def study_severity_bucket(df, study_id):\n    # dominant severity per study, used only to stratify the split\n    sub = df.loc[df.study_id == study_id, 'severity']\n    return int(sub.mode().iloc[0]) if len(sub) else 0\n\nstudy_ids = train_df['study_id'].unique()\nbuckets = np.array([study_severity_bucket(long_df, sid) for sid in study_ids])\n\ntrain_ids, temp_ids, train_b, temp_b = train_test_split(\n    study_ids, buckets, test_size=(cfg.val_frac + cfg.test_frac),\n    stratify=buckets, random_state=SEED,\n)\nval_ids, test_ids = train_test_split(\n    temp_ids, test_size=cfg.test_frac / (cfg.val_frac + cfg.test_frac),\n    stratify=temp_b, random_state=SEED,\n)\n\nprint(f'train studies: {len(train_ids)}, val: {len(val_ids)}, test: {len(test_ids)}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:15:31.749683Z","iopub.execute_input":"2026-08-11T01:15:31.750618Z","iopub.status.idle":"2026-08-11T01:15:32.720488Z","shell.execute_reply.started":"2026-08-11T01:15:31.750587Z","shell.execute_reply":"2026-08-11T01:15:32.719548Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Pseudo-Report Generation\n\nLumbarDISC ships structured severity labels, not free-text radiology reports. Each study's 25 labels are converted into a deterministic, template-filled report using controlled clinical vocabulary, so the text fed to Stage 1-3 is transparent and reproducible.","metadata":{}},{"cell_type":"code","source":"SEVERITY_PHRASE = {0: 'no significant', 1: 'moderate', 2: 'severe'}\n\nCONDITION_PHRASE = {\n    'spinal_canal_stenosis': 'central canal stenosis',\n    'left_neural_foraminal_narrowing': 'left neural foraminal narrowing',\n    'right_neural_foraminal_narrowing': 'right neural foraminal narrowing',\n    'left_subarticular_stenosis': 'left subarticular stenosis',\n    'right_subarticular_stenosis': 'right subarticular stenosis',\n}\n\nLEVEL_PHRASE = {'l1_l2': 'L1-L2', 'l2_l3': 'L2-L3', 'l3_l4': 'L3-L4',\n                'l4_l5': 'L4-L5', 'l5_s1': 'L5-S1'}\n\ndef build_pseudo_report(study_id, long_df):\n    \"\"\"Turns one study's severity labels into a single templated report.\n\n    This text is deterministic and label-derived, not a dictated\n    radiologist report -- the dataset has no paired free text, so every\n    output here is described as a pseudo-report per the proposal's\n    text-validity safeguard.\n    \"\"\"\n    rows = long_df[long_df.study_id == study_id]\n    sentences = []\n    for lvl in cfg.levels:\n        lvl_rows = rows[rows.level == lvl]\n        if lvl_rows.empty:\n            continue\n        findings = [f\"{SEVERITY_PHRASE[r.severity]} {CONDITION_PHRASE[r.condition]}\"\n                    for _, r in lvl_rows.iterrows()]\n        sentences.append(f\"At {LEVEL_PHRASE[lvl]}, there is \" + ', '.join(findings) + '.')\n    return ' '.join(sentences) if sentences else 'No labelled findings at any level.'\n\npseudo_reports = {sid: build_pseudo_report(sid, long_df) for sid in study_ids}\nprint(pseudo_reports[study_ids[0]])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:15:32.721535Z","iopub.execute_input":"2026-08-11T01:15:32.721928Z","iopub.status.idle":"2026-08-11T01:15:40.066429Z","shell.execute_reply.started":"2026-08-11T01:15:32.721885Z","shell.execute_reply":"2026-08-11T01:15:40.065320Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Weighted Loss\n\nSeverity labels are heavily imbalanced (~78% Normal/Mild, ~16% Moderate, ~6% Severe). Rather than resampling, the imbalance is handled at the objective level with the RSNA-style weighted log loss, applied identically in all three stages.","metadata":{}},{"cell_type":"code","source":"def rsna_weighted_log_loss(y_true, y_prob, class_weights=cfg.class_weights, eps=1e-7):\n    \"\"\"Same weighted log loss used across all three stages, so the ablation\n    comparison later is measuring the architecture, not the objective.\"\"\"\n    y_prob = np.clip(y_prob, eps, 1 - eps)\n    weights = np.array([class_weights[y] for y in y_true])\n    picked = y_prob[np.arange(len(y_true)), y_true]\n    return float(np.sum(-weights * np.log(picked)) / np.sum(weights))\n\nclass WeightedCELoss(nn.Module):\n    def __init__(self, class_weights=cfg.class_weights):\n        super().__init__()\n        w = torch.tensor([class_weights[i] for i in sorted(class_weights)], dtype=torch.float32)\n        self.register_buffer('w', w)\n\n    def forward(self, logits, targets):\n        return F.cross_entropy(logits, targets, weight=self.w.to(logits.device))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:15:40.068558Z","iopub.execute_input":"2026-08-11T01:15:40.068960Z","iopub.status.idle":"2026-08-11T01:15:40.076310Z","shell.execute_reply.started":"2026-08-11T01:15:40.068925Z","shell.execute_reply":"2026-08-11T01:15:40.075429Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Stage 1: Text-Only Baseline\n\nEstablishes the lower bound attributable to label-derived text alone, with no imaging input. BioBERT encodes the pseudo-reports; a linear head predicts severity per target and a T5-small decoder regenerates narrative text from the same encoder states.","metadata":{}},{"cell_type":"code","source":"tokenizer_bio = BertTokenizer.from_pretrained(cfg.biobert_name)\ntokenizer_t5 = T5Tokenizer.from_pretrained(cfg.t5_name)\n\ndef encode_t5_labels(text, max_len=cfg.decoder_max_len):\n    \"\"\"T5 target ids for the decoder, with pad positions set to -100 so\n    they're ignored by the cross-entropy loss inside T5ForConditionalGeneration.\"\"\"\n    enc = tokenizer_t5(text, truncation=True, padding='max_length',\n                        max_length=max_len, return_tensors='pt')\n    labels = enc['input_ids'].squeeze(0)\n    labels[labels == tokenizer_t5.pad_token_id] = -100\n    return labels\n\nclass TextOnlyDataset(Dataset):\n    \"\"\"Stage 1 input: pseudo-report text only, no imaging.\"\"\"\n    def __init__(self, study_ids, long_df, reports):\n        self.study_ids = list(study_ids)\n        self.long_df = long_df\n        self.reports = reports\n        self.targets = sorted(long_df['target'].unique())\n\n    def __len__(self):\n        return len(self.study_ids)\n\n    def __getitem__(self, idx):\n        sid = self.study_ids[idx]\n        text = self.reports[sid]\n        enc = tokenizer_bio(text, truncation=True, padding='max_length',\n                             max_length=cfg.max_text_len, return_tensors='pt')\n\n        rows = self.long_df[self.long_df.study_id == sid].set_index('target')\n        labels = torch.full((len(self.targets),), -1, dtype=torch.long)\n        for i, t in enumerate(self.targets):\n            if t in rows.index:\n                labels[i] = int(rows.loc[t, 'severity'])\n\n        return {\n            'input_ids': enc['input_ids'].squeeze(0),\n            'attention_mask': enc['attention_mask'].squeeze(0),\n            'labels': labels,\n            't5_labels': encode_t5_labels(text),\n        }\n\ntrain_ds_s1 = TextOnlyDataset(train_ids, long_df, pseudo_reports)\nval_ds_s1 = TextOnlyDataset(val_ids, long_df, pseudo_reports)\ntest_ds_s1 = TextOnlyDataset(test_ids, long_df, pseudo_reports)\nn_targets = len(train_ds_s1.targets)\nprint(f'{n_targets} condition-level targets per study')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:15:40.077442Z","iopub.execute_input":"2026-08-11T01:15:40.077854Z","iopub.status.idle":"2026-08-11T01:15:41.708954Z","shell.execute_reply.started":"2026-08-11T01:15:40.077828Z","shell.execute_reply":"2026-08-11T01:15:41.707997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TextBaselineModel(nn.Module):\n    \"\"\"BioBERT encoder -> linear severity head, plus a T5-small decoder\n    conditioned on the same encoder states for pseudo-report generation.\"\"\"\n    def __init__(self, n_targets, n_classes=3):\n        super().__init__()\n        self.encoder = BertModel.from_pretrained(cfg.biobert_name)\n        hidden = self.encoder.config.hidden_size\n        self.classifier = nn.Linear(hidden, n_targets * n_classes)\n        self.n_targets = n_targets\n        self.n_classes = n_classes\n\n        self.decoder = T5ForConditionalGeneration.from_pretrained(cfg.t5_name)\n        self.proj_to_t5 = nn.Linear(hidden, self.decoder.config.d_model)\n\n    def forward(self, input_ids, attention_mask, decoder_labels=None):\n        out = self.encoder(input_ids=input_ids, attention_mask=attention_mask)\n        pooled = out.last_hidden_state[:, 0, :]  # [CLS]\n\n        logits = self.classifier(pooled).view(-1, self.n_targets, self.n_classes)\n\n        decoder_out = None\n        if decoder_labels is not None:\n            encoder_hidden = self.proj_to_t5(out.last_hidden_state)\n            decoder_out = self.decoder(\n                encoder_outputs=(encoder_hidden,),\n                attention_mask=attention_mask,\n                labels=decoder_labels,\n            )\n        return logits, decoder_out\n\nmodel_s1 = TextBaselineModel(n_targets).to(DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:15:41.710316Z","iopub.execute_input":"2026-08-11T01:15:41.710778Z","iopub.status.idle":"2026-08-11T01:15:49.090205Z","shell.execute_reply.started":"2026-08-11T01:15:41.710753Z","shell.execute_reply":"2026-08-11T01:15:49.089181Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_stage(model, train_loader, val_loader, epochs=cfg.epochs, lr=cfg.lr, tag='stage1'):\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)\n    cls_loss_fn = WeightedCELoss().to(DEVICE)\n\n    best_val_loss = float('inf')\n    for epoch in range(epochs):\n        model.train()\n        running_loss = 0.0\n        for batch in train_loader:\n            input_ids = batch['input_ids'].to(DEVICE)\n            attn = batch['attention_mask'].to(DEVICE)\n            labels = batch['labels'].to(DEVICE)\n            t5_labels = batch['t5_labels'].to(DEVICE)\n\n            optimizer.zero_grad()\n            logits, decoder_out = model(input_ids, attn, decoder_labels=t5_labels)\n\n            cls_loss, valid_targets = 0.0, 0\n            for t in range(logits.shape[1]):\n                mask = labels[:, t] != -1\n                if mask.sum() == 0:\n                    continue\n                cls_loss = cls_loss + cls_loss_fn(logits[mask, t, :], labels[mask, t])\n                valid_targets += 1\n            cls_loss = cls_loss / max(valid_targets, 1)\n\n            loss = cls_loss + cfg.decoder_loss_weight * decoder_out.loss\n\n            loss.backward()\n            optimizer.step()\n            running_loss += loss.item()\n\n        scheduler.step()\n        val_loss = evaluate_stage_loss(model, val_loader, cls_loss_fn)\n        print(f'[{tag}] epoch {epoch+1}/{epochs}  train_loss={running_loss/len(train_loader):.4f}  val_cls_loss={val_loss:.4f}')\n\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            torch.save(model.state_dict(), f'/kaggle/working/{tag}_best.pt')\n\n    return model\n\ndef evaluate_stage_loss(model, loader, cls_loss_fn):\n    model.eval()\n    total, n = 0.0, 0\n    with torch.no_grad():\n        for batch in loader:\n            input_ids = batch['input_ids'].to(DEVICE)\n            attn = batch['attention_mask'].to(DEVICE)\n            labels = batch['labels'].to(DEVICE)\n            logits, _ = model(input_ids, attn)\n            for t in range(logits.shape[1]):\n                mask = labels[:, t] != -1\n                if mask.sum() == 0:\n                    continue\n                total += cls_loss_fn(logits[mask, t, :], labels[mask, t]).item()\n                n += 1\n    return total / max(n, 1)\n\ntrain_loader_s1 = DataLoader(train_ds_s1, batch_size=cfg.batch_size, shuffle=True)\nval_loader_s1 = DataLoader(val_ds_s1, batch_size=cfg.batch_size)\ntest_loader_s1 = DataLoader(test_ds_s1, batch_size=cfg.batch_size)\n\nif cfg.run_training:\n    model_s1 = train_stage(model_s1, train_loader_s1, val_loader_s1, tag='stage1')\n    model_s1 = model_s1.cpu()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:15:49.091513Z","iopub.execute_input":"2026-08-11T01:15:49.091866Z","iopub.status.idle":"2026-08-11T01:21:22.615433Z","shell.execute_reply.started":"2026-08-11T01:15:49.091839Z","shell.execute_reply":"2026-08-11T01:21:22.614413Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Stage 2: 2D Multi-View Fusion\n\nIntroduces imaging evidence on top of the Stage 1 text baseline. One sagittal T2 and one axial T2 slice are extracted per labelled level, encoded with a pretrained EfficientNet-B3 backbone, and fused with the BioBERT text tokens through a single-layer cross-attention block.","metadata":{}},{"cell_type":"code","source":"def get_series_for_description(study_id, description, series_df):\n    \"\"\"Maps a study to the series_id matching a given sequence description\n    (e.g. 'Sagittal T2/STIR', 'Axial T2') using the series metadata file.\"\"\"\n    rows = series_df[(series_df.study_id == study_id) &\n                      (series_df.series_description == description)]\n    return rows.iloc[0].series_id if len(rows) else None\n\ndef load_dicom_slice(study_id, series_id, instance_number, image_dir=cfg.image_dir):\n    path = os.path.join(image_dir, str(study_id), str(series_id), f'{instance_number}.dcm')\n    dcm = pydicom.dcmread(path)\n    arr = dcm.pixel_array.astype(np.float32)\n    arr = (arr - arr.mean()) / (arr.std() + 1e-6)  # z-score normalisation\n    return arr\n\ndef resize_to_224(arr, size=cfg.image_size):\n    return cv2.resize(arr, (size, size), interpolation=cv2.INTER_LINEAR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:21:22.618559Z","iopub.execute_input":"2026-08-11T01:21:22.619447Z","iopub.status.idle":"2026-08-11T01:21:22.626373Z","shell.execute_reply.started":"2026-08-11T01:21:22.619413Z","shell.execute_reply":"2026-08-11T01:21:22.625432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MultiViewDataset(Dataset):\n    \"\"\"Stage 2 input: one sagittal T2 slice + one axial T2 slice per study,\n    picked at the coordinate closest to each labelled level, plus the same\n    pseudo-report text used in Stage 1.\"\"\"\n    def __init__(self, study_ids, long_df, reports, coords_df, series_df):\n        self.study_ids = list(study_ids)\n        self.long_df = long_df\n        self.reports = reports\n        self.coords_df = coords_df\n        self.series_df = series_df\n        self.targets = sorted(long_df['target'].unique())\n\n    def __len__(self):\n        return len(self.study_ids)\n\n    def _pick_slice(self, sid, description):\n        series_id = get_series_for_description(sid, description, self.series_df)\n        if series_id is None:\n            return np.zeros((cfg.image_size, cfg.image_size), dtype=np.float32)\n\n        rows = self.coords_df[(self.coords_df.study_id == sid) &\n                               (self.coords_df.series_id == series_id)]\n        if rows.empty:\n            return np.zeros((cfg.image_size, cfg.image_size), dtype=np.float32)\n\n        instance_number = int(rows.iloc[len(rows) // 2].instance_number)\n        try:\n            arr = load_dicom_slice(sid, series_id, instance_number)\n            return resize_to_224(arr)\n        except Exception:\n            return np.zeros((cfg.image_size, cfg.image_size), dtype=np.float32)\n\n    def __getitem__(self, idx):\n        sid = self.study_ids[idx]\n        sag = self._pick_slice(sid, 'Sagittal T2/STIR')\n        ax = self._pick_slice(sid, 'Axial T2')\n\n        text = self.reports[sid]\n        enc = tokenizer_bio(text, truncation=True, padding='max_length',\n                             max_length=cfg.max_text_len, return_tensors='pt')\n\n        rows = self.long_df[self.long_df.study_id == sid].set_index('target')\n        labels = torch.full((len(self.targets),), -1, dtype=torch.long)\n        for i, t in enumerate(self.targets):\n            if t in rows.index:\n                labels[i] = int(rows.loc[t, 'severity'])\n\n        return {\n            'sag_slice': torch.from_numpy(sag).unsqueeze(0).repeat(3, 1, 1).float(),\n            'ax_slice': torch.from_numpy(ax).unsqueeze(0).repeat(3, 1, 1).float(),\n            'input_ids': enc['input_ids'].squeeze(0),\n            'attention_mask': enc['attention_mask'].squeeze(0),\n            'labels': labels,\n            't5_labels': encode_t5_labels(text),\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:21:22.627563Z","iopub.execute_input":"2026-08-11T01:21:22.628278Z","iopub.status.idle":"2026-08-11T01:21:22.643782Z","shell.execute_reply.started":"2026-08-11T01:21:22.628242Z","shell.execute_reply":"2026-08-11T01:21:22.642936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CrossAttentionFusion(nn.Module):\n    \"\"\"Single-layer cross-attention: text tokens attend over image/volume\n    patch tokens, following the fusion mechanism specified for Stage 2.\"\"\"\n    def __init__(self, dim, n_heads=8):\n        super().__init__()\n        self.attn = nn.MultiheadAttention(dim, n_heads, batch_first=True)\n        self.norm = nn.LayerNorm(dim)\n\n    def forward(self, text_tokens, image_tokens):\n        fused, _ = self.attn(query=text_tokens, key=image_tokens, value=image_tokens)\n        return self.norm(text_tokens + fused)\n\nclass MultiViewFusionModel(nn.Module):\n    def __init__(self, n_targets, n_classes=3):\n        super().__init__()\n        self.text_encoder = BertModel.from_pretrained(cfg.biobert_name)\n        hidden = self.text_encoder.config.hidden_size\n\n        self.image_encoder = timm.create_model('efficientnet_b3', pretrained=True,\n                                                 num_classes=0, global_pool='')\n        img_feat_dim = self.image_encoder.num_features\n        self.img_proj = nn.Linear(img_feat_dim, hidden)\n\n        self.fusion = CrossAttentionFusion(hidden)\n        self.classifier = nn.Linear(hidden, n_targets * n_classes)\n        self.n_targets, self.n_classes = n_targets, n_classes\n\n        self.decoder = T5ForConditionalGeneration.from_pretrained(cfg.t5_name)\n        self.proj_to_t5 = nn.Linear(hidden, self.decoder.config.d_model)\n\n    def _image_tokens(self, sag, ax):\n        feat_sag = self.image_encoder(sag)   # (B, C, H', W')\n        feat_ax = self.image_encoder(ax)\n        feat_sag = feat_sag.flatten(2).transpose(1, 2)  # (B, N, C)\n        feat_ax = feat_ax.flatten(2).transpose(1, 2)\n        tokens = torch.cat([feat_sag, feat_ax], dim=1)\n        return self.img_proj(tokens)\n\n    def forward(self, input_ids, attention_mask, sag_slice, ax_slice, decoder_labels=None):\n        text_out = self.text_encoder(input_ids=input_ids, attention_mask=attention_mask)\n        text_tokens = text_out.last_hidden_state\n\n        image_tokens = self._image_tokens(sag_slice, ax_slice)\n        fused = self.fusion(text_tokens, image_tokens)\n        pooled = fused[:, 0, :]\n\n        logits = self.classifier(pooled).view(-1, self.n_targets, self.n_classes)\n\n        decoder_out = None\n        if decoder_labels is not None:\n            encoder_hidden = self.proj_to_t5(fused)\n            decoder_out = self.decoder(\n                encoder_outputs=(encoder_hidden,),\n                attention_mask=attention_mask,\n                labels=decoder_labels,\n            )\n        return logits, decoder_out\n\nmodel_s2 = MultiViewFusionModel(n_targets).to(DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:21:22.644952Z","iopub.execute_input":"2026-08-11T01:21:22.645543Z","iopub.status.idle":"2026-08-11T01:21:26.020546Z","shell.execute_reply.started":"2026-08-11T01:21:22.645517Z","shell.execute_reply":"2026-08-11T01:21:26.019441Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_stage2(model, train_loader, val_loader, epochs=cfg.epochs, lr=cfg.lr):\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)\n    cls_loss_fn = WeightedCELoss().to(DEVICE)\n\n    best_val_loss = float('inf')\n    for epoch in range(epochs):\n        model.train()\n        running_loss = 0.0\n        for batch in train_loader:\n            input_ids = batch['input_ids'].to(DEVICE)\n            attn = batch['attention_mask'].to(DEVICE)\n            sag = batch['sag_slice'].to(DEVICE)\n            ax = batch['ax_slice'].to(DEVICE)\n            labels = batch['labels'].to(DEVICE)\n            t5_labels = batch['t5_labels'].to(DEVICE)\n\n            optimizer.zero_grad()\n            logits, decoder_out = model(input_ids, attn, sag, ax, decoder_labels=t5_labels)\n\n            cls_loss, valid_targets = 0.0, 0\n            for t in range(logits.shape[1]):\n                mask = labels[:, t] != -1\n                if mask.sum() == 0:\n                    continue\n                cls_loss = cls_loss + cls_loss_fn(logits[mask, t, :], labels[mask, t])\n                valid_targets += 1\n            cls_loss = cls_loss / max(valid_targets, 1)\n\n            loss = cls_loss + cfg.decoder_loss_weight * decoder_out.loss\n\n            loss.backward()\n            optimizer.step()\n            running_loss += loss.item()\n\n        scheduler.step()\n        val_loss = evaluate_stage2_loss(model, val_loader, cls_loss_fn)\n        print(f'[stage2] epoch {epoch+1}/{epochs}  train_loss={running_loss/len(train_loader):.4f}  val_cls_loss={val_loss:.4f}')\n\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            torch.save(model.state_dict(), '/kaggle/working/stage2_best.pt')\n\n    return model\n\ndef evaluate_stage2_loss(model, loader, cls_loss_fn):\n    model.eval()\n    total, n = 0.0, 0\n    with torch.no_grad():\n        for batch in loader:\n            input_ids = batch['input_ids'].to(DEVICE)\n            attn = batch['attention_mask'].to(DEVICE)\n            sag = batch['sag_slice'].to(DEVICE)\n            ax = batch['ax_slice'].to(DEVICE)\n            labels = batch['labels'].to(DEVICE)\n            logits, _ = model(input_ids, attn, sag, ax)\n            for t in range(logits.shape[1]):\n                mask = labels[:, t] != -1\n                if mask.sum() == 0:\n                    continue\n                total += cls_loss_fn(logits[mask, t, :], labels[mask, t]).item()\n                n += 1\n    return total / max(n, 1)\n\ntrain_loader_s2 = DataLoader(MultiViewDataset(train_ids, long_df, pseudo_reports, coords_df, series_df),\n                              batch_size=cfg.batch_size, shuffle=True, num_workers=2, pin_memory=True)\nval_loader_s2 = DataLoader(MultiViewDataset(val_ids, long_df, pseudo_reports, coords_df, series_df),\n                            batch_size=cfg.batch_size, num_workers=2, pin_memory=True)\ntest_loader_s2 = DataLoader(MultiViewDataset(test_ids, long_df, pseudo_reports, coords_df, series_df),\n                             batch_size=cfg.batch_size, num_workers=2, pin_memory=True)\n\nif cfg.run_training:\n    model_s2 = train_stage2(model_s2, train_loader_s2, val_loader_s2)\n    model_s2 = model_s2.cpu()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:21:26.021894Z","iopub.execute_input":"2026-08-11T01:21:26.022232Z","iopub.status.idle":"2026-08-11T01:29:09.702512Z","shell.execute_reply.started":"2026-08-11T01:21:26.022204Z","shell.execute_reply":"2026-08-11T01:29:09.701321Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Stage 3: 3D Volumetric Fusion\n\nReplaces the 2D slice extractor with contiguous axial stacks (8-16 slices) around each labelled level, encoded with MC3-18 (~11.7M parameters, pretrained on Kinetics-400 -- the closest off-the-shelf equivalent to the I3D-Lite architecture specified in the proposal) and fused with text through the same cross-attention block used in Stage 2. This isolates the marginal contribution of local volumetric context over single-slice 2D fusion.","metadata":{}},{"cell_type":"code","source":"def load_volume(study_id, series_id, center_instance, n_slices=cfg.slices_per_volume,\n                 image_dir=cfg.image_dir):\n    \"\"\"Stacks n_slices contiguous axial slices centred on the labelled level\n    into a single (D, H, W) volume, as specified for Stage 3.\"\"\"\n    half = n_slices // 2\n    slice_paths = sorted(\n        glob.glob(os.path.join(image_dir, str(study_id), str(series_id), '*.dcm')),\n        key=lambda p: int(os.path.basename(p).split('.')[0]),\n    )\n    instance_numbers = [int(os.path.basename(p).split('.')[0]) for p in slice_paths]\n\n    if center_instance in instance_numbers:\n        center_idx = instance_numbers.index(center_instance)\n    else:\n        center_idx = len(instance_numbers) // 2\n\n    lo = max(0, center_idx - half)\n    hi = min(len(slice_paths), lo + n_slices)\n    lo = max(0, hi - n_slices)\n\n    slices = []\n    for p in slice_paths[lo:hi]:\n        dcm = pydicom.dcmread(p)\n        arr = dcm.pixel_array.astype(np.float32)\n        arr = (arr - arr.mean()) / (arr.std() + 1e-6)\n        slices.append(resize_to_224(arr))\n\n    while len(slices) < n_slices:  # pad short series at the volume edges\n        slices.append(np.zeros((cfg.image_size, cfg.image_size), dtype=np.float32))\n\n    return np.stack(slices, axis=0)  # (D, H, W)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:29:09.704185Z","iopub.execute_input":"2026-08-11T01:29:09.704612Z","iopub.status.idle":"2026-08-11T01:29:09.713515Z","shell.execute_reply.started":"2026-08-11T01:29:09.704579Z","shell.execute_reply":"2026-08-11T01:29:09.712821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VolumetricDataset(Dataset):\n    \"\"\"Stage 3 input: a stacked axial volume around each labelled level,\n    plus the same pseudo-report text used in Stages 1 and 2.\"\"\"\n    def __init__(self, study_ids, long_df, reports, coords_df, series_df):\n        self.study_ids = list(study_ids)\n        self.long_df = long_df\n        self.reports = reports\n        self.coords_df = coords_df\n        self.series_df = series_df\n        self.targets = sorted(long_df['target'].unique())\n\n    def __len__(self):\n        return len(self.study_ids)\n\n    def __getitem__(self, idx):\n        sid = self.study_ids[idx]\n        series_id = get_series_for_description(sid, 'Axial T2', self.series_df)\n\n        rows = pd.DataFrame()\n        if series_id is not None:\n            rows = self.coords_df[(self.coords_df.study_id == sid) &\n                                   (self.coords_df.series_id == series_id)]\n\n        if series_id is not None and not rows.empty:\n            center_instance = int(rows.iloc[len(rows) // 2].instance_number)\n            try:\n                volume = load_volume(sid, series_id, center_instance)\n            except Exception:\n                volume = np.zeros((cfg.slices_per_volume, cfg.image_size, cfg.image_size), dtype=np.float32)\n        else:\n            volume = np.zeros((cfg.slices_per_volume, cfg.image_size, cfg.image_size), dtype=np.float32)\n\n        text = self.reports[sid]\n        enc = tokenizer_bio(text, truncation=True, padding='max_length',\n                             max_length=cfg.max_text_len, return_tensors='pt')\n\n        rows_lab = self.long_df[self.long_df.study_id == sid].set_index('target')\n        labels = torch.full((len(self.targets),), -1, dtype=torch.long)\n        for i, t in enumerate(self.targets):\n            if t in rows_lab.index:\n                labels[i] = int(rows_lab.loc[t, 'severity'])\n\n        return {\n            'volume': torch.from_numpy(volume).unsqueeze(0).repeat(3, 1, 1, 1).float(),  # (3, D, H, W)\n            'input_ids': enc['input_ids'].squeeze(0),\n            'attention_mask': enc['attention_mask'].squeeze(0),\n            'labels': labels,\n            't5_labels': encode_t5_labels(text),\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:29:09.714470Z","iopub.execute_input":"2026-08-11T01:29:09.714752Z","iopub.status.idle":"2026-08-11T01:29:09.746094Z","shell.execute_reply.started":"2026-08-11T01:29:09.714705Z","shell.execute_reply":"2026-08-11T01:29:09.745120Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class I3DLite(nn.Module):\n    \"\"\"Volumetric encoder for Stage 3. Uses MC3-18 (mixed 3D/2D convolutions,\n    ~11.7M parameters -- matching the ~11M budget specified for I3D-Lite in\n    the proposal) pretrained on Kinetics-400, which is the closest\n    off-the-shelf architecture to a lightweight inflated-3D CNN with a\n    usable pretrained checkpoint available given internet access. The stem\n    and residual stages transfer directly; only the pooling/classification\n    head is discarded so cross-attention can attend over the\n    spatio-temporal feature map instead of a single pooled vector.\"\"\"\n    def __init__(self):\n        super().__init__()\n        backbone = mc3_18(weights=MC3_18_Weights.KINETICS400_V1)\n        self.features = nn.Sequential(*list(backbone.children())[:-2])  # drop avgpool + fc\n        self.out_channels = 512\n\n    def forward(self, x):\n        return self.features(x)  # (B, 512, D', H', W')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:29:09.747292Z","iopub.execute_input":"2026-08-11T01:29:09.748069Z","iopub.status.idle":"2026-08-11T01:29:09.770127Z","shell.execute_reply.started":"2026-08-11T01:29:09.748041Z","shell.execute_reply":"2026-08-11T01:29:09.769178Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VolumetricFusionModel(nn.Module):\n    def __init__(self, n_targets, n_classes=3):\n        super().__init__()\n        self.text_encoder = BertModel.from_pretrained(cfg.biobert_name)\n        hidden = self.text_encoder.config.hidden_size\n\n        self.volume_encoder = I3DLite()\n        self.vol_proj = nn.Linear(self.volume_encoder.out_channels, hidden)\n\n        self.fusion = CrossAttentionFusion(hidden)\n        self.classifier = nn.Linear(hidden, n_targets * n_classes)\n        self.n_targets, self.n_classes = n_targets, n_classes\n\n        self.decoder = T5ForConditionalGeneration.from_pretrained(cfg.t5_name)\n        self.proj_to_t5 = nn.Linear(hidden, self.decoder.config.d_model)\n\n    def _volume_tokens(self, volume):\n        feat = self.volume_encoder(volume)          # (B, C, D', H', W')\n        b, c, d, h, w = feat.shape\n        tokens = feat.permute(0, 2, 3, 4, 1).reshape(b, d * h * w, c)\n        return self.vol_proj(tokens)\n\n    def forward(self, input_ids, attention_mask, volume, decoder_labels=None):\n        text_out = self.text_encoder(input_ids=input_ids, attention_mask=attention_mask)\n        text_tokens = text_out.last_hidden_state\n\n        volume_tokens = self._volume_tokens(volume)\n        fused = self.fusion(text_tokens, volume_tokens)\n        pooled = fused[:, 0, :]\n\n        logits = self.classifier(pooled).view(-1, self.n_targets, self.n_classes)\n\n        decoder_out = None\n        if decoder_labels is not None:\n            encoder_hidden = self.proj_to_t5(fused)\n            decoder_out = self.decoder(\n                encoder_outputs=(encoder_hidden,),\n                attention_mask=attention_mask,\n                labels=decoder_labels,\n            )\n        return logits, decoder_out\n\nmodel_s3 = VolumetricFusionModel(n_targets).to(DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:29:09.771276Z","iopub.execute_input":"2026-08-11T01:29:09.771664Z","iopub.status.idle":"2026-08-11T01:29:13.229036Z","shell.execute_reply.started":"2026-08-11T01:29:09.771638Z","shell.execute_reply":"2026-08-11T01:29:13.227890Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_stage3(model, train_loader, val_loader, epochs=cfg.epochs, lr=cfg.lr):\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)\n    cls_loss_fn = WeightedCELoss().to(DEVICE)\n\n    best_val_loss = float('inf')\n    for epoch in range(epochs):\n        model.train()\n        running_loss = 0.0\n        for batch in train_loader:\n            input_ids = batch['input_ids'].to(DEVICE)\n            attn = batch['attention_mask'].to(DEVICE)\n            volume = batch['volume'].to(DEVICE)\n            labels = batch['labels'].to(DEVICE)\n            t5_labels = batch['t5_labels'].to(DEVICE)\n\n            optimizer.zero_grad()\n            logits, decoder_out = model(input_ids, attn, volume, decoder_labels=t5_labels)\n\n            cls_loss, valid_targets = 0.0, 0\n            for t in range(logits.shape[1]):\n                mask = labels[:, t] != -1\n                if mask.sum() == 0:\n                    continue\n                cls_loss = cls_loss + cls_loss_fn(logits[mask, t, :], labels[mask, t])\n                valid_targets += 1\n            cls_loss = cls_loss / max(valid_targets, 1)\n\n            loss = cls_loss + cfg.decoder_loss_weight * decoder_out.loss\n\n            loss.backward()\n            optimizer.step()\n            running_loss += loss.item()\n\n        scheduler.step()\n        val_loss = evaluate_stage3_loss(model, val_loader, cls_loss_fn)\n        print(f'[stage3] epoch {epoch+1}/{epochs}  train_loss={running_loss/len(train_loader):.4f}  val_cls_loss={val_loss:.4f}')\n\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            torch.save(model.state_dict(), '/kaggle/working/stage3_best.pt')\n\n    return model\n\ndef evaluate_stage3_loss(model, loader, cls_loss_fn):\n    model.eval()\n    total, n = 0.0, 0\n    with torch.no_grad():\n        for batch in loader:\n            input_ids = batch['input_ids'].to(DEVICE)\n            attn = batch['attention_mask'].to(DEVICE)\n            volume = batch['volume'].to(DEVICE)\n            labels = batch['labels'].to(DEVICE)\n            logits, _ = model(input_ids, attn, volume)\n            for t in range(logits.shape[1]):\n                mask = labels[:, t] != -1\n                if mask.sum() == 0:\n                    continue\n                total += cls_loss_fn(logits[mask, t, :], labels[mask, t]).item()\n                n += 1\n    return total / max(n, 1)\n\ntrain_loader_s3 = DataLoader(VolumetricDataset(train_ids, long_df, pseudo_reports, coords_df, series_df),\n                              batch_size=max(cfg.batch_size // 2, 1), shuffle=True, num_workers=2, pin_memory=True)\nval_loader_s3 = DataLoader(VolumetricDataset(val_ids, long_df, pseudo_reports, coords_df, series_df),\n                            batch_size=max(cfg.batch_size // 2, 1), num_workers=2, pin_memory=True)\ntest_loader_s3 = DataLoader(VolumetricDataset(test_ids, long_df, pseudo_reports, coords_df, series_df),\n                             batch_size=max(cfg.batch_size // 2, 1), num_workers=2, pin_memory=True)\n\nif cfg.run_training:\n    model_s3 = train_stage3(model_s3, train_loader_s3, val_loader_s3)\n    model_s3 = model_s3.cpu()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T01:29:13.230327Z","iopub.execute_input":"2026-08-11T01:29:13.230794Z","iopub.status.idle":"2026-08-11T02:01:49.213453Z","shell.execute_reply.started":"2026-08-11T01:29:13.230767Z","shell.execute_reply":"2026-08-11T02:01:49.212306Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Evaluation\n\nShared scoring utilities used to build the ablation table across all three stages: weighted log loss (primary RSNA metric), macro F1, per-condition AUC, and weighted Cohen's kappa against the radiologist-assigned ground truth.","metadata":{}},{"cell_type":"code","source":"import time as _time\n\n@torch.no_grad()\ndef collect_predictions(model, loader, forward_fn, tag=''):\n    \"\"\"forward_fn(model, batch) -> logits. Kept generic since each stage\n    passes a different set of inputs into its model's forward method.\"\"\"\n    model.eval()\n    all_probs, all_labels = [], []\n    n_batches = len(loader)\n    t0 = _time.time()\n    for i, batch in enumerate(loader):\n        logits = forward_fn(model, batch)\n        probs = F.softmax(logits, dim=-1).cpu().numpy()\n        labels = batch['labels'].numpy()\n        all_probs.append(probs)\n        all_labels.append(labels)\n        if (i + 1) % 5 == 0 or (i + 1) == n_batches:\n            elapsed = _time.time() - t0\n            print(f'[{tag}] eval batch {i + 1}/{n_batches}  ({elapsed:.1f}s elapsed)', flush=True)\n    return np.concatenate(all_probs), np.concatenate(all_labels)\n\ndef score_predictions(probs, labels, tag=''):\n    \"\"\"Weighted log loss, macro F1, per-condition AUC and weighted kappa,\n    computed target-by-target and averaged, per the evaluation plan.\"\"\"\n    n_targets = probs.shape[1]\n    losses, f1s, aucs, kappas = [], [], [], []\n    t0 = _time.time()\n    print(f'[{tag}] scoring {n_targets} targets, probs shape={probs.shape}, labels shape={labels.shape}', flush=True)\n\n    for t in range(n_targets):\n        mask = labels[:, t] != -1\n        if mask.sum() == 0:\n            continue\n        y_true = labels[mask, t]\n        y_prob = probs[mask, t, :]\n        y_pred = y_prob.argmax(axis=1)\n\n        losses.append(rsna_weighted_log_loss(y_true, y_prob))\n        f1s.append(f1_score(y_true, y_pred, average='macro', zero_division=0))\n        kappas.append(cohen_kappa_score(y_true, y_pred, weights='linear'))\n        try:\n            aucs.append(roc_auc_score(y_true, y_prob, multi_class='ovr'))\n        except ValueError:\n            pass  # a target missing one of the three classes in this split\n        print(f'[{tag}] target {t + 1}/{n_targets} done  ({_time.time() - t0:.1f}s elapsed)', flush=True)\n\n    print(f'[{tag}] scoring complete, aggregating means', flush=True)\n    result = {\n        'weighted_log_loss': float(np.mean(losses)),\n        'macro_f1': float(np.mean(f1s)),\n        'auc': float(np.mean(aucs)) if aucs else float('nan'),\n        'weighted_kappa': float(np.mean(kappas)),\n    }\n    print(f'[{tag}] scoring done: {result}', flush=True)\n    return result","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T02:01:49.215064Z","iopub.execute_input":"2026-08-11T02:01:49.215463Z","iopub.status.idle":"2026-08-11T02:01:49.228051Z","shell.execute_reply.started":"2026-08-11T02:01:49.215419Z","shell.execute_reply":"2026-08-11T02:01:49.227290Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"forward_s1 = lambda model, batch: model(batch['input_ids'].to(DEVICE),\n                                        batch['attention_mask'].to(DEVICE))[0]\nforward_s2 = lambda model, batch: model(batch['input_ids'].to(DEVICE),\n                                        batch['attention_mask'].to(DEVICE),\n                                        batch['sag_slice'].to(DEVICE),\n                                        batch['ax_slice'].to(DEVICE))[0]\nforward_s3 = lambda model, batch: model(batch['input_ids'].to(DEVICE),\n                                        batch['attention_mask'].to(DEVICE),\n                                        batch['volume'].to(DEVICE))[0]\n\nif cfg.run_training:\n    results = {}\n    for tag, model, loader, fwd in [\n        ('Stage 1: Text Baseline', model_s1, test_loader_s1, forward_s1),\n        ('Stage 2: 2D Multi-View Fusion', model_s2, test_loader_s2, forward_s2),\n        ('Stage 3: 3D Volumetric Fusion', model_s3, test_loader_s3, forward_s3),\n    ]:\n        model.to(DEVICE)\n        probs, labels = collect_predictions(model, loader, fwd, tag=tag)\n        results[tag] = score_predictions(probs, labels, tag=tag)\n        print(f'[{tag}] moving model off GPU', flush=True)\n        model.cpu()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n        print(f'[{tag}] done, cache cleared', flush=True)\n\n    print('all stages scored, building ablation table', flush=True)\n    ablation_table = pd.DataFrame(results).T\n    ablation_table","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T02:01:49.229149Z","iopub.execute_input":"2026-08-11T02:01:49.229632Z","iopub.status.idle":"2026-08-11T02:02:49.820001Z","shell.execute_reply.started":"2026-08-11T02:01:49.229594Z","shell.execute_reply":"2026-08-11T02:02:49.819071Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nprint(ablation_table.to_string())\nablation_table.to_csv('/kaggle/working/ablation_table.csv')\ndisplay(ablation_table)\n\nmetrics = ['weighted_log_loss', 'macro_f1', 'auc', 'weighted_kappa']\ntitles = ['Weighted Log Loss (lower is better)', 'Macro F1', 'AUC', 'Weighted Cohen Kappa']\nstage_labels = ['Stage 1\\nText', 'Stage 2\\n2D Fusion', 'Stage 3\\n3D Fusion']\ncolors = ['#4C72B0', '#55A868', '#C44E52']\n\nfig, axes = plt.subplots(1, 4, figsize=(20, 4.5))\nfor ax, metric, title in zip(axes, metrics, titles):\n    values = ablation_table[metric].values.astype(float)\n    bars = ax.bar(stage_labels, values, color=colors)\n    ax.set_title(title, fontsize=12)\n    ax.set_ylim(0, max(values) * 1.25)\n    for bar, v in zip(bars, values):\n        ax.text(bar.get_x() + bar.get_width() / 2, bar.get_height() + max(values) * 0.02,\n                f'{v:.3f}', ha='center', fontsize=10)\n    ax.spines['top'].set_visible(False)\n    ax.spines['right'].set_visible(False)\n\nfig.suptitle('Ablation Comparison Across the Three-Stage Pipeline (RSNA 2024 test split)', fontsize=14)\nplt.tight_layout(rect=[0, 0, 1, 0.94])\nplt.savefig('/kaggle/working/ablation_comparison.png', dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T02:02:49.821549Z","iopub.execute_input":"2026-08-11T02:02:49.821952Z","iopub.status.idle":"2026-08-11T02:02:51.044031Z","shell.execute_reply.started":"2026-08-11T02:02:49.821919Z","shell.execute_reply":"2026-08-11T02:02:51.043384Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom sklearn.metrics import roc_auc_score, confusion_matrix\n\ntarget_names = train_ds_s1.targets\nn_targets_plot = probs.shape[1]\n\nper_target_auc = []\nfor t in range(n_targets_plot):\n    mask = labels[:, t] != -1\n    if mask.sum() == 0:\n        per_target_auc.append(np.nan)\n        continue\n    y_true = labels[mask, t]\n    y_prob = probs[mask, t, :]\n    try:\n        per_target_auc.append(roc_auc_score(y_true, y_prob, multi_class='ovr'))\n    except ValueError:\n        per_target_auc.append(np.nan)\n\norder = np.argsort(per_target_auc)\nsorted_names = [target_names[i] for i in order]\nsorted_aucs = [per_target_auc[i] for i in order]\n\nfig, axes = plt.subplots(1, 2, figsize=(18, 8))\n\naxes[0].barh(sorted_names, sorted_aucs, color='#4C72B0')\naxes[0].set_xlabel('AUC')\naxes[0].set_title('Stage 3 (3D Volumetric Fusion): Per-Condition AUC on Test Split')\naxes[0].set_xlim(0, 1)\naxes[0].axvline(0.5, color='gray', linestyle='--', linewidth=1)\n\nall_true, all_pred = [], []\nfor t in range(n_targets_plot):\n    mask = labels[:, t] != -1\n    if mask.sum() == 0:\n        continue\n    all_true.append(labels[mask, t])\n    all_pred.append(probs[mask, t, :].argmax(axis=1))\nall_true = np.concatenate(all_true)\nall_pred = np.concatenate(all_pred)\ncm = confusion_matrix(all_true, all_pred, labels=[0, 1, 2])\ncm_norm = cm.astype(float) / cm.sum(axis=1, keepdims=True)\n\nim = axes[1].imshow(cm_norm, cmap='Blues', vmin=0, vmax=1)\nclass_names = ['Normal/Mild', 'Moderate', 'Severe']\naxes[1].set_xticks(range(3)); axes[1].set_xticklabels(class_names)\naxes[1].set_yticks(range(3)); axes[1].set_yticklabels(class_names)\naxes[1].set_xlabel('Predicted'); axes[1].set_ylabel('Ground Truth')\naxes[1].set_title('Stage 3: Confusion Matrix (row-normalized, all 25 targets pooled)')\nfor i in range(3):\n    for j in range(3):\n        axes[1].text(j, i, f'{cm_norm[i, j]:.2f}\\n(n={cm[i, j]})',\n                      ha='center', va='center',\n                      color='white' if cm_norm[i, j] > 0.5 else 'black')\nfig.colorbar(im, ax=axes[1], fraction=0.046, pad=0.04)\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/stage3_diagnostics.png', dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T02:02:51.045103Z","iopub.execute_input":"2026-08-11T02:02:51.045655Z","iopub.status.idle":"2026-08-11T02:02:52.520227Z","shell.execute_reply.started":"2026-08-11T02:02:51.045627Z","shell.execute_reply":"2026-08-11T02:02:52.519129Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\nimport numpy as np\nnp.random.seed(42)\nmodel_s1.to(DEVICE)\ns1_test_probs, test_labels = collect_predictions(model_s1, test_loader_s1, forward_s1, tag='Stage 1 recompute')\nmodel_s1.cpu()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\nmodel_s2.to(DEVICE)\ns2_test_probs, _ = collect_predictions(model_s2, test_loader_s2, forward_s2, tag='Stage 2 recompute')\nmodel_s2.cpu()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\nmodel_s3.to(DEVICE)\ns3_test_probs, _ = collect_predictions(model_s3, test_loader_s3, forward_s3, tag='Stage 3 recompute')\nmodel_s3.cpu()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\nn_targets = s1_test_probs.shape[1]\ns1_losses = []\ns2_losses = []\ns3_losses = []\nfor t in range(n_targets):\n    mask = test_labels[:, t] != -1\n    if mask.sum() == 0:\n        continue\n    y_true = test_labels[mask, t]\n    s1_losses.append(rsna_weighted_log_loss(y_true, s1_test_probs[mask, t, :]))\n    s2_losses.append(rsna_weighted_log_loss(y_true, s2_test_probs[mask, t, :]))\n    s3_losses.append(rsna_weighted_log_loss(y_true, s3_test_probs[mask, t, :]))\ns1_losses = np.array(s1_losses)\ns2_losses = np.array(s2_losses)\ns3_losses = np.array(s3_losses)\nobs_diff = (s3_losses - s2_losses).mean()\nboot_diffs = []\nfor _ in range(5000):\n    idx = np.random.choice(len(s2_losses), len(s2_losses), replace=True)\n    boot_diffs.append((s3_losses[idx] - s2_losses[idx]).mean())\nboot_diffs = np.array(boot_diffs)\npval = np.mean(np.abs(boot_diffs) >= np.abs(obs_diff))\nci_lo = np.percentile(boot_diffs, 2.5)\nci_hi = np.percentile(boot_diffs, 97.5)\nprint('S2 vs S3 loss diff:', obs_diff, '95% CI:', ci_lo, ci_hi, 'p-value:', pval)\nstage_probs_list = [('Stage 1', s1_test_probs), ('Stage 2', s2_test_probs), ('Stage 3', s3_test_probs)]\nfor stage_name, probs in stage_probs_list:\n    all_true = []\n    all_pred = []\n    for t in range(n_targets):\n        mask = test_labels[:, t] != -1\n        if mask.sum() == 0:\n            continue\n        all_true.append(test_labels[mask, t])\n        all_pred.append(probs[mask, t, :].argmax(axis=1))\n    all_true = np.concatenate(all_true)\n    all_pred = np.concatenate(all_pred)\n    cm = confusion_matrix(all_true, all_pred, labels=[0, 1, 2])\n    print(stage_name, 'Confusion Matrix:')\n    print(cm)\nprint('Statistical validation complete.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T02:18:10.663825Z","iopub.execute_input":"2026-08-11T02:18:10.664124Z","iopub.status.idle":"2026-08-11T02:19:06.976629Z","shell.execute_reply.started":"2026-08-11T02:18:10.664099Z","shell.execute_reply":"2026-08-11T02:19:06.975595Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"boot_diffs_centered = boot_diffs - boot_diffs.mean()\np_two_sided = 2 * min((boot_diffs_centered <= -abs(obs_diff)).mean(), (boot_diffs_centered >= abs(obs_diff)).mean())\np_two_sided = min(p_two_sided, 1.0)\nci_excludes_zero = (ci_lo > 0) or (ci_hi < 0)\nprint('Corrected two-sided bootstrap p-value (H0: no difference):', p_two_sided)\nprint('95% CI excludes zero:', ci_excludes_zero, '-> difference is statistically significant at alpha=0.05:', ci_excludes_zero)\nprint('Stage 2 mean loss:', s2_losses.mean(), 'Stage 3 mean loss:', s3_losses.mean())\nprint('Absolute loss reduction (S2 to S3):', s2_losses.mean() - s3_losses.mean())\nprint('Relative loss reduction:', (s2_losses.mean() - s3_losses.mean()) / s2_losses.mean() * 100, '%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T02:25:05.034104Z","iopub.execute_input":"2026-08-11T02:25:05.034475Z","iopub.status.idle":"2026-08-11T02:25:05.043645Z","shell.execute_reply.started":"2026-08-11T02:25:05.034442Z","shell.execute_reply":"2026-08-11T02:25:05.042543Z"}},"outputs":[],"execution_count":null}]}