{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"e1e6745c-70b2-4d0a-8d95-96958a0d6299","cell_type":"markdown","source":"# Lumbar Spine Degenerative Classification — Explainable 5-Stage Pipeline\n\n**Dataset:** RSNA 2024 Lumbar Spine Degenerative Classification (Kaggle)\n\n**Pipeline:**\n1. Anatomical Localization Network (3D ConvNeXt) → instance / depth / coordinate prediction → ROI extraction\n2. Graph-based Severity Prediction (EfficientNetV2-L + Slice/Level/Neuro graphs + Cross-Graph Attention + Hierarchical Graph Transformer)\n3. Multi-Level Explainable AI (GradCAM++, Integrated Gradients, GraphCAM, GNNExplainer-lite, attention/node/edge importance, MC-Dropout uncertainty)\n4. Neuro-Symbolic Clinical Reasoning (simplified, honest rule-based consistency checking + confidence calibration — **not** a full clinical knowledge-graph system)\n5. Explainable Clinical Report Generator (structured template, optionally narrated by an LLM — **not** a fine-tuned medical LLM)\n\n> **Scope note:** Stages 1–3 are real, trainable deep-learning components you can fit on the RSNA data.\n> Stages 4–5 are intentionally implemented as transparent, inspectable rule/template systems rather than\n> claiming a genuine clinical-grade knowledge graph or medical LLM — those would require curated clinical\n> ontologies and licensed medical language models that are out of scope for a thesis-level Kaggle notebook.\n","metadata":{}},{"id":"d572277d-249c-49d9-a5fb-324c0d5c7ab9","cell_type":"code","source":"# Kaggle base image already ships torch, torchvision, numpy, pandas, pydicom, pre-built and\n# matched to Kaggle's specific GPU (T4/P100). Do NOT let pip touch torch/torchvision/torchaudio --\n# a generic PyPI torch wheel can silently replace Kaggle's GPU-matched build with one that lacks\n# compiled kernels for that GPU, causing \"no kernel image is available for execution on the\n# device\" the moment any CUDA op runs, often several cells later. Every install below uses\n# --no-deps for exactly this reason; each package's own non-torch dependencies are installed\n# separately, explicitly, without touching torch.\n!pip install -q --no-deps timm\n!pip install -q --no-deps grad-cam ttach\n!pip install -q --no-deps captum\n!pip install -q --no-deps torch-geometric\n!pip install -q networkx\n!pip install -q --no-deps safetensors huggingface_hub\n\nimport importlib\nimport torch\n\n# ---- Verify every package actually installed & imports correctly -------------\n_pkgs_to_check = {\n    \"timm\": \"timm\",\n    \"grad-cam\": \"pytorch_grad_cam\",\n    \"captum\": \"captum\",\n    \"torch-geometric\": \"torch_geometric\",\n    \"networkx\": \"networkx\",\n}\nprint(\"Checking installed packages:\")\n_all_ok = True\nfor _pip_name, _import_name in _pkgs_to_check.items():\n    try:\n        _mod = importlib.import_module(_import_name)\n        _ver = getattr(_mod, \"__version__\", \"unknown\")\n        print(f\"  OK   {_pip_name:<16} -> import {_import_name} (version {_ver})\")\n    except Exception as e:\n        _all_ok = False\n        print(f\"  FAIL {_pip_name:<16} -> {type(e).__name__}: {e}\")\nprint(\"CELL 1 status:\", \"PASSED — all packages importable\" if _all_ok else \"FAILED — see FAIL lines above, re-run pip installs or restart the session\")\nassert _all_ok, \"One or more packages failed to install/import — fix before continuing.\"\n\n# ---- Diagnostics: print exactly what torch build vs. what GPU is assigned ---\n# If the CUDA probe below fails, this block tells us WHY: either the GPU's compute capability\n# isn't in torch's compiled arch list (a real torch/GPU build mismatch -- not fixable by pip\n# installs in this notebook, needs a different torch build or a different GPU type in Kaggle\n# Settings), or something else entirely (e.g. the session was not actually restarted).\nprint(f\"torch.__version__       = {torch.__version__}\")\nprint(f\"torch.version.cuda      = {torch.version.cuda}\")\nprint(f\"torch.cuda.is_available = {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    _arch_list = torch.cuda.get_arch_list()\n    _device_name = torch.cuda.get_device_name(0)\n    _capability = torch.cuda.get_device_capability(0)\n    print(f\"torch.cuda.get_arch_list() = {_arch_list}\")\n    print(f\"GPU device name          = {_device_name}\")\n    print(f\"GPU compute capability   = {_capability}\")\n\n    # Fail fast with an ACTIONABLE message instead of a cryptic CUDA error several lines later,\n    # if this specific GPU's compute capability isn't in what this torch build was compiled for.\n    # This is a real, known situation on Kaggle: recent torch wheels have dropped support for\n    # older GPU generations (e.g. Pascal / P100, sm_60) that Kaggle still offers as an\n    # accelerator option -- no pip install or code change here can fix that; the only fix is\n    # picking a different accelerator (T4 is currently well within the supported range).\n    _min_supported_sm = min(int(a.split(\"_\")[1]) for a in _arch_list if a.startswith(\"sm_\"))\n    _device_sm = _capability[0] * 10 + _capability[1]\n    if _device_sm < _min_supported_sm:\n        raise RuntimeError(\n            f\"GPU '{_device_name}' has compute capability sm_{_device_sm}, but this torch build \"\n            f\"only supports sm_{_min_supported_sm} and above ({_arch_list}). This GPU generation is \"\n            f\"too old for the installed torch -- no pip install or restart fixes this. \"\n            f\"In Kaggle: Settings > Accelerator, switch away from P100 to GPU T4 x2, Save, then \"\n            f\"Run > Restart session and start again from this cell.\"\n        )\nprint()\n!nvidia-smi\n\n# ---- Verify CUDA actually works end-to-end (catches kernel-image mismatches\n# immediately, here, instead of many cells later mid-training) ----------------\nif torch.cuda.is_available():\n    try:\n        _probe = torch.randn(4, 4, device=\"cuda\") @ torch.randn(4, 4, device=\"cuda\")\n        torch.cuda.synchronize()\n        print(f\"CUDA kernel probe OK: {_probe.shape} tensor computed on {_probe.device}\")\n        del _probe\n    except RuntimeError as e:\n        raise RuntimeError(\n            f\"CUDA is reported available but a basic GPU matmul failed ({e}). This means torch's \"\n            f\"CUDA kernels do not match this GPU (see 'no kernel image is available' errors) -- \"\n            f\"restart the session before continuing; if it recurs, the pip installs above are the \"\n            f\"likely cause even with --no-deps.\"\n        ) from e\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:41:41.622958Z","iopub.execute_input":"2026-08-02T05:41:41.623472Z","iopub.status.idle":"2026-08-02T05:42:28.099518Z","shell.execute_reply.started":"2026-08-02T05:41:41.623438Z","shell.execute_reply":"2026-08-02T05:42:28.098433Z"}},"outputs":[],"execution_count":null},{"id":"e15ba188-589d-4fbb-abce-45966fcea151","cell_type":"code","source":"# =====================================================================\n# CELL 2 — Imports and global configuration\n# =====================================================================\nimport os, sys, json, math, glob, random, zipfile, time, warnings, gc\nwarnings.filterwarnings(\"ignore\")\n\n# Reduces GPU memory fragmentation (this is PyTorch's own suggested fix, quoted directly in the\n# OutOfMemoryError message). Must be set before the first CUDA allocation, so it goes here at\n# the very top of the imports cell.\nos.environ.setdefault(\"PYTORCH_CUDA_ALLOC_CONF\", \"expandable_segments:True\")\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\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport timm\nfrom torch_geometric.nn import GATv2Conv, global_mean_pool\nfrom torch_geometric.data import Data, Batch\nimport networkx as nx\nimport resource\n\n\ndef host_ram_gb():\n    '''Peak resident-set size (host RAM, NOT GPU) of this process so far, in GB -- Linux's\n    ru_maxrss is reported in KB. Used to pinpoint exactly which line causes a host-RAM spike,\n    since Kaggle\\'s \"tried to allocate more memory than is available\" OOM message refers to\n    system RAM, not GPU memory, and gives no traceback when it kills the kernel.'''\n    return resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / 1e6\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\n\n# STRICT_GPU_ONLY=True means this notebook REFUSES to run any model computation on CPU: if no\n# GPU is available it raises immediately instead of silently falling back to a (very slow) CPU\n# run. Data loading (DICOM decode via pydicom, resizing via cv2, numpy preprocessing) still runs\n# on CPU inside DataLoader workers regardless of this flag -- that is standard practice for every\n# PyTorch pipeline and is not something CUDA can do; it is also a tiny fraction of total runtime\n# compared to model forward/backward passes, which this flag guarantees stay on GPU.\nSTRICT_GPU_ONLY = True\n\nif STRICT_GPU_ONLY and not torch.cuda.is_available():\n    raise RuntimeError(\n        \"STRICT_GPU_ONLY is set and no GPU is available (torch.cuda.is_available() is False). \"\n        \"In Kaggle: right panel > Settings > Accelerator > GPU T4 x2 (or P100), then Save and \"\n        \"restart the session. Set STRICT_GPU_ONLY = False above if you actually want a CPU run.\"\n    )\n\nDEVICE = torch.device(\"cuda\")  # hardcoded -- STRICT_GPU_ONLY above already guarantees CUDA is available by this point\n\n\ndef assert_on_gpu(module_or_tensor, name=\"\"):\n    '''Call after every .to(DEVICE) / model construction to fail loudly, right at the spot it\n    happened, if something ended up on CPU instead of GPU. Raises RuntimeError (never silently\n    moves anything) identifying the variable name, its shape, its current device, and the\n    expected device -- per the STRICT GPU POLICY, no silent CPU fallback is permitted anywhere.'''\n    if not STRICT_GPU_ONLY:\n        return\n    label = name or \"tensor/model\"\n    if isinstance(module_or_tensor, nn.Module):\n        p = next(module_or_tensor.parameters())\n        dev, shape = p.device, tuple(p.shape)\n    else:\n        dev, shape = module_or_tensor.device, tuple(module_or_tensor.shape)\n    if dev.type != \"cuda\":\n        raise RuntimeError(\n            f\"GPU POLICY VIOLATION -- variable: '{label}', shape: {shape}, \"\n            f\"current device: {dev}, expected device: cuda. STRICT_GPU_ONLY forbids silent CPU \"\n            f\"fallback; fix the .to(...) call at the source instead of relying on this check.\"\n        )\n\n\ndef print_gpu_monitoring_report():\n    '''Prints the exact pre-training GPU status block required by the strict GPU policy: CUDA\n    version, PyTorch build's CUDA version, GPU count/names/memory capacity, and execution mode.'''\n    print(\"=\" * 60)\n    print(f\"CUDA Available   : {torch.cuda.is_available()}\")\n    print(f\"CUDA Version     : {torch.version.cuda}\")\n    print(f\"PyTorch Version  : {torch.__version__}\")\n    print(f\"Detected GPUs    : {N_GPUS}\")\n    for i in range(N_GPUS):\n        _props = torch.cuda.get_device_properties(i)\n        print(f\"GPU {i}:\")\n        print(f\"    {_props.name}\")\n        print(f\"    {_props.total_memory / 1e9:.1f} GB\")\n    print(f\"Execution Mode   : {'Dual-GPU Training (DataParallel)' if N_GPUS > 1 else 'Single-GPU Training'}\")\n    print(f\"CPU Fallback     : DISABLED (STRICT_GPU_ONLY={STRICT_GPU_ONLY})\")\n    print(\"=\" * 60)\n\n\ndef log_and_print(msg, log_path):\n    '''Prints msg as usual AND appends it to log_path on disk immediately (flushed), so if the\n    session dies or disconnects mid-epoch, everything up to the last completed epoch survives on\n    disk in CFG[\"work_dir\"] and can be read/uploaded afterward -- no more losing output because\n    a terminal scrollback wasn't copied in time.'''\n    print(msg)\n    with open(log_path, \"a\") as f:\n        f.write(msg + \"\\n\")\n        f.flush()\n\n\ndef print_per_gpu_memory_report(prefix=\"\"):\n    '''Prints allocated/reserved/peak memory for EVERY visible GPU -- called after every epoch\n    so it is directly verifiable that both GPUs are actually participating, not just GPU 0.'''\n    for i in range(N_GPUS):\n        alloc = torch.cuda.memory_allocated(i) / 1e9\n        reserved = torch.cuda.memory_reserved(i) / 1e9\n        peak = torch.cuda.max_memory_allocated(i) / 1e9\n        print(f\"{prefix}GPU {i} allocated memory : {alloc:.2f} GB\")\n        print(f\"{prefix}GPU {i} reserved memory  : {reserved:.2f} GB\")\n        print(f\"{prefix}GPU {i} peak memory      : {peak:.2f} GB\")\n\n# ---- Kaggle path layout -------------------------------------------------\n# /kaggle/input/<dataset-slug>/          -> read-only competition data\n# /kaggle/working/                       -> writable, persists in notebook output (max ~20GB)\n# /kaggle/temp/                          -> writable scratch space, NOT saved in notebook output\nCFG = {\n    \"data_root\": \"/kaggle/input/competitions/rsna-2024-lumbar-spine-degenerative-classification\",\n    \"work_dir\": \"/kaggle/working\",\n    \"temp_dir\": \"/kaggle/temp\",\n    \"img_size_2d\": 224,            # ROI crop size fed to EfficientNetV2-L\n    \"vol_size_3d\": (32, 128, 128), # (D, H, W) resampled volume for Stage 1 localization net\n    \"num_levels\": 5,               # L1/L2 .. L5/S1\n    \"conditions\": [\n        \"spinal_canal_stenosis\",\n        \"left_neural_foraminal_narrowing\",\n        \"right_neural_foraminal_narrowing\",\n        \"left_subarticular_stenosis\",\n        \"right_subarticular_stenosis\",\n    ],\n    \"num_classes\": 3,              # Normal/Mild, Moderate, Severe\n    \"levels\": [\"l1_l2\", \"l2_l3\", \"l3_l4\", \"l4_l5\", \"l5_s1\"],\n    \"batch_size_stage1\": 2,\n    \"batch_size_stage2\": 2,  # reduced from 4 -- with dual-GPU DataParallel this means 1 study per GPU per step instead of 2, to rule out the batch-of-4-per-GPU-pair load as the cause of the silent kernel crash\n    \"lr_stage1\": 3e-4,\n    \"lr_stage2\": 1e-4,\n    \"epochs_stage1\": 15,\n    \"epochs_stage2\": 20,\n    \"num_workers\": max(2, (os.cpu_count() or 2) - 1),  # leave 1 core free for the main process\n    \"graph_hidden\": 256,\n    \"graph_heads\": 4,\n    \"roi_embed_chunk_size\": 16,\n}\n\nos.makedirs(CFG[\"work_dir\"], exist_ok=True)\nos.makedirs(CFG[\"temp_dir\"], exist_ok=True)\nprint(\"Device:\", DEVICE)\nprint(json.dumps({k: v for k, v in CFG.items() if k not in (\"data_root\",)}, indent=2))\n\n# ---- GPU diagnostics -- check this FIRST, before running anything else ------\n# This tells you immediately whether the GPU is actually available to this notebook and how much\n# of it is already in use BEFORE any of our code has allocated a single tensor. If \"free\" here is\n# already small, that memory is NOT from this notebook's code -- it is left over from a previous\n# session that was not fully restarted (or another session/tab still holds the GPU), and no code\n# fix here will free it; you need to actually restart the Kaggle session (see below).\nif not torch.cuda.is_available():\n    print(\"\\n\" + \"!\" * 70)\n    print(\"WARNING: torch.cuda.is_available() is False -- this notebook will run ENTIRELY ON CPU.\")\n    print(\"Training will be extremely slow. In Kaggle: Settings (right panel) > Accelerator > GPU T4 x2 (or P100),\")\n    print(\"then Save & re-run. Also double-check under Settings that a GPU quota is actually available to you.\")\n    print(\"!\" * 70 + \"\\n\")\nelse:\n    _free_bytes, _total_bytes = torch.cuda.mem_get_info()\n    _free_gb, _total_gb = _free_bytes / 1e9, _total_bytes / 1e9\n    print(f\"\\nGPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"GPU memory: {_free_gb:.2f} GB free / {_total_gb:.2f} GB total \"\n          f\"(already allocated by PyTorch in this process: {torch.cuda.memory_allocated() / 1e9:.2f} GB)\")\n    if _free_gb < _total_gb * 0.5:\n        print(\"\\n\" + \"!\" * 70)\n        print(\"WARNING: less than half the GPU is free at the very start of this notebook, before any\")\n        print(\"training has happened. This is almost always LEFTOVER memory from a previous session that\")\n        print(\"was not fully restarted. Fix: Kaggle menu -> Run -> 'Restart session' (not just re-running\")\n        print(\"cells), then start again from this cell. If it's still full after a real restart, check\")\n        print(\"Kaggle's session list (top right) for another running session of this same notebook and\")\n        print(\"stop it, since two sessions can share/exhaust the same GPU quota.\")\n        print(\"!\" * 70 + \"\\n\")\n\ndef cell_status(name, ok=True, extra=\"\"):\n    '''Shared helper used at the end of every cell below to print a clear, greppable\n    PASSED/FAILED line, so you can scroll the notebook and immediately see which cells\n    ran successfully without reading every line of output.'''\n    tag = \"PASSED\" if ok else \"FAILED\"\n    msg = f\"[{tag}] {name}\"\n    if extra:\n        msg += f\" — {extra}\"\n    print(msg)\n\ncell_status(\"CELL 2 — imports & global configuration\", extra=f\"device={DEVICE}\")\n\n# =====================================================================\n# Multi-GPU + fast-DataLoader helpers (used by every training/eval cell below)\n# =====================================================================\nN_GPUS = torch.cuda.device_count() if torch.cuda.is_available() else 0\n\ndef report_gpu_status(prefix=\"\"):\n    if N_GPUS == 0:\n        print(f\"{prefix}No GPU detected.\")\n        return\n    print(f\"{prefix}Detected GPUs: {N_GPUS}\")\n    if N_GPUS > 1:\n        print(f\"{prefix}Using Multi-GPU Training (torch.nn.DataParallel)\")\n    for i in range(N_GPUS):\n        print(f\"{prefix}GPU {i}: {torch.cuda.get_device_name(i)}\")\n\n\ndef wrap_for_multi_gpu(model):\n    '''Wraps model in nn.DataParallel across every visible GPU when more than one is available.\n    True DistributedDataParallel needs a separate process per GPU (torchrun/mp.spawn with its\n    own rendezvous), which is not practical inside a single interactive Kaggle notebook cell --\n    DataParallel is the documented fallback for exactly this situation and works within one\n    process, splitting each batch across GPUs and gathering gradients automatically.'''\n    if N_GPUS > 1:\n        return nn.DataParallel(model)\n    return model\n\n\ndef unwrap_model(model):\n    '''Returns the underlying model whether or not it is wrapped in DataParallel, so checkpoint\n    saving/loading and single-GPU inference always use a plain (non-wrapped) state_dict.'''\n    return model.module if isinstance(model, nn.DataParallel) else model\n\n\ndef make_fast_loader_kwargs(num_workers):\n    '''Standard DataLoader speed knobs: pinned host memory for fast non_blocking H2D copies,\n    persistent workers to avoid respawning worker processes every epoch, and a modest prefetch\n    factor to keep the GPU fed. persistent_workers/prefetch_factor are only valid when\n    num_workers > 0, so they are omitted entirely otherwise (passing them with num_workers=0\n    raises an error in PyTorch).'''\n    kwargs = {\"pin_memory\": torch.cuda.is_available()}\n    if num_workers > 0:\n        kwargs[\"persistent_workers\"] = True\n        kwargs[\"prefetch_factor\"] = 4\n    return kwargs\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:44:36.060785Z","iopub.execute_input":"2026-08-02T05:44:36.061178Z","iopub.status.idle":"2026-08-02T05:44:36.547701Z","shell.execute_reply.started":"2026-08-02T05:44:36.061142Z","shell.execute_reply":"2026-08-02T05:44:36.546979Z"}},"outputs":[],"execution_count":null},{"id":"f2d96901-dc44-440b-8de0-3b314056ea55","cell_type":"markdown","source":"## Stage 0 — Dataset layer\n\nRSNA 2024 Lumbar Spine Degenerative Classification provides, per `study_id`:\n- `train.csv` — one row per study, columns `<condition>_<level>` with values `Normal/Mild`, `Moderate`, `Severe`\n- `train_series_descriptions.csv` — `study_id, series_id, series_description` (Sagittal T1 / Sagittal T2-STIR / Axial T2)\n- `train_label_coordinates.csv` — `study_id, series_id, instance_number, condition, level, x, y` pixel coordinates used to supervise Stage 1\n- DICOM images under `train_images/<study_id>/<series_id>/<instance_number>.dcm`\n\nWe build two datasets: one that yields whole (resampled) volumes for Stage 1, and one that yields\nalready-localized ROI crops for Stage 2 (crops are produced on the fly from Stage 1 predictions at\ninference time, or from ground-truth coordinates during Stage 2 training/pretraining).\n","metadata":{}},{"id":"180658d3-d511-440c-bb9c-de5af0f614f5","cell_type":"code","source":"# =====================================================================\n# CELL 3 — Low level IO helpers\n# =====================================================================\n\ndef load_dicom_series(study_id, series_id, data_root=CFG[\"data_root\"], split=\"train\"):\n    '''Load every DICOM instance of one series, sorted by instance number.\n    Returns:\n      volume        (D, H, W) float32 in [0, 1], resized in-plane to CFG[\"vol_size_3d\"][1:]\n      instance_nums list[int] original DICOM InstanceNumber for each slice, in volume order\n      orig_hw       (H_orig, W_orig) of the raw DICOM pixel array, BEFORE resizing.\n                     train_label_coordinates.csv gives x, y in this original pixel space, so any\n                     code normalizing those coordinates must divide by orig_hw, not by volume.shape.\n    '''\n    series_dir = os.path.join(data_root, f\"{split}_images\", str(study_id), str(series_id))\n    paths = sorted(glob.glob(os.path.join(series_dir, \"*.dcm\")),\n                    key=lambda p: int(os.path.splitext(os.path.basename(p))[0]))\n    slices = []\n    orig_hw = None\n    for p in paths:\n        dcm = pydicom.dcmread(p)\n        img = dcm.pixel_array.astype(np.float32)\n        if orig_hw is None:\n            orig_hw = img.shape  # (H_orig, W_orig) — same for every slice in a series\n        # Min-max normalize each slice independently (robust to scanner intensity drift)\n        img = (img - img.min()) / (img.max() - img.min() + 1e-6)\n        # Guard against corrupt DICOM pixel data (NaN/Inf) -- a single bad slice like this\n        # would otherwise poison every downstream loss for the whole training epoch.\n        img = np.nan_to_num(img, nan=0.0, posinf=1.0, neginf=0.0)\n        slices.append(img)\n    if len(slices) == 0:\n        return None, [], None\n    # Resize every slice to a common in-plane size before stacking\n    h, w = CFG[\"vol_size_3d\"][1], CFG[\"vol_size_3d\"][2]\n    slices = [cv2.resize(s, (w, h), interpolation=cv2.INTER_LINEAR) for s in slices]\n    volume = np.stack(slices, axis=0)  # (D, H, W)\n    instance_nums = [int(os.path.splitext(os.path.basename(p))[0]) for p in paths]\n    return volume, instance_nums, orig_hw\n\n\ndef resample_depth(volume, target_d):\n    '''Resample the depth (slice) axis to a fixed number of slices via linear interpolation\n    on indices — keeps the network input shape constant across studies with a different\n    number of DICOM slices.'''\n    d = volume.shape[0]\n    if d == target_d:\n        return volume\n    idx = np.linspace(0, d - 1, target_d)\n    idx_floor = np.floor(idx).astype(int)\n    idx_ceil = np.clip(idx_floor + 1, 0, d - 1)\n    frac = (idx - idx_floor)[:, None, None]\n    resampled = volume[idx_floor] * (1 - frac) + volume[idx_ceil] * frac\n    return resampled.astype(np.float32)\n\n\ndef severity_to_label(value):\n    mapping = {\"Normal/Mild\": 0, \"Moderate\": 1, \"Severe\": 2}\n    return mapping.get(value, 0)\n\n# ---- Self-test with synthetic data (no DICOM files needed) -------------------\ntry:\n    _fake_vol = np.random.rand(20, 64, 64).astype(np.float32)\n    _resampled = resample_depth(_fake_vol, CFG[\"vol_size_3d\"][0])\n    assert _resampled.shape == (CFG[\"vol_size_3d\"][0], 64, 64), f\"unexpected shape {_resampled.shape}\"\n    assert severity_to_label(\"Normal/Mild\") == 0\n    assert severity_to_label(\"Moderate\") == 1\n    assert severity_to_label(\"Severe\") == 2\n    cell_status(\"CELL 3 — IO helpers\", extra=f\"resample_depth {_fake_vol.shape} -> {_resampled.shape}, severity_to_label OK\")\nexcept Exception as e:\n    cell_status(\"CELL 3 — IO helpers\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:44:44.325671Z","iopub.execute_input":"2026-08-02T05:44:44.326336Z","iopub.status.idle":"2026-08-02T05:44:44.344100Z","shell.execute_reply.started":"2026-08-02T05:44:44.326300Z","shell.execute_reply":"2026-08-02T05:44:44.343448Z"}},"outputs":[],"execution_count":null},{"id":"06f09e71-47f7-4070-afe5-6773a193563c","cell_type":"code","source":"# =====================================================================\n# CELL 4 — Stage 1 dataset: whole-volume localization targets\n# =====================================================================\nclass LocalizationDataset(Dataset):\n    '''Yields a resampled 3D volume plus, for every of the 5 disc levels, the ground-truth\n    (instance_index, y, x) coordinate normalized to [0, 1]. Instance-presence and depth targets\n    let the network learn *which* slice and *how deep* a level sits before we crop the ROI.'''\n\n    def __init__(self, df_coords, df_series, data_root=CFG[\"data_root\"], split=\"train\"):\n        self.data_root = data_root\n        self.split = split\n        self.df_series = df_series\n        # one sample per (study_id, series_id) that has at least one labelled coordinate\n        self.groups = list(df_coords.groupby([\"study_id\", \"series_id\"]))\n        # ON-DISK cache only (.npz files under CFG[\"temp_dir\"]) for the DECODED + RESAMPLED\n        # volume per (study_id, series_id) -- by far the most expensive part of __getitem__\n        # (pydicom decode of every slice + cv2 resize + depth resampling). Deliberately NOT also\n        # kept in an in-memory Python dict: an earlier version added an unbounded per-worker-\n        # process in-memory cache on TOP of this disk cache, and with num_workers > 1 each worker\n        # accumulates its own full copy with no eviction, which exhausted system RAM and got a\n        # DataLoader worker SIGKILL'd by the Linux OOM killer. The disk cache alone already gives\n        # the real fix (shared across every worker process, unlike an in-memory dict), and the\n        # Linux page cache transparently keeps recently-read .npz files in RAM anyway -- with\n        # proper eviction under memory pressure, which our own dict could never do.\n        self._disk_cache_dir = os.path.join(CFG[\"temp_dir\"], \"loc_volume_cache\")\n        os.makedirs(self._disk_cache_dir, exist_ok=True)\n\n    def _disk_cache_path(self, study_id, series_id):\n        return os.path.join(self._disk_cache_dir, f\"{study_id}_{series_id}.npz\")\n\n    def __len__(self):\n        return len(self.groups)\n\n    def __getitem__(self, i):\n        (study_id, series_id), rows = self.groups[i]\n        cache_key = (study_id, series_id)\n\n        _empty_sample = {\n            \"volume\": torch.zeros(1, *CFG[\"vol_size_3d\"], dtype=torch.float32),\n            \"coords\": torch.zeros(CFG[\"num_levels\"], 3, dtype=torch.float32),\n            \"presence\": torch.zeros(CFG[\"num_levels\"], dtype=torch.float32),\n            \"study_id\": study_id, \"series_id\": series_id,\n        }\n\n        disk_path = self._disk_cache_path(study_id, series_id)\n        if os.path.exists(disk_path):\n            with np.load(disk_path) as npz:\n                volume = npz[\"volume\"]\n                instance_numbers = npz[\"instance_numbers\"].tolist()\n                orig_hw = tuple(npz[\"orig_hw\"].tolist())\n                d_orig = int(npz[\"d_orig\"])\n        else:\n            volume, instance_numbers, orig_hw = load_dicom_series(study_id, series_id, self.data_root, self.split)\n            if volume is None:\n                # Series had no readable DICOM files — return an empty (all-zero, all-absent)\n                # sample rather than crashing the DataLoader; the loss masks out absent levels.\n                # (Deliberately NOT disk-cached: a transient read failure shouldn't be baked into\n                # a permanent \"this series is empty\" file on disk.)\n                return _empty_sample\n            d_orig = len(instance_numbers)\n            volume = resample_depth(volume, CFG[\"vol_size_3d\"][0])\n            try:\n                np.savez(disk_path, volume=volume, instance_numbers=np.array(instance_numbers),\n                          orig_hw=np.array(orig_hw), d_orig=d_orig)\n            except OSError as e:\n                print(f\"  [WARN] Could not write disk cache for {cache_key}: {e} (continuing without it)\")\n\n        orig_h, orig_w = orig_hw\n\n        coord_target = np.zeros((CFG[\"num_levels\"], 3), dtype=np.float32)  # (depth, y, x)\n        presence = np.zeros((CFG[\"num_levels\"],), dtype=np.float32)\n        level_to_idx = {lvl: i for i, lvl in enumerate(CFG[\"levels\"])}\n\n        for _, r in rows.iterrows():\n            lvl = str(r[\"level\"]).lower().replace(\"/\", \"_\")\n            if lvl not in level_to_idx:\n                continue\n            li = level_to_idx[lvl]\n            inst_pos = instance_numbers.index(int(r[\"instance_number\"])) if int(r[\"instance_number\"]) in instance_numbers else 0\n            depth_norm = inst_pos / max(1, d_orig - 1)\n            depth_resampled = depth_norm  # normalized depth is invariant to resampling\n            # IMPORTANT: x, y in train_label_coordinates.csv are in the ORIGINAL DICOM pixel\n            # space, not the resized (vol_size_3d) space — normalize by orig_h / orig_w, not by\n            # volume.shape, so the [0,1] coordinate is correct regardless of the resize target.\n            y_norm = float(r[\"y\"]) / orig_h\n            x_norm = float(r[\"x\"]) / orig_w\n            coord_target[li] = [depth_resampled, y_norm, x_norm]\n            presence[li] = 1.0\n\n        vol_t = torch.from_numpy(volume).unsqueeze(0).float()  # (1, D, H, W)\n        return {\n            \"volume\": vol_t,\n            \"coords\": torch.from_numpy(coord_target),\n            \"presence\": torch.from_numpy(presence),\n            \"study_id\": study_id,\n            \"series_id\": series_id,\n        }\n\n# ---- Self-test: class is importable and has the expected interface ----------\ntry:\n    assert issubclass(LocalizationDataset, Dataset)\n    assert hasattr(LocalizationDataset, \"__getitem__\") and hasattr(LocalizationDataset, \"__len__\")\n    cell_status(\"CELL 4 — LocalizationDataset\", extra=\"class defined correctly (no DICOM data touched yet)\")\nexcept Exception as e:\n    cell_status(\"CELL 4 — LocalizationDataset\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:44:49.798877Z","iopub.execute_input":"2026-08-02T05:44:49.799671Z","iopub.status.idle":"2026-08-02T05:44:49.814159Z","shell.execute_reply.started":"2026-08-02T05:44:49.799641Z","shell.execute_reply":"2026-08-02T05:44:49.813606Z"}},"outputs":[],"execution_count":null},{"id":"9f02d4c8-fa5c-4590-83f0-9948738654eb","cell_type":"markdown","source":"## Stage 1 — Anatomical Localization Network\n\nA lightweight **3D ConvNeXt** encodes the whole MRI volume once. Three heads share the same\nbackbone features (multi-task learning improves localization because \"which level is present\",\n\"which slice it's on\" and \"where in-plane it is\" are correlated signals):\n\n- **Instance Prediction** — per-level presence/confidence (is this level visible in this series?)\n- **Depth Prediction** — normalized slice index of each level\n- **Coordinate Prediction** — normalized (y, x) in-plane location of each level\n\nThe predicted (depth, y, x) triplets are then used to crop a small 3D ROI around each disc level,\nwhich becomes the input to Stage 2.\n","metadata":{}},{"id":"ef497276-7931-4eec-b2bd-9783bb0c2597","cell_type":"code","source":"# =====================================================================\n# CELL 5 — 3D ConvNeXt backbone (implemented from scratch: torchvision only\n# ships the 2D ConvNeXt, so we port the block design to 3D convolutions)\n# =====================================================================\n\nclass ConvNeXtBlock3D(nn.Module):\n    '''Depthwise 7x7x7 conv -> LayerNorm -> pointwise MLP (4x expansion) -> residual.\n    Mirrors the 2D ConvNeXt block design (Liu et al., 2022) with 3D depthwise convolutions.'''\n\n    def __init__(self, dim, drop_path=0.0):\n        super().__init__()\n        self.dwconv = nn.Conv3d(dim, dim, kernel_size=7, padding=3, groups=dim)\n        self.norm = nn.LayerNorm(dim, eps=1e-6)\n        self.pwconv1 = nn.Linear(dim, 4 * dim)\n        self.act = nn.GELU()\n        self.pwconv2 = nn.Linear(4 * dim, dim)\n        self.gamma = nn.Parameter(1e-6 * torch.ones(dim), requires_grad=True)\n        self.drop_path = nn.Identity()  # kept simple; swap for stochastic depth if needed\n\n    def forward(self, x):\n        inp = x\n        x = self.dwconv(x)\n        x = x.permute(0, 2, 3, 4, 1)          # (B, D, H, W, C) for channel-last LayerNorm\n        x = self.norm(x)\n        x = self.pwconv1(x)\n        x = self.act(x)\n        x = self.pwconv2(x)\n        x = self.gamma * x\n        x = x.permute(0, 4, 1, 2, 3)          # back to (B, C, D, H, W)\n        return inp + self.drop_path(x)\n\n\nclass ConvNeXt3D(nn.Module):\n    '''4-stage 3D ConvNeXt encoder. Depths/dims follow the \"Tiny\" recipe scaled down for the\n    small volume sizes used here (32x128x128), keeping the model trainable on a single Kaggle GPU.'''\n\n    def __init__(self, in_chans=1, depths=(2, 2, 4, 2), dims=(32, 64, 128, 256)):\n        super().__init__()\n        self.downsample_layers = nn.ModuleList()\n        stem = nn.Sequential(\n            nn.Conv3d(in_chans, dims[0], kernel_size=4, stride=4),\n            LayerNormChannelsFirst3D(dims[0]),\n        )\n        self.downsample_layers.append(stem)\n        for i in range(3):\n            self.downsample_layers.append(nn.Sequential(\n                LayerNormChannelsFirst3D(dims[i]),\n                nn.Conv3d(dims[i], dims[i + 1], kernel_size=2, stride=2),\n            ))\n\n        self.stages = nn.ModuleList()\n        for i in range(4):\n            stage = nn.Sequential(*[ConvNeXtBlock3D(dims[i]) for _ in range(depths[i])])\n            self.stages.append(stage)\n\n        self.out_dim = dims[-1]\n\n    def forward(self, x):\n        feats = []\n        for i in range(4):\n            x = self.downsample_layers[i](x)\n            x = self.stages[i](x)\n            feats.append(x)\n        return x, feats  # x: (B, C, D', H', W') final feature map; feats: multi-scale features\n\n\nclass LayerNormChannelsFirst3D(nn.Module):\n    '''LayerNorm over the channel dim for (B, C, D, H, W) tensors.'''\n    def __init__(self, num_channels, eps=1e-6):\n        super().__init__()\n        self.weight = nn.Parameter(torch.ones(num_channels))\n        self.bias = nn.Parameter(torch.zeros(num_channels))\n        self.eps = eps\n\n    def forward(self, x):\n        u = x.mean(1, keepdim=True)\n        s = (x - u).pow(2).mean(1, keepdim=True)\n        x = (x - u) / torch.sqrt(s + self.eps)\n        x = self.weight[None, :, None, None, None] * x + self.bias[None, :, None, None, None]\n        return x\n\n# ---- Self-test: forward pass on a random volume of the configured size ------\ntry:\n    _test_backbone = ConvNeXt3D(in_chans=1)\n    _test_input = torch.randn(1, 1, *CFG[\"vol_size_3d\"])\n    with torch.no_grad():\n        _out, _feats = _test_backbone(_test_input)\n    _n_params = sum(p.numel() for p in _test_backbone.parameters())\n    cell_status(\"CELL 5 — 3D ConvNeXt backbone\", extra=f\"input {tuple(_test_input.shape)} -> output {tuple(_out.shape)}, {_n_params:,} params\")\n    del _test_backbone, _test_input, _out, _feats\nexcept Exception as e:\n    cell_status(\"CELL 5 — 3D ConvNeXt backbone\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:44:54.121444Z","iopub.execute_input":"2026-08-02T05:44:54.121848Z","iopub.status.idle":"2026-08-02T05:44:54.400576Z","shell.execute_reply.started":"2026-08-02T05:44:54.121820Z","shell.execute_reply":"2026-08-02T05:44:54.399733Z"}},"outputs":[],"execution_count":null},{"id":"2514c43e-073f-43dc-b156-871331657f46","cell_type":"code","source":"# =====================================================================\n# CELL 6 — Stage 1 multi-task heads + full localization model\n# =====================================================================\nclass AnatomicalLocalizationNet(nn.Module):\n    def __init__(self, num_levels=CFG[\"num_levels\"]):\n        super().__init__()\n        self.backbone = ConvNeXt3D(in_chans=1)\n        feat_dim = self.backbone.out_dim\n        self.pool = nn.AdaptiveAvgPool3d(1)\n\n        # Shared trunk before splitting into task-specific heads\n        self.trunk = nn.Sequential(nn.Linear(feat_dim, 256), nn.GELU(), nn.Dropout(0.2))\n\n        # Instance Prediction: presence probability of each disc level in this series\n        self.instance_head = nn.Linear(256, num_levels)\n        # Depth Prediction: normalized slice index in [0, 1] for each level\n        self.depth_head = nn.Linear(256, num_levels)\n        # Coordinate Prediction: normalized (y, x) in-plane location for each level\n        self.coord_head = nn.Linear(256, num_levels * 2)\n\n        self.num_levels = num_levels\n\n    def forward(self, volume):\n        feat_map, _ = self.backbone(volume)          # (B, C, D', H', W')\n        pooled = self.pool(feat_map).flatten(1)       # (B, C)\n        z = self.trunk(pooled)\n\n        instance_logits = self.instance_head(z)                          # (B, L)\n        depth_pred = torch.sigmoid(self.depth_head(z))                   # (B, L) in [0,1]\n        coord_pred = torch.sigmoid(self.coord_head(z)).view(-1, self.num_levels, 2)  # (B, L, 2) -> (y, x)\n\n        return {\n            \"instance_logits\": instance_logits,\n            \"depth_pred\": depth_pred,\n            \"coord_pred\": coord_pred,\n            \"feat_map\": feat_map,   # exposed for GradCAM++ in Stage 3\n        }\n\n# ---- Self-test: full Stage 1 model forward pass + output shape checks -------\ntry:\n    _test_model = AnatomicalLocalizationNet()\n    _test_vol = torch.randn(2, 1, *CFG[\"vol_size_3d\"])\n    with torch.no_grad():\n        _out = _test_model(_test_vol)\n    assert _out[\"instance_logits\"].shape == (2, CFG[\"num_levels\"]), _out[\"instance_logits\"].shape\n    assert _out[\"depth_pred\"].shape == (2, CFG[\"num_levels\"]), _out[\"depth_pred\"].shape\n    assert _out[\"coord_pred\"].shape == (2, CFG[\"num_levels\"], 2), _out[\"coord_pred\"].shape\n    assert (_out[\"depth_pred\"] >= 0).all() and (_out[\"depth_pred\"] <= 1).all(), \"depth_pred not in [0,1]\"\n    _n_params = sum(p.numel() for p in _test_model.parameters())\n    cell_status(\"CELL 6 — AnatomicalLocalizationNet\", extra=f\"batch=2 forward OK, all head shapes correct, {_n_params:,} params\")\n    del _test_model, _test_vol, _out\nexcept Exception as e:\n    cell_status(\"CELL 6 — AnatomicalLocalizationNet\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:44:59.955240Z","iopub.execute_input":"2026-08-02T05:44:59.955982Z","iopub.status.idle":"2026-08-02T05:45:00.066440Z","shell.execute_reply.started":"2026-08-02T05:44:59.955949Z","shell.execute_reply":"2026-08-02T05:45:00.065596Z"}},"outputs":[],"execution_count":null},{"id":"fafce2b9-f9dd-40aa-997f-799470bc5115","cell_type":"code","source":"# =====================================================================\n# CELL 7 — Stage 1 loss and training loop\n# =====================================================================\ndef stage1_loss(outputs, coords_gt, presence_gt):\n    # Instance presence: binary cross-entropy per level\n    loss_instance = F.binary_cross_entropy_with_logits(outputs[\"instance_logits\"], presence_gt)\n\n    # Depth / coordinate regression only where the level is actually present (masked L1 loss)\n    mask = presence_gt.unsqueeze(-1)                              # (B, L, 1)\n    depth_gt = coords_gt[..., 0]\n    yx_gt = coords_gt[..., 1:]\n\n    loss_depth = (F.l1_loss(outputs[\"depth_pred\"], depth_gt, reduction=\"none\") * presence_gt).sum() \\\n                 / presence_gt.sum().clamp(min=1)\n    loss_coord = (F.l1_loss(outputs[\"coord_pred\"], yx_gt, reduction=\"none\") * mask).sum() \\\n                 / mask.sum().clamp(min=1)\n\n    total = loss_instance + 2.0 * loss_depth + 2.0 * loss_coord\n    return total, {\"instance\": loss_instance.item(), \"depth\": loss_depth.item(), \"coord\": loss_coord.item()}\n\n\ndef stage1_accuracy(outputs, presence_gt):\n    '''Instance-presence accuracy: fraction of (study, level) cells where the predicted\n    presence/absence (sigmoid(instance_logits) > 0.5) matches the ground truth. Depth/coordinate\n    regression has no natural \"accuracy\" (it is continuous), so this is the metric reported as\n    Stage 1's accuracy each epoch.'''\n    with torch.no_grad():\n        pred = (torch.sigmoid(outputs[\"instance_logits\"]) > 0.5).float()\n        return (pred == presence_gt).float().mean().item()\n\n\ndef train_stage1(model, train_loader, val_loader, cfg=CFG, epochs=None, grad_clip=1.0):\n    # NOTE: Stage 1 trains in plain fp32, NOT mixed precision. The hand-written 3D ConvNeXt uses\n    # several custom LayerNorm ops that are prone to overflow/underflow under fp16 autocast --\n    # that is the most likely cause of the widespread \"all parts nan\" losses seen in fp16 runs.\n    # This model is small enough that fp32 costs little extra memory/time, so correctness wins.\n    epochs = epochs or cfg[\"epochs_stage1\"]\n    model.to(DEVICE)\n    assert_on_gpu(model, \"Stage1 model\")\n    print_gpu_monitoring_report()\n    model = wrap_for_multi_gpu(model)\n    print(f\"[Stage1] Training on device: {next(model.parameters()).device}\"\n          f\"{' (DataParallel across ' + str(N_GPUS) + ' GPUs)' if isinstance(model, nn.DataParallel) else ''}\")\n    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg[\"lr_stage1\"], weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)\n\n    log_path = os.path.join(cfg[\"work_dir\"], \"stage1_training_log.txt\")\n    with open(log_path, \"a\") as f:\n        f.write(f\"\\n===== New Stage 1 run started, {epochs} epochs =====\\n\")\n    print(f\"[Stage1] Logging every epoch to: {log_path}\")\n\n    history = {\"train_loss\": [], \"val_loss\": [], \"train_acc\": [], \"val_acc\": []}\n    best_epoch, best_val_loss = -1, float(\"inf\")\n    for epoch in range(epochs):\n        epoch_start = time.time()\n        model.train()\n        running, running_acc = 0.0, 0.0\n        n_valid_batches = 0\n        n_skipped = 0\n        for batch_idx, batch in enumerate(train_loader):\n            vol = batch[\"volume\"].to(DEVICE, non_blocking=True)\n            coords = batch[\"coords\"].to(DEVICE, non_blocking=True)\n            presence = batch[\"presence\"].to(DEVICE, non_blocking=True)\n\n            # Guard the INPUT too, not just the loss -- if a volume slipped through with NaN/Inf\n            # (e.g. an unusual DICOM RescaleSlope/Intercept), catch it before it ever reaches the\n            # model, since by the time the loss is NaN it's too late to tell input vs. numerical\n            # instability apart.\n            if not torch.isfinite(vol).all():\n                n_skipped += 1\n                log_and_print(f\"  [WARN] Stage1 epoch {epoch+1} batch {batch_idx}: non-finite INPUT volume — skipping\", log_path)\n                continue\n\n            optimizer.zero_grad(set_to_none=True)\n            out = model(vol)\n            loss, parts = stage1_loss(out, coords, presence)\n\n            # A single corrupt sample (bad DICOM, no labels at all, etc.) producing a NaN/Inf\n            # loss would otherwise poison the running sum for the WHOLE epoch (NaN + anything =\n            # NaN forever) -- skip just that batch instead of letting it wreck every reported\n            # metric, and warn so it's visible which batch/study caused it.\n            if not torch.isfinite(loss):\n                n_skipped += 1\n                log_and_print(f\"  [WARN] Stage1 epoch {epoch+1} batch {batch_idx}: non-finite loss \"\n                              f\"({loss.item()!r}, parts={parts}) — skipping this batch's update\", log_path)\n                continue\n\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip)\n            optimizer.step()\n            running += loss.item()\n            running_acc += stage1_accuracy(out, presence)\n            n_valid_batches += 1\n\n        scheduler.step()\n        train_loss = running / max(1, n_valid_batches)\n        train_acc = running_acc / max(1, n_valid_batches)\n        if n_skipped > 0:\n            log_and_print(f\"  Stage1 epoch {epoch+1}: skipped {n_skipped}/{len(train_loader)} batches due to non-finite loss\", log_path)\n\n        model.eval()\n        val_running, val_running_acc = 0.0, 0.0\n        n_valid_val = 0\n        with torch.no_grad():\n            for batch in val_loader:\n                vol = batch[\"volume\"].to(DEVICE, non_blocking=True)\n                coords = batch[\"coords\"].to(DEVICE, non_blocking=True)\n                presence = batch[\"presence\"].to(DEVICE, non_blocking=True)\n                out = model(vol)\n                loss, _ = stage1_loss(out, coords, presence)\n                if torch.isfinite(loss):\n                    val_running += loss.item()\n                    val_running_acc += stage1_accuracy(out, presence)\n                    n_valid_val += 1\n        val_loss = val_running / max(1, n_valid_val)\n        val_acc = val_running_acc / max(1, n_valid_val)\n\n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss)\n        history[\"train_acc\"].append(train_acc)\n        history[\"val_acc\"].append(val_acc)\n        elapsed = time.time() - epoch_start\n\n        log_and_print(f\"[Stage1][Epoch {epoch+1}/{epochs}]\", log_path)\n        log_and_print(f\"  Train Loss     : {train_loss:.4f}\", log_path)\n        log_and_print(f\"  Train Accuracy : {train_acc:.2%}\", log_path)\n        log_and_print(f\"  Val Loss       : {val_loss:.4f}\", log_path)\n        log_and_print(f\"  Val Accuracy   : {val_acc:.2%}\", log_path)\n        log_and_print(f\"  Elapsed Time   : {elapsed:.1f}s\", log_path)\n\n        plain_model = unwrap_model(model)\n        torch.save(plain_model.state_dict(), os.path.join(cfg[\"work_dir\"], \"stage1_last.pt\"))\n        save_status = \"last checkpoint saved\"\n        if val_loss < best_val_loss:\n            best_val_loss, best_epoch = val_loss, epoch + 1\n            torch.save(plain_model.state_dict(), os.path.join(cfg[\"work_dir\"], \"stage1_best.pt\"))\n            save_status = \"last + BEST checkpoint saved\"\n        log_and_print(f\"  Best Epoch So Far : {best_epoch} (val_loss={best_val_loss:.4f})\", log_path)\n        log_and_print(f\"  Model Save Status : {save_status}\", log_path)\n        for i in range(N_GPUS):\n            alloc = torch.cuda.memory_allocated(i) / 1e9\n            reserved = torch.cuda.memory_reserved(i) / 1e9\n            peak = torch.cuda.max_memory_allocated(i) / 1e9\n            log_and_print(f\"  GPU {i} allocated memory : {alloc:.2f} GB\", log_path)\n            log_and_print(f\"  GPU {i} reserved memory  : {reserved:.2f} GB\", log_path)\n            log_and_print(f\"  GPU {i} peak memory      : {peak:.2f} GB\", log_path)\n\n\ndef evaluate_stage1(model, data_loader, cfg=CFG):\n    '''Collects presence/absence predictions over the full data_loader and computes\n    per-class Precision, Recall, F1, overall accuracy, and a confusion matrix for the\n    instance-presence head (binary: level visible / level not visible per disc level).\n    Runs entirely under torch.no_grad() to avoid accumulating gradients.\n    Returns: dict with keys \"report\" (sklearn classification_report string),\n             \"confusion_matrix\" (2x2 numpy array), \"accuracy\" (float).'''\n    from sklearn.metrics import classification_report, confusion_matrix, accuracy_score\n    model.eval()\n    all_preds, all_targets = [], []\n    with torch.no_grad():\n        for batch in data_loader:\n            vol = batch[\"volume\"].to(DEVICE, non_blocking=True)\n            presence = batch[\"presence\"].to(DEVICE, non_blocking=True)\n            out = model(vol)\n            preds = (torch.sigmoid(out[\"instance_logits\"]) > 0.5).long().view(-1).cpu()\n            targets = presence.long().view(-1).cpu()\n            all_preds.extend(preds.tolist())\n            all_targets.extend(targets.tolist())\n\n    report = classification_report(\n        all_targets, all_preds,\n        labels=[0, 1],\n        target_names=[\"Absent (level not visible)\", \"Present (level visible)\"],\n        zero_division=0,\n    )\n    cm = confusion_matrix(all_targets, all_preds, labels=[0, 1])\n    acc = accuracy_score(all_targets, all_preds)\n    return {\"report\": report, \"confusion_matrix\": cm, \"accuracy\": acc}\n\n\n    return history\n\n# ---- Self-test: stage1_loss / stage1_accuracy run and return finite values on fake data ----\ntry:\n    _fake_out = {\n        \"instance_logits\": torch.randn(2, CFG[\"num_levels\"]),\n        \"depth_pred\": torch.rand(2, CFG[\"num_levels\"]),\n        \"coord_pred\": torch.rand(2, CFG[\"num_levels\"], 2),\n    }\n    _fake_coords = torch.rand(2, CFG[\"num_levels\"], 3)\n    _fake_presence = torch.randint(0, 2, (2, CFG[\"num_levels\"])).float()\n    _loss, _parts = stage1_loss(_fake_out, _fake_coords, _fake_presence)\n    _acc = stage1_accuracy(_fake_out, _fake_presence)\n    assert torch.isfinite(_loss), \"loss is not finite\"\n    assert 0.0 <= _acc <= 1.0, f\"accuracy out of range: {_acc}\"\n    cell_status(\"CELL 7 — Stage 1 loss & training loop\",\n                extra=f\"stage1_loss OK loss={_loss.item():.4f}, stage1_accuracy OK acc={_acc:.2%}\")\n    del _fake_out, _fake_coords, _fake_presence, _loss, _parts, _acc\nexcept Exception as e:\n    cell_status(\"CELL 7 — Stage 1 loss & training loop\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:45:06.059847Z","iopub.execute_input":"2026-08-02T05:45:06.060356Z","iopub.status.idle":"2026-08-02T05:45:06.102939Z","shell.execute_reply.started":"2026-08-02T05:45:06.060328Z","shell.execute_reply":"2026-08-02T05:45:06.102333Z"}},"outputs":[],"execution_count":null},{"id":"398341c6-bdd5-4b3a-a680-084111229ff8","cell_type":"code","source":"# =====================================================================\n# CELL 8 — ROI extraction from Stage 1 predictions\n# =====================================================================\ndef extract_rois(volume, coord_pred, instance_logits, crop_size=(5, 96, 96), presence_thresh=0.5):\n    '''Crop a small (d, h, w) sub-volume around each predicted disc-level location.\n    volume: (D, H, W) numpy array (single series, already resampled)\n    coord_pred: (L, 3) normalized (depth, y, x) predictions\n    instance_logits: (L,) presence logits\n    Returns: dict[level_idx] -> (crop, present_bool)\n    '''\n    D, H, W = volume.shape\n    cd, ch, cw = crop_size\n    presence = torch.sigmoid(torch.as_tensor(instance_logits)).numpy()\n\n    rois = {}\n    for li in range(coord_pred.shape[0]):\n        present = presence[li] > presence_thresh\n        depth_n, y_n, x_n = coord_pred[li]\n        d = int(np.clip(depth_n * (D - 1), 0, D - 1))\n        y = int(np.clip(y_n * H, ch // 2, H - ch // 2))\n        x = int(np.clip(x_n * W, cw // 2, W - cw // 2))\n\n        d0, d1 = max(0, d - cd // 2), min(D, d + cd // 2 + 1)\n        crop = volume[d0:d1, y - ch // 2:y + ch // 2, x - cw // 2:x + cw // 2]\n        if crop.shape[0] < cd:\n            pad = np.zeros((cd - crop.shape[0], ch, cw), dtype=crop.dtype)\n            crop = np.concatenate([crop, pad], axis=0)\n        rois[li] = (crop.astype(np.float32), bool(present))\n    return rois\n\n# ---- Self-test: extract_rois on a synthetic volume + synthetic predictions --\ntry:\n    _fake_volume = np.random.rand(*CFG[\"vol_size_3d\"]).astype(np.float32)\n    _fake_coord_pred = np.random.rand(CFG[\"num_levels\"], 3).astype(np.float32)\n    _fake_inst_logits = torch.randn(CFG[\"num_levels\"])\n    _rois = extract_rois(_fake_volume, _fake_coord_pred, _fake_inst_logits)\n    assert len(_rois) == CFG[\"num_levels\"]\n    _present_count = sum(1 for _, present in _rois.values() if present)\n    cell_status(\"CELL 8 — ROI extraction\", extra=f\"extracted {len(_rois)} level slots, {_present_count} marked present\")\n    del _fake_volume, _fake_coord_pred, _fake_inst_logits, _rois\nexcept Exception as e:\n    cell_status(\"CELL 8 — ROI extraction\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:45:11.844472Z","iopub.execute_input":"2026-08-02T05:45:11.844923Z","iopub.status.idle":"2026-08-02T05:45:11.861806Z","shell.execute_reply.started":"2026-08-02T05:45:11.844896Z","shell.execute_reply":"2026-08-02T05:45:11.860792Z"}},"outputs":[],"execution_count":null},{"id":"7989f20c-7f78-4c51-b5ed-82ed46abca72","cell_type":"markdown","source":"## Stage 2 — Graph-based Severity Prediction\n\nEach of the 5 levels (from Stage 1 ROIs, across the 3 series types) is embedded with\n**EfficientNetV2-L**. We then build three complementary graphs over the level embeddings:\n\n- **Slice Graph** — connects adjacent slices *within* a level's ROI stack (local continuity)\n- **Level Graph** — connects adjacent vertebral levels (L1/L2 ↔ L2/L3 ↔ … ↔ L5/S1), since\n  degeneration at one level correlates with its neighbours\n- **Neuro Graph** — connects the 5 condition-specific nodes *within* a level (canal stenosis,\n  left/right foraminal narrowing, left/right subarticular stenosis share the same anatomy)\n\n**Cross-Graph Attention** lets information flow between the three graphs, and a\n**Hierarchical Graph Transformer** aggregates node embeddings (slice → level → whole-study) before\na final classification head predicts the 3-way severity for every condition × level.\n","metadata":{}},{"id":"e4e21c77-8d18-46e9-baa9-94d3c10b73ad","cell_type":"code","source":"# =====================================================================\n# CELL 9 — EfficientNetV2-L ROI embedding extractor\n# =====================================================================\nclass ROIEmbedder(nn.Module):\n    '''Encodes 2D ROI slices with EfficientNetV2-L pretrained on ImageNet.\n\n    The real cause of the Stage 2 CUDA OOM is here, not in batch_size_stage2: SeverityGraphModel\n    concatenates every ROI slice of every level of every study in a batch into ONE tensor of\n    shape (total_slices, 1, H, W) and used to run it through this backbone in a single forward\n    call. With batch_size_stage2=4, num_levels=5, max_slices_per_level=8, total_slices can reach\n    4*5*8=160 -- i.e. a single forward/backward pass through a ~120M-parameter EfficientNetV2-L\n    with an effective batch of 160 images, which is far larger than this model is normally\n    trained with even on multi-GPU setups. That single call is what allocates the ~14 GiB seen in\n    the OOM trace, not fragmentation and not the optimizer state.\n\n    Two orthogonal, architecture-preserving fixes are applied, both standard practice for exactly\n    this situation:\n    1. Gradient checkpointing on the backbone (timm's built-in `set_grad_checkpointing`) --\n       activations inside the backbone are recomputed during backward instead of stored, cutting\n       backbone activation memory dramatically at the cost of ~20-30% extra compute time. This\n       changes nothing about what is computed -- forward output and gradients are identical.\n    2. Chunked forward -- `x` is split into groups of `chunk_size` images and each group is run\n       through the (checkpointed) backbone sequentially, with results concatenated back together\n       before the projection head. Because this is plain Python looping over ordinary tensor\n       ops (no `.detach()`, no `torch.no_grad()` forced here), autograd builds one continuous\n       graph across all chunks exactly as if the whole tensor had been passed in at once -- the\n       forward output, the loss, and every gradient are numerically identical to the unchunked\n       version. Only peak memory changes: instead of holding activations for all 160 images at\n       once, at most `chunk_size` images are ever in flight through the backbone simultaneously.\n    Neither change alters the model's parameters, forward semantics, or the research\n    methodology -- both are purely about *when* memory is allocated/freed, not *what* is computed.\n\n    A THIRD fix is required alongside these two: EfficientNetV2-L's BatchNorm layers, left in\n    default train() mode, compute normalization statistics from whatever is in the CURRENT batch\n    -- so a chunk of 16 images and the full 160-image batch get numerically different statistics,\n    making chunked and unchunked forward passes diverge (this is not a chunking bug; it is a\n    property of BatchNorm itself, and would silently make the *unchunked* run's per-batch BN\n    statistics noisy and inconsistent across steps too, since total_slices varies study to study).\n    BatchNorm running statistics are frozen to the ImageNet-pretrained values (a standard,\n    well-established transfer-learning practice, not a research-methodology compromise) so the\n    embedder's output for a given image no longer depends on what else happens to be in its\n    chunk/batch -- this makes chunking exactly reproducible AND removes a source of training\n    noise that had nothing to do with the actual research design.\n    '''\n\n    def __init__(self, out_dim=CFG[\"graph_hidden\"], pretrained=True, backbone_name=\"tf_efficientnetv2_l\",\n                 grad_checkpointing=True, chunk_size=CFG[\"roi_embed_chunk_size\"]):\n        super().__init__()\n        self.backbone = timm.create_model(\n            backbone_name, pretrained=pretrained, in_chans=1, num_classes=0, global_pool=\"avg\"\n        )\n        if grad_checkpointing and hasattr(self.backbone, \"set_grad_checkpointing\"):\n            self.backbone.set_grad_checkpointing(enable=True)\n        self.chunk_size = chunk_size\n        feat_dim = self.backbone.num_features\n        self.proj = nn.Linear(feat_dim, out_dim)\n        self._freeze_batchnorm()\n\n    def _freeze_batchnorm(self):\n        \"\"\"Forces every BatchNorm layer in the backbone to always use its frozen running\n        mean/var (never per-batch statistics), regardless of train()/eval() mode. Affine\n        weight/bias parameters are untouched and remain trainable.\"\"\"\n        for m in self.backbone.modules():\n            if isinstance(m, (nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d)):\n                m.eval()\n\n    def train(self, mode=True):\n        super().train(mode)\n        self._freeze_batchnorm()  # re-freeze after PyTorch's recursive train() flips BN back on\n        return self\n\n    def forward(self, x):\n        # x: (N, 1, H, W) — a batch of 2D slices (already resized to CFG[\"img_size_2d\"])\n        if x.shape[0] <= self.chunk_size:\n            feat = self.backbone(x)\n        else:\n            feat = torch.cat([self.backbone(chunk) for chunk in x.split(self.chunk_size, dim=0)], dim=0)\n        return self.proj(feat)  # (N, out_dim)\n\n\ndef roi_to_slices(roi_crop, img_size=CFG[\"img_size_2d\"]):\n    '''Turn a (d, h, w) ROI crop into a stack of resized 2D slices ready for the embedder.'''\n    d = roi_crop.shape[0]\n    out = np.zeros((d, img_size, img_size), dtype=np.float32)\n    for i in range(d):\n        out[i] = cv2.resize(roi_crop[i], (img_size, img_size), interpolation=cv2.INTER_LINEAR)\n    return out  # (d, H, W)\n\n# ---- Self-test: forward pass with pretrained=False (no internet download needed\n# for this check -- the real training run should use pretrained=True), including a\n# batch LARGER than chunk_size to actually exercise the chunking path -----------\ntry:\n    _test_embedder = ROIEmbedder(pretrained=False, chunk_size=4)\n    _test_slices = torch.randn(3, 1, CFG[\"img_size_2d\"], CFG[\"img_size_2d\"])\n    with torch.no_grad():\n        _emb = _test_embedder(_test_slices)\n    assert _emb.shape == (3, CFG[\"graph_hidden\"]), _emb.shape\n\n    _test_slices_big = torch.randn(10, 1, CFG[\"img_size_2d\"], CFG[\"img_size_2d\"])\n    with torch.no_grad():\n        _emb_big = _test_embedder(_test_slices_big)\n    assert _emb_big.shape == (10, CFG[\"graph_hidden\"]), _emb_big.shape\n    _test_embedder.train()  # confirm BN stays frozen even after an explicit train() call\n    assert all(m.training == False for m in _test_embedder.backbone.modules()\n               if isinstance(m, (nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d))), \"BatchNorm did not stay frozen after train()\"\n    _emb_unchunked = _test_embedder.proj(_test_embedder.backbone(_test_slices_big))\n    assert torch.allclose(_emb_big, _emb_unchunked, atol=1e-5), \"chunked forward diverged from unchunked forward\"\n\n    _fake_roi = np.random.rand(5, 40, 40).astype(np.float32)\n    _resized = roi_to_slices(_fake_roi)\n    assert _resized.shape == (5, CFG[\"img_size_2d\"], CFG[\"img_size_2d\"]), _resized.shape\n    cell_status(\"CELL 9 — ROIEmbedder\", extra=f\"embed shape {tuple(_emb.shape)}, chunked forward matches unchunked, roi_to_slices OK\")\n    del _test_embedder, _test_slices, _emb, _test_slices_big, _emb_big, _emb_unchunked, _fake_roi, _resized\nexcept Exception as e:\n    cell_status(\"CELL 9 — ROIEmbedder\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:45:16.276375Z","iopub.execute_input":"2026-08-02T05:45:16.277092Z","iopub.status.idle":"2026-08-02T05:45:24.482743Z","shell.execute_reply.started":"2026-08-02T05:45:16.277062Z","shell.execute_reply":"2026-08-02T05:45:24.482054Z"}},"outputs":[],"execution_count":null},{"id":"a3a04d6d-ab5c-44de-871e-cd4f00e3f1f6","cell_type":"code","source":"# =====================================================================\n# CELL 10 — Graph construction: Slice Graph, Level Graph, Neuro Graph\n# =====================================================================\nNUM_LEVELS = CFG[\"num_levels\"]\nNUM_CONDITIONS = len(CFG[\"conditions\"])\n\ndef build_slice_edge_index(num_slices_per_node):\n    '''Chain graph connecting adjacent slices within one node's ROI stack (bidirectional).'''\n    src, dst = [], []\n    for i in range(num_slices_per_node - 1):\n        src += [i, i + 1]\n        dst += [i + 1, i]\n    return torch.tensor([src, dst], dtype=torch.long) if src else torch.zeros((2, 0), dtype=torch.long)\n\n\ndef build_level_edge_index(num_levels=NUM_LEVELS):\n    '''Chain graph over the 5 vertebral levels (L1/L2 - L2/L3 - ... - L5/S1), since adjacent\n    levels are anatomically coupled (a bulging disc often affects the neighbouring level too).'''\n    src, dst = [], []\n    for i in range(num_levels - 1):\n        src += [i, i + 1]\n        dst += [i + 1, i]\n    return torch.tensor([src, dst], dtype=torch.long)\n\n\ndef build_neuro_edge_index(num_conditions=NUM_CONDITIONS):\n    '''Fully-connected graph over the 5 condition nodes belonging to the same level — canal\n    stenosis and the four foraminal/subarticular conditions all reflect the same cross-section\n    of anatomy, so we let every condition attend to every other condition at that level.'''\n    src, dst = [], []\n    for i in range(num_conditions):\n        for j in range(num_conditions):\n            if i != j:\n                src.append(i); dst.append(j)\n    return torch.tensor([src, dst], dtype=torch.long)\n\n\nLEVEL_EDGE_INDEX = build_level_edge_index()\nNEURO_EDGE_INDEX = build_neuro_edge_index()\n\n# ---- Self-test: edge_index shapes for all three graph types -----------------\ntry:\n    _slice_ei = build_slice_edge_index(6)\n    assert _slice_ei.shape == (2, 10), _slice_ei.shape  # 5 slices -> 10 directed edges\n    assert LEVEL_EDGE_INDEX.shape == (2, 2 * (CFG[\"num_levels\"] - 1)), LEVEL_EDGE_INDEX.shape\n    assert NEURO_EDGE_INDEX.shape == (2, len(CFG[\"conditions\"]) * (len(CFG[\"conditions\"]) - 1)), NEURO_EDGE_INDEX.shape\n    cell_status(\"CELL 10 — Graph construction\", extra=f\"slice/level/neuro edge_index shapes all correct\")\n    del _slice_ei\nexcept Exception as e:\n    cell_status(\"CELL 10 — Graph construction\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n\n# ---- Self-test: edge_index shapes for all three graph types -----------------\ntry:\n    _slice_ei = build_slice_edge_index(6)\n    assert _slice_ei.shape == (2, 10), _slice_ei.shape  # 5 slices -> 10 directed edges\n    assert LEVEL_EDGE_INDEX.shape == (2, 2 * (CFG[\"num_levels\"] - 1)), LEVEL_EDGE_INDEX.shape\n    assert NEURO_EDGE_INDEX.shape == (2, len(CFG[\"conditions\"]) * (len(CFG[\"conditions\"]) - 1)), NEURO_EDGE_INDEX.shape\n    cell_status(\"CELL 10 — Graph construction\", extra=f\"slice/level/neuro edge_index shapes all correct\")\n    del _slice_ei\nexcept Exception as e:\n    cell_status(\"CELL 10 — Graph construction\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:45:30.143528Z","iopub.execute_input":"2026-08-02T05:45:30.143971Z","iopub.status.idle":"2026-08-02T05:45:30.155495Z","shell.execute_reply.started":"2026-08-02T05:45:30.143944Z","shell.execute_reply":"2026-08-02T05:45:30.154743Z"}},"outputs":[],"execution_count":null},{"id":"a3db3520-1aa8-4d07-8516-f401350ab47a","cell_type":"code","source":"# =====================================================================\n# CELL 11 — Graph Encoders (Slice / Level / Neuro)\n# CrossGraphAttention and HierarchicalGraphTransformer have been\n# removed. Each graph is encoded independently with a small GATv2\n# stack; outputs are mean-pooled into a per-level, per-condition\n# embedding and then fed straight to the classification heads.\n# This keeps the three anatomically-motivated graph structures\n# (slice continuity, level adjacency, condition co-occurrence)\n# while drastically reducing model complexity and memory usage.\n# =====================================================================\nclass GraphEncoderGATv2(nn.Module):\n    '''Two-layer GATv2 encoder shared across all three graph types.\n    Returns refined node embeddings + the final-layer attention\n    weights (used by Stage 3 XAI for edge/node importance).'''\n\n    def __init__(self, dim=CFG[\"graph_hidden\"], heads=CFG[\"graph_heads\"], layers=2):\n        super().__init__()\n        self.convs = nn.ModuleList([\n            GATv2Conv(dim, dim // heads, heads=heads, concat=True, add_self_loops=True)\n            for _ in range(layers)\n        ])\n        self.norms = nn.ModuleList([nn.LayerNorm(dim) for _ in range(layers)])\n\n    def forward(self, x, edge_index):\n        last_attn = None\n        for conv, norm in zip(self.convs, self.norms):\n            out, (ei_used, alpha) = conv(x, edge_index, return_attention_weights=True)\n            x = norm(x + F.gelu(out))\n            last_attn = (ei_used, alpha)\n        return x, last_attn   # node embeddings, (edge_index, attention weights)\n\n\n# ---- Self-test -----------------------------------------------------------\ntry:\n    _dim, _n = CFG[\"graph_hidden\"], 5\n    _enc = GraphEncoderGATv2(dim=_dim)\n    _x   = torch.randn(_n, _dim)\n    _ei  = torch.tensor([[0,1,2,3],[1,2,3,4]], dtype=torch.long)\n    with torch.no_grad():\n        _out, _attn = _enc(_x, _ei)\n    assert _out.shape == (_n, _dim)\n    cell_status(\"CELL 11 — GraphEncoderGATv2\",\n                extra=f\"forward OK shape={tuple(_out.shape)}\")\n    del _dim, _n, _enc, _x, _ei, _out, _attn\nexcept Exception as e:\n    cell_status(\"CELL 11 — GraphEncoderGATv2\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:45:36.459253Z","iopub.execute_input":"2026-08-02T05:45:36.459693Z","iopub.status.idle":"2026-08-02T05:45:36.555552Z","shell.execute_reply.started":"2026-08-02T05:45:36.459665Z","shell.execute_reply":"2026-08-02T05:45:36.554942Z"}},"outputs":[],"execution_count":null},{"id":"56d8d327-248d-4d70-ab65-58dc1af4ff77","cell_type":"code","source":"# =====================================================================\n# CELL 12 — SeverityGraphModel (simplified)\n# Architecture:\n#   ROI slices → EfficientNetV2-L embedder\n#   → Slice Graph  (GATv2, per-level slice chain)\n#   → Level Graph  (GATv2, L1/L2 ↔ L2/L3 ↔ … ↔ L5/S1)\n#   → Neuro Graph  (GATv2, fully-connected within each level)\n#   → mean-pool readout\n#   → per-condition classification heads\n# CrossGraphAttention and HierarchicalGraphTransformer removed.\n# =====================================================================\nclass SeverityGraphModel(nn.Module):\n    def __init__(self, dim=CFG[\"graph_hidden\"], num_classes=CFG[\"num_classes\"],\n                 pretrained=True, backbone_name=\"tf_efficientnetv2_l\",\n                 embed_chunk_size=CFG[\"roi_embed_chunk_size\"]):\n        super().__init__()\n        self.embedder = ROIEmbedder(out_dim=dim, pretrained=pretrained,\n                                     backbone_name=backbone_name,\n                                     chunk_size=embed_chunk_size)\n        # Learnable per-(level, condition) positional offset — same purpose as\n        # the old level_condition_embed but survives without the transformer.\n        self.level_condition_embed = nn.Parameter(\n            torch.randn(NUM_LEVELS, NUM_CONDITIONS, dim) * 0.02)\n\n        # Three independent GATv2 encoders, one per graph type.\n        self.slice_encoder = GraphEncoderGATv2(dim)\n        self.level_encoder = GraphEncoderGATv2(dim)\n        self.neuro_encoder = GraphEncoderGATv2(dim)\n\n        # Per-condition classification heads (not shared across conditions).\n        self.cls_heads = nn.ModuleList(\n            [nn.Linear(dim, num_classes) for _ in range(NUM_CONDITIONS)])\n\n    def forward(self, slice_stack, slice_counts):\n        '''\n        slice_stack  : (total_slices, 1, H, W)\n        slice_counts : list[int]  number of slices per (study x level) node\n        Returns dict with keys: logits, attn_maps, node_embeds\n        '''\n        # 1) Embed every ROI slice through EfficientNetV2-L (chunked)\n        slice_embeds = self.embedder(slice_stack)   # (total_slices, dim)\n        device = slice_embeds.device\n\n        # 2) Build slice_batch and slice edge_index (chain within each level-node)\n        slice_batch, e_src, e_dst, offset = [], [], [], 0\n        for node_id, n in enumerate(slice_counts):\n            slice_batch += [node_id] * n\n            for i in range(n - 1):\n                e_src += [offset + i, offset + i + 1]\n                e_dst += [offset + i + 1, offset + i]\n            offset += n\n        slice_batch_t = torch.tensor(slice_batch, dtype=torch.long, device=device)\n        slice_ei = (torch.tensor([e_src, e_dst], dtype=torch.long, device=device)\n                    if e_src else torch.zeros((2, 0), dtype=torch.long, device=device))\n\n        # 3) Slice Graph: refine per-slice embeddings within each level-node\n        slice_refined, slice_attn = self.slice_encoder(slice_embeds, slice_ei)\n\n        # 4) Pool slices → one embedding per (study, level) node\n        B = len(slice_counts) // NUM_LEVELS\n        level_embeds = global_mean_pool(slice_refined, slice_batch_t).view(B, NUM_LEVELS, -1)\n\n        # 5) Level Graph: propagate across the 5 vertebral levels\n        level_flat = level_embeds.reshape(B * NUM_LEVELS, -1)\n        level_ei   = torch.cat(\n            [LEVEL_EDGE_INDEX + i * NUM_LEVELS for i in range(B)], dim=1\n        ).to(device)\n        level_refined, level_attn = self.level_encoder(level_flat, level_ei)\n        level_refined = level_refined.view(B, NUM_LEVELS, -1)\n\n        # 6) Neuro Graph: propagate across the 5 conditions within each level\n        #    Initial node features = level embedding + learnable position offset\n        neuro_init = (level_refined.unsqueeze(2)\n                      + self.level_condition_embed.unsqueeze(0))   # (B, L, C, dim)\n        neuro_flat = neuro_init.reshape(B * NUM_LEVELS * NUM_CONDITIONS, -1)\n        neuro_ei   = torch.cat(\n            [NEURO_EDGE_INDEX + i * NUM_CONDITIONS for i in range(B * NUM_LEVELS)], dim=1\n        ).to(device)\n        neuro_refined, neuro_attn = self.neuro_encoder(neuro_flat, neuro_ei)\n        neuro_refined = neuro_refined.view(B, NUM_LEVELS, NUM_CONDITIONS, -1)  # (B,L,C,dim)\n\n        # 7) Classification: one head per condition\n        logits = torch.zeros(B, NUM_LEVELS, NUM_CONDITIONS,\n                              len(self.cls_heads[0].bias), device=device)\n        for c in range(NUM_CONDITIONS):\n            logits[:, :, c, :] = self.cls_heads[c](neuro_refined[:, :, c, :])\n\n        attn_maps = {\n            \"slice_attn\": slice_attn,\n            \"level_attn\": level_attn,\n            \"neuro_attn\": neuro_attn,\n        }\n        return {\"logits\": logits, \"attn_maps\": attn_maps,\n                \"node_embeds\": neuro_refined}\n\n\n# ---- Self-test (pretrained=False → no download) ----\ntry:\n    _m2 = SeverityGraphModel(pretrained=False)\n    _sc = [2] * NUM_LEVELS           # 1 study, 2 slices per level\n    _ss = torch.randn(sum(_sc), 1, CFG[\"img_size_2d\"], CFG[\"img_size_2d\"])\n    with torch.no_grad():\n        _o2 = _m2(_ss, _sc)\n    assert _o2[\"logits\"].shape == (1, NUM_LEVELS, NUM_CONDITIONS, CFG[\"num_classes\"])\n    _np = sum(p.numel() for p in _m2.parameters())\n    cell_status(\"CELL 12 — SeverityGraphModel (simplified)\",\n                extra=f\"logits {tuple(_o2['logits'].shape)}, {_np:,} params\")\n    del _m2, _sc, _ss, _o2, _np\nexcept Exception as e:\n    cell_status(\"CELL 12 — SeverityGraphModel (simplified)\", ok=False,\n                extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:45:41.099070Z","iopub.execute_input":"2026-08-02T05:45:41.099731Z","iopub.status.idle":"2026-08-02T05:45:44.675039Z","shell.execute_reply.started":"2026-08-02T05:45:41.099701Z","shell.execute_reply":"2026-08-02T05:45:44.674401Z"}},"outputs":[],"execution_count":null},{"id":"6334f6fa-d7da-48ac-8190-4f5b62d35fe9","cell_type":"code","source":"# =====================================================================\n# CELL 13 — Stage 2 loss (competition uses sample-weighted log loss:\n# Normal/Mild=1, Moderate=2, Severe=4) and training loop\n# =====================================================================\nSEVERITY_SAMPLE_WEIGHT = torch.tensor([1.0, 2.0, 4.0])  # official RSNA metric weighting\n\ndef stage2_loss(logits, targets, valid_mask):\n    '''logits: (B, L, C, num_classes); targets: (B, L, C) int64 in {0,1,2}; valid_mask: (B, L, C) bool\n    (some condition/level combinations are missing for a given study and must be excluded).'''\n    B, L, C, K = logits.shape\n    logits_flat = logits.view(-1, K)\n    targets_flat = targets.view(-1)\n    weights = SEVERITY_SAMPLE_WEIGHT.to(logits.device)[targets_flat.clamp(min=0)]\n    mask_flat = valid_mask.view(-1).float()\n\n    loss_per_sample = F.cross_entropy(logits_flat, targets_flat.clamp(min=0), reduction=\"none\")\n    loss = (loss_per_sample * weights * mask_flat).sum() / (weights * mask_flat).sum().clamp(min=1)\n    return loss\n\n\ndef stage2_accuracy(logits, targets, valid_mask):\n    '''Argmax classification accuracy over every valid (level, condition) cell.'''\n    with torch.no_grad():\n        preds = logits.argmax(dim=-1)\n        correct = (preds == targets) & valid_mask\n        return correct.float().sum().item() / valid_mask.float().sum().clamp(min=1).item()\n\n\nclass SeverityDataParallel(nn.DataParallel):\n    '''nn.DataParallel by default splits every positional TENSOR argument along dim 0 and passes\n    non-tensor arguments (like a Python list) UNCHANGED to every replica. SeverityGraphModel's\n    forward(slice_stack, slice_counts) takes exactly one of each -- naive DataParallel would split\n    slice_stack across GPUs but hand each replica the FULL, un-split slice_counts list, breaking\n    the slice-to-level grouping and silently producing wrong results. This subclass splits\n    slice_stack and slice_counts together, in lock-step, by STUDY (not by raw row count, since a\n    study's row count varies with how many ROI slices it contributed), so each GPU replica\n    receives a self-consistent (slice_stack_chunk, slice_counts_chunk) pair -- this is what makes\n    multi-GPU training numerically correct here, not just \"technically using 2 GPUs\".'''\n\n    def scatter(self, inputs, kwargs, device_ids):\n        slice_stack, slice_counts = inputs\n        num_levels = CFG[\"num_levels\"]\n        B = len(slice_counts) // num_levels\n        n_gpus = len(device_ids)\n        studies_per_gpu = [B // n_gpus + (1 if i < B % n_gpus else 0) for i in range(n_gpus)]\n\n        scattered_inputs, scattered_kwargs = [], []\n        study_offset, row_offset = 0, 0\n        for gpu_idx, n_studies in enumerate(studies_per_gpu):\n            target_device = f\"cuda:{device_ids[gpu_idx]}\"\n            if n_studies == 0:\n                empty = torch.empty(0, *slice_stack.shape[1:], device=target_device, dtype=slice_stack.dtype)\n                scattered_inputs.append(((empty, []), {}))\n                continue\n            group_start = study_offset * num_levels\n            group_end = (study_offset + n_studies) * num_levels\n            counts_chunk = slice_counts[group_start:group_end]\n            n_rows = sum(counts_chunk)\n            stack_chunk = slice_stack[row_offset:row_offset + n_rows].to(target_device, non_blocking=True)\n            scattered_inputs.append(((stack_chunk, counts_chunk), {}))\n            study_offset += n_studies\n            row_offset += n_rows\n        return [i for i, _ in scattered_inputs], [k for _, k in scattered_inputs]\n\n    def gather(self, outputs, output_device):\n        # Training only needs \"logits\" for the loss -- attn_maps/node_embeds contain graph-\n        # structure tensors (e.g. edge_index) that are local to each replica's sub-batch and\n        # cannot be meaningfully concatenated across GPUs, so they are intentionally dropped\n        # here. Inference (Stage 3/4 XAI, run_full_pipeline) always runs on the single unwrapped\n        # model, never through this DataParallel path, so nothing downstream loses access to them.\n        logits = torch.cat([o[\"logits\"].to(output_device) for o in outputs if o[\"logits\"].shape[0] > 0], dim=0)\n        return {\"logits\": logits}\n\n\ndef train_stage2(model, train_loader, val_loader, cfg=CFG, epochs=None, grad_clip=1.0):\n    epochs = epochs or cfg[\"epochs_stage2\"]\n    model.to(DEVICE)\n    assert_on_gpu(model, \"Stage2 model\")\n    print_gpu_monitoring_report()\n    if N_GPUS > 1:\n        model = SeverityDataParallel(model)\n    print(f\"[Stage2] Training on device: {next(model.parameters()).device}\"\n          f\"{' (custom-scatter DataParallel across ' + str(N_GPUS) + ' GPUs)' if isinstance(model, nn.DataParallel) else ''}\")\n    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg[\"lr_stage2\"], weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(\n        optimizer, max_lr=cfg[\"lr_stage2\"], epochs=epochs, steps_per_epoch=max(1, len(train_loader))\n    )\n    scaler = GradScaler()\n\n    log_path = os.path.join(cfg[\"work_dir\"], \"stage2_training_log.txt\")\n    with open(log_path, \"a\") as f:\n        f.write(f\"\\n===== New Stage 2 run started, {epochs} epochs =====\\n\")\n    print(f\"[Stage2] Logging every epoch to: {log_path}\")\n\n    history = {\"train_loss\": [], \"val_loss\": [], \"train_acc\": [], \"val_acc\": []}\n    best_epoch, best_val_loss = -1, float(\"inf\")\n    for epoch in range(epochs):\n        epoch_start = time.time()\n        model.train()\n        running, running_acc = 0.0, 0.0\n        n_valid_batches = 0\n        n_skipped = 0\n        for batch_idx, batch in enumerate(train_loader):\n            slice_stack = batch[\"slice_stack\"].to(DEVICE, non_blocking=True)\n            targets = batch[\"targets\"].to(DEVICE, non_blocking=True)\n            valid_mask = batch[\"valid_mask\"].to(DEVICE, non_blocking=True)\n\n            _debug = batch_idx <= 2  # covers the exact window where the silent crash has happened\n            if _debug:\n                log_and_print(f\"  [DEBUG] Epoch {epoch+1} batch {batch_idx}: host RAM={host_ram_gb():.2f} GB \"\n                              f\"(after DataLoader fetch), slice_stack.shape={tuple(slice_stack.shape)}, \"\n                              f\"len(slice_counts)={len(batch['slice_counts'])}, \"\n                              f\"sum(slice_counts)={sum(batch['slice_counts'])}\", log_path)\n\n            try:\n                optimizer.zero_grad(set_to_none=True)\n                with autocast():\n                    out = model(slice_stack, batch[\"slice_counts\"])\n                    loss = stage2_loss(out[\"logits\"], targets, valid_mask)\n                if _debug:\n                    log_and_print(f\"  [DEBUG] Epoch {epoch+1} batch {batch_idx}: host RAM={host_ram_gb():.2f} GB \"\n                                  f\"(after forward+loss, loss={loss.item():.4f})\", log_path)\n\n                if not torch.isfinite(loss):\n                    n_skipped += 1\n                    log_and_print(f\"  [WARN] Stage2 epoch {epoch+1} batch {batch_idx}: non-finite loss — skipping this batch\", log_path)\n                    continue\n\n                scaler.scale(loss).backward()\n                if _debug:\n                    log_and_print(f\"  [DEBUG] Epoch {epoch+1} batch {batch_idx}: host RAM={host_ram_gb():.2f} GB (after backward)\", log_path)\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip)\n                scaler.step(optimizer)\n                scaler.update()\n                scheduler.step()\n                if _debug:\n                    log_and_print(f\"  [DEBUG] Epoch {epoch+1} batch {batch_idx}: host RAM={host_ram_gb():.2f} GB (after optimizer.step)\", log_path)\n                running += loss.item()\n                running_acc += stage2_accuracy(out[\"logits\"], targets, valid_mask)\n                n_valid_batches += 1\n\n                # Explicitly release the large tensors from this batch right here -- the Python\n                # garbage collector is not guaranteed to run between batches, and with\n                # SeverityDataset's disk-cache-building phase still running in DataLoader worker\n                # processes, host RAM grows continuously until GC gets a chance. Explicit del\n                # gives it that chance every batch instead of waiting until it can't allocate.\n                del out, loss\n                if n_valid_batches % 10 == 0:\n                    import gc as _gc; _gc.collect()\n\n            except torch.cuda.OutOfMemoryError as e:\n                # Free whatever we can and skip this batch rather than crashing the whole run --\n                # if this happens often, lower CFG[\"batch_size_stage2\"], cap SeverityDataset's\n                # max_slices_per_level, or switch roi_embedder_backbone to a smaller model.\n                optimizer.zero_grad(set_to_none=True)\n                torch.cuda.empty_cache()\n                n_skipped += 1\n                log_and_print(f\"  [WARN] Stage2 epoch {epoch+1} batch {batch_idx}: CUDA OOM ({e}) — \"\n                              f\"skipping this batch and freeing cache. Consider a smaller batch_size_stage2, \"\n                              f\"a lower max_slices_per_level, or a smaller roi_embedder_backbone if this repeats.\", log_path)\n                continue\n\n        if n_skipped > 0:\n            log_and_print(f\"  Stage2 epoch {epoch+1}: skipped {n_skipped}/{len(train_loader)} batches (non-finite loss or OOM)\", log_path)\n        train_loss = running / max(1, n_valid_batches)\n        train_acc = running_acc / max(1, n_valid_batches)\n\n        model.eval()\n        val_running, val_running_acc = 0.0, 0.0\n        n_valid_val = 0\n        with torch.no_grad():\n            for batch in val_loader:\n                try:\n                    slice_stack = batch[\"slice_stack\"].to(DEVICE, non_blocking=True)\n                    targets = batch[\"targets\"].to(DEVICE, non_blocking=True)\n                    valid_mask = batch[\"valid_mask\"].to(DEVICE, non_blocking=True)\n                    out = model(slice_stack, batch[\"slice_counts\"])\n                    loss = stage2_loss(out[\"logits\"], targets, valid_mask)\n                    if torch.isfinite(loss):\n                        val_running += loss.item()\n                        val_running_acc += stage2_accuracy(out[\"logits\"], targets, valid_mask)\n                        n_valid_val += 1\n                except torch.cuda.OutOfMemoryError:\n                    torch.cuda.empty_cache()\n                    continue\n        val_loss = val_running / max(1, n_valid_val)\n        val_acc = val_running_acc / max(1, n_valid_val)\n\n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss)\n        history[\"train_acc\"].append(train_acc)\n        history[\"val_acc\"].append(val_acc)\n        elapsed = time.time() - epoch_start\n\n        _ram_now = host_ram_gb()\n        if _ram_now >= 28:\n            log_and_print(f\"  [WARNING] Host RAM = {_ram_now:.1f} GB  >=28 GB -- risk of OOM\", log_path)\n        elif _ram_now >= 25:\n            log_and_print(f\"  [NOTICE] Host RAM = {_ram_now:.1f} GB (approaching limit)\", log_path)\n        log_and_print(f\"[Stage2][Epoch {epoch+1}/{epochs}]\", log_path)\n        log_and_print(f\"  Train Loss     : {train_loss:.4f}\", log_path)\n        log_and_print(f\"  Train Accuracy : {train_acc:.2%}\", log_path)\n        log_and_print(f\"  Val Loss       : {val_loss:.4f}\", log_path)\n        log_and_print(f\"  Val Accuracy   : {val_acc:.2%}\", log_path)\n        log_and_print(f\"  Elapsed Time   : {elapsed:.1f}s\", log_path)\n\n        plain_model = unwrap_model(model)\n        torch.save(plain_model.state_dict(), os.path.join(cfg[\"work_dir\"], \"stage2_last.pt\"))\n        save_status = \"last checkpoint saved\"\n        if val_loss < best_val_loss:\n            best_val_loss, best_epoch = val_loss, epoch + 1\n            torch.save(plain_model.state_dict(), os.path.join(cfg[\"work_dir\"], \"stage2_best.pt\"))\n            save_status = \"last + BEST checkpoint saved\"\n        log_and_print(f\"  Host RAM          : {host_ram_gb():.1f} GB\", log_path)\n        log_and_print(f\"  Best Epoch So Far : {best_epoch} (val_loss={best_val_loss:.4f})\", log_path)\n        log_and_print(f\"  Model Save Status : {save_status}\", log_path)\n        for i in range(N_GPUS):\n            alloc = torch.cuda.memory_allocated(i) / 1e9\n            reserved = torch.cuda.memory_reserved(i) / 1e9\n            peak = torch.cuda.max_memory_allocated(i) / 1e9\n            log_and_print(f\"  GPU {i} allocated memory : {alloc:.2f} GB\", log_path)\n            log_and_print(f\"  GPU {i} reserved memory  : {reserved:.2f} GB\", log_path)\n            log_and_print(f\"  GPU {i} peak memory      : {peak:.2f} GB\", log_path)\n\n    return history\n\n# ---- Self-test: stage2_loss / stage2_accuracy run and return sane values on fake data ----\ntry:\n    _fake_logits = torch.randn(2, CFG[\"num_levels\"], len(CFG[\"conditions\"]), CFG[\"num_classes\"])\n    _fake_targets = torch.randint(0, CFG[\"num_classes\"], (2, CFG[\"num_levels\"], len(CFG[\"conditions\"])))\n    _fake_valid = torch.ones(2, CFG[\"num_levels\"], len(CFG[\"conditions\"]), dtype=torch.bool)\n    _loss2 = stage2_loss(_fake_logits, _fake_targets, _fake_valid)\n    _acc2 = stage2_accuracy(_fake_logits, _fake_targets, _fake_valid)\n    assert torch.isfinite(_loss2), \"loss is not finite\"\n    assert 0.0 <= _acc2 <= 1.0, f\"accuracy out of range: {_acc2}\"\n    cell_status(\"CELL 13 — Stage 2 loss & training loop\",\n                extra=f\"stage2_loss OK loss={_loss2.item():.4f}, stage2_accuracy OK acc={_acc2:.2%}\")\n    del _fake_logits, _fake_targets, _fake_valid, _loss2, _acc2\nexcept Exception as e:\n    cell_status(\"CELL 13 — Stage 2 loss & training loop\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:45:49.587684Z","iopub.execute_input":"2026-08-02T05:45:49.588383Z","iopub.status.idle":"2026-08-02T05:45:49.626393Z","shell.execute_reply.started":"2026-08-02T05:45:49.588356Z","shell.execute_reply":"2026-08-02T05:45:49.625500Z"}},"outputs":[],"execution_count":null},{"id":"31e3d6e1-4578-4f2a-aefc-2020ba3b543d","cell_type":"code","source":"# =====================================================================\n# CELL 13b — Stage 2 dataset: ROI crops from ground-truth coordinates\n# =====================================================================\nclass SeverityDataset(Dataset):\n    '''Builds one training sample per study: for every of the 5 levels, crops a small ROI stack\n    around the ground-truth (x, y, instance_number) from train_label_coordinates.csv, across every\n    series that has a labelled point for that level. Targets come from train.csv\n    (\"<condition>_<level>\" columns). Training Stage 2 this way (ground-truth ROIs) is standard\n    practice: it lets Stage 2 converge independently of Stage 1's localization accuracy; at\n    inference time the same model is fed Stage 1's *predicted* ROIs instead (see run_full_pipeline).\n    '''\n\n    def __init__(self, df_train, df_coords, data_root=CFG[\"data_root\"], split=\"train\", crop=(5, 96, 96),\n                 max_slices_per_level=8):\n        self.data_root = data_root\n        self.split = split\n        self.crop = crop\n        # A level can have a labelled coordinate on MANY axial slices (axial series often has one\n        # point per level per slice across a whole stack), so the ROI count per level is\n        # unbounded unless capped -- left unbounded this is the single biggest driver of the\n        # CUDA OOM seen in Stage 2 training, since slice_stack size scales directly with it.\n        self.max_slices_per_level = max_slices_per_level\n        self.df_train = df_train.set_index(\"study_id\")\n        self.df_coords = df_coords\n        self.study_ids = list(self.df_train.index)\n        self.level_to_idx = {lvl: i for i, lvl in enumerate(CFG[\"levels\"])}\n        # ON-DISK cache only (.npz files under CFG[\"temp_dir\"]) for the RAW extracted ROI crops\n        # per study -- i.e. everything expensive (DICOM decode via pydicom, per-slice cv2\n        # resize). The random max_slices_per_level SUBSAMPLING still happens fresh on every\n        # __getitem__ call (after the cache lookup), so training keeps its intended per-epoch\n        # stochasticity; only the expensive I/O + resize is cached, not the sampled example.\n        # Deliberately NOT also kept in an in-memory Python dict: doing so previously caused an\n        # OOM (Linux OOM-killer SIGKILL'ing a DataLoader worker) once several workers each\n        # accumulated their own unbounded, un-evictable copy in RAM. The on-disk cache alone\n        # already gives the real fix (shared across every worker process), and the Linux page\n        # cache transparently keeps recently-read .npz files warm in RAM with proper eviction\n        # under memory pressure, which our own dict could never do.\n        self._disk_cache_dir = os.path.join(CFG[\"temp_dir\"], \"severity_roi_cache\")\n        os.makedirs(self._disk_cache_dir, exist_ok=True)\n\n    def _disk_cache_path(self, study_id):\n        return os.path.join(self._disk_cache_dir, f\"{study_id}.npz\")\n\n    def __len__(self):\n        return len(self.study_ids)\n\n    def _crop_roi(self, volume, orig_hw, instance_numbers, instance_number, x, y):\n        cd, ch, cw = self.crop\n        orig_h, orig_w = orig_hw\n        if instance_number not in instance_numbers:\n            return None\n        idx = instance_numbers.index(instance_number)\n        # Coordinates are in ORIGINAL pixel space -> rescale into the resized volume's H, W\n        y_r = float(y) / orig_h * volume.shape[1]\n        x_r = float(x) / orig_w * volume.shape[2]\n        y_i = int(np.clip(y_r, ch // 2, max(ch // 2, volume.shape[1] - ch // 2)))\n        x_i = int(np.clip(x_r, cw // 2, max(cw // 2, volume.shape[2] - cw // 2)))\n        d0, d1 = max(0, idx - cd // 2), min(volume.shape[0], idx + cd // 2 + 1)\n        crop = volume[d0:d1, y_i - ch // 2:y_i + ch // 2, x_i - cw // 2:x_i + cw // 2]\n        if crop.shape[0] < cd:\n            pad = np.zeros((cd - crop.shape[0], ch, cw), dtype=crop.dtype)\n            crop = np.concatenate([crop, pad], axis=0)\n        return crop.astype(np.float32)\n\n    def __getitem__(self, i):\n        study_id = self.study_ids[i]\n        row = self.df_train.loc[study_id]\n\n        targets = np.zeros((CFG[\"num_levels\"], len(CFG[\"conditions\"])), dtype=np.int64)\n        valid_mask = np.zeros((CFG[\"num_levels\"], len(CFG[\"conditions\"])), dtype=bool)\n        for li, level in enumerate(CFG[\"levels\"]):\n            for ci, cond in enumerate(CFG[\"conditions\"]):\n                col = f\"{cond}_{level}\"\n                if col in row.index and pd.notna(row[col]):\n                    targets[li, ci] = severity_to_label(row[col])\n                    valid_mask[li, ci] = True\n\n        disk_path = self._disk_cache_path(study_id)\n        if os.path.exists(disk_path):\n            with np.load(disk_path) as npz:\n                all_rois, counts = npz[\"all_rois\"], npz[\"counts\"].tolist()\n            rois_per_level, offset = {}, 0\n            for li, c in enumerate(counts):\n                rois_per_level[li] = [all_rois[offset + k] for k in range(c)]\n                offset += c\n        else:\n            study_coords = self.df_coords[self.df_coords[\"study_id\"] == study_id]\n            rois_per_level = {li: [] for li in range(CFG[\"num_levels\"])}\n\n            for series_id, grp in study_coords.groupby(\"series_id\"):\n                volume, instance_numbers, orig_hw = load_dicom_series(study_id, series_id, self.data_root, self.split)\n                if volume is None:\n                    continue\n                for _, r in grp.iterrows():\n                    lvl = str(r[\"level\"]).lower().replace(\"/\", \"_\")\n                    if lvl not in self.level_to_idx:\n                        continue\n                    li = self.level_to_idx[lvl]\n                    crop = self._crop_roi(volume, orig_hw, instance_numbers, int(r[\"instance_number\"]), r[\"x\"], r[\"y\"])\n                    if crop is not None:\n                        rois_per_level[li].append(roi_to_slices(crop, img_size=CFG[\"img_size_2d\"]))\n\n            counts = [len(rois_per_level[li]) for li in range(CFG[\"num_levels\"])]\n            _empty_part = np.zeros((0, self.crop[0], CFG[\"img_size_2d\"], CFG[\"img_size_2d\"]), dtype=np.float32)\n            _parts = [np.stack(rois_per_level[li], axis=0) if rois_per_level[li] else _empty_part\n                      for li in range(CFG[\"num_levels\"])]\n            all_rois = np.concatenate(_parts, axis=0)\n            try:\n                np.savez(disk_path, all_rois=all_rois, counts=np.array(counts))\n            except OSError as e:\n                print(f\"  [WARN] Could not write disk cache for study {study_id}: {e} (continuing without it)\")\n\n        slice_stacks, slice_counts = [], []\n        for li in range(CFG[\"num_levels\"]):\n            if rois_per_level[li]:\n                # Cap how many ROI crops feed the model for this level -- random subsampling\n                # (not just \"first N\") so we don't systematically bias toward whichever series\n                # happened to be listed first in train_label_coordinates.csv.\n                _level_rois = rois_per_level[li]\n                if len(_level_rois) > self.max_slices_per_level:\n                    _keep_idx = np.random.choice(len(_level_rois), self.max_slices_per_level, replace=False)\n                    _level_rois = [_level_rois[k] for k in _keep_idx]\n                stacked = np.concatenate(_level_rois, axis=0)\n            else:\n                # No labelled coordinate for this level in this study -> one blank slice as a\n                # placeholder; valid_mask already marks the corresponding targets invalid where\n                # train.csv had no label, but a level can still be masked-in with no ROI found,\n                # so we feed a zero slice rather than crashing the batch collation.\n                stacked = np.zeros((1, CFG[\"img_size_2d\"], CFG[\"img_size_2d\"]), dtype=np.float32)\n            slice_counts.append(stacked.shape[0])\n            slice_stacks.append(stacked)\n        slice_stack = np.concatenate(slice_stacks, axis=0)\n\n        return {\n            \"slice_stack\": torch.from_numpy(slice_stack).unsqueeze(1).float(),\n            \"slice_counts\": slice_counts,\n            \"targets\": torch.from_numpy(targets),\n            \"valid_mask\": torch.from_numpy(valid_mask),\n            \"study_id\": study_id,\n        }\n\n\ndef severity_collate_fn(batch):\n    '''Concatenates every sample's slice stack along dim 0 and flattens slice_counts to length\n    B * num_levels, matching what SeverityGraphModel.forward expects.'''\n    slice_stack = torch.cat([b[\"slice_stack\"] for b in batch], dim=0)\n    slice_counts = []\n    for b in batch:\n        slice_counts.extend(b[\"slice_counts\"])\n    targets = torch.stack([b[\"targets\"] for b in batch], dim=0)\n    valid_mask = torch.stack([b[\"valid_mask\"] for b in batch], dim=0)\n    return {\"slice_stack\": slice_stack, \"slice_counts\": slice_counts, \"targets\": targets, \"valid_mask\": valid_mask}\n\n# ---- Self-test: class interface + collate_fn on fabricated samples ----------\ntry:\n    assert issubclass(SeverityDataset, Dataset)\n    _fake_samples = []\n    for _ in range(2):\n        _counts = [2] * CFG[\"num_levels\"]\n        _fake_samples.append({\n            \"slice_stack\": torch.randn(sum(_counts), 1, CFG[\"img_size_2d\"], CFG[\"img_size_2d\"]),\n            \"slice_counts\": _counts,\n            \"targets\": torch.randint(0, CFG[\"num_classes\"], (CFG[\"num_levels\"], len(CFG[\"conditions\"]))),\n            \"valid_mask\": torch.ones(CFG[\"num_levels\"], len(CFG[\"conditions\"]), dtype=torch.bool),\n            \"study_id\": 0,\n        })\n    _batch = severity_collate_fn(_fake_samples)\n    assert len(_batch[\"slice_counts\"]) == 2 * CFG[\"num_levels\"]\n    assert _batch[\"targets\"].shape == (2, CFG[\"num_levels\"], len(CFG[\"conditions\"]))\n    cell_status(\"CELL 13b — SeverityDataset & collate_fn\", extra=\"collate_fn produces correctly shaped batch\")\n    del _fake_samples, _batch\nexcept Exception as e:\n    cell_status(\"CELL 13b — SeverityDataset & collate_fn\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:45:59.169154Z","iopub.execute_input":"2026-08-02T05:45:59.169457Z","iopub.status.idle":"2026-08-02T05:45:59.201872Z","shell.execute_reply.started":"2026-08-02T05:45:59.169433Z","shell.execute_reply":"2026-08-02T05:45:59.201194Z"}},"outputs":[],"execution_count":null},{"id":"9dcc537d-ebe8-4a9e-aa7f-234c96e3aa28","cell_type":"markdown","source":"## Stage 3 — Multi-Level Explainable AI\n\nEvery explanation method below is real and runnable, but each targets a *different* part of the\nnetwork:\n\n| Method | Applies to | What it explains |\n|---|---|---|\n| GradCAM++ | EfficientNetV2-L (CNN) | which pixels of a 2D ROI slice drove the embedding |\n| Integrated Gradients | EfficientNetV2-L (CNN) | pixel attribution with a theoretical completeness guarantee |\n| GraphCAM | GATv2 graph encoders | which *nodes* (slices/levels/conditions) drove the final logit |\n| GNNExplainer-lite | GATv2 graph encoders | a learned soft mask over nodes/edges that best preserves the prediction |\n| Attention / node / edge importance | GATv2 + Cross-Graph Attention | raw attention weights, read directly off the model (no extra computation) |\n| MC-Dropout uncertainty | full model | predictive variance across stochastic forward passes |\n","metadata":{}},{"id":"7574a1c1-7327-485c-928e-5bc7ae83e167","cell_type":"code","source":"# =====================================================================\n# CELL 14 — GradCAM++ and Integrated Gradients on the CNN embedder\n# =====================================================================\nfrom pytorch_grad_cam import GradCAMPlusPlus\nfrom pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\nfrom captum.attr import IntegratedGradients\n\n\ndef run_gradcampp(embedder_with_head, roi_slice_tensor, target_class):\n    '''embedder_with_head: a small nn.Module wrapping ROIEmbedder + one classification head, so\n    GradCAM++ has a scalar target to differentiate. roi_slice_tensor: (1, 1, H, W).'''\n    target_layer = embedder_with_head.embedder.backbone.conv_head  # last conv layer of EfficientNetV2\n    cam = GradCAMPlusPlus(model=embedder_with_head, target_layers=[target_layer])\n    grayscale_cam = cam(input_tensor=roi_slice_tensor, targets=[ClassifierOutputTarget(target_class)])\n    return grayscale_cam[0]  # (H, W) heatmap in [0, 1]\n\n\ndef run_integrated_gradients(embedder_with_head, roi_slice_tensor, target_class, baseline=None):\n    ig = IntegratedGradients(embedder_with_head)\n    baseline = baseline if baseline is not None else torch.zeros_like(roi_slice_tensor)\n    attributions, delta = ig.attribute(\n        roi_slice_tensor, baselines=baseline, target=target_class,\n        n_steps=64, return_convergence_delta=True,\n    )\n    return attributions.squeeze().detach().cpu().numpy(), float(delta.abs().mean())\n\n\nclass EmbedderWithHead(nn.Module):\n    '''Thin wrapper so pixel-attribution tools (which expect a single classifier) can target\n    one condition's classification head through the shared ROI embedder.'''\n    def __init__(self, embedder, cls_head):\n        super().__init__()\n        self.embedder = embedder\n        self.cls_head = cls_head\n\n    def forward(self, x):\n        return self.cls_head(self.embedder(x))\n\n# ---- Lightweight check: functions/classes defined, callable objects --------\ntry:\n    assert callable(run_gradcampp) and callable(run_integrated_gradients)\n    assert issubclass(EmbedderWithHead, nn.Module)\n    cell_status(\"CELL 14 — GradCAM++ & Integrated Gradients\", extra=\"functions defined (exercised later on a trained model + real ROI, not tested standalone here)\")\nexcept Exception as e:\n    cell_status(\"CELL 14 — GradCAM++ & Integrated Gradients\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:46:06.446482Z","iopub.execute_input":"2026-08-02T05:46:06.446971Z","iopub.status.idle":"2026-08-02T05:46:06.456043Z","shell.execute_reply.started":"2026-08-02T05:46:06.446941Z","shell.execute_reply":"2026-08-02T05:46:06.455446Z"}},"outputs":[],"execution_count":null},{"id":"a740d0a5-402b-4230-9799-23184bcfa490","cell_type":"code","source":"# =====================================================================\n# CELL 15 — GraphCAM and a lightweight GNNExplainer over the GATv2 graphs\n# =====================================================================\ndef graphcam(node_embeds, logits, target_class_idx):\n    '''Gradient-weighted node importance (GraphCAM): backprop the target logit to the node\n    embeddings and weight each node's activation by its gradient, analogous to GradCAM but for\n    graph nodes instead of conv feature maps.'''\n    node_embeds.retain_grad()\n    target = logits[..., target_class_idx].sum()\n    target.backward(retain_graph=True)\n    grad = node_embeds.grad  # (B, L, C, dim)\n    importance = F.relu((grad * node_embeds).sum(-1))  # (B, L, C)\n    return importance.detach()\n\n\ndef gnn_explainer_lite(model, slice_stack, slice_counts, target_level, target_condition,\n                        target_class, steps=100, lr=0.05, l1_lambda=0.01):\n    '''A simplified GNNExplainer: learns a soft edge mask (per graph) that maximizes the\n    target logit while being sparse (L1 penalty), so the surviving high-weight edges/nodes are the\n    ones the model actually relies on. This mirrors the intent of Ying et al. (2019) GNNExplainer\n    without depending on torch_geometric's explainer API, which expects a single homogeneous graph.'''\n    model.eval()\n    for p in model.parameters():\n        p.requires_grad_(False)\n\n    out = model(slice_stack, slice_counts)\n    num_slices = slice_stack.shape[0]\n    edge_mask = nn.Parameter(torch.ones(num_slices, device=slice_stack.device) * 0.5)\n    optimizer = torch.optim.Adam([edge_mask], lr=lr)\n\n    for step in range(steps):\n        optimizer.zero_grad()\n        masked_stack = slice_stack * torch.sigmoid(edge_mask).view(-1, 1, 1, 1)\n        out = model(masked_stack, slice_counts)\n        target_logit = out[\"logits\"][0, target_level, target_condition, target_class]\n        loss = -target_logit + l1_lambda * torch.sigmoid(edge_mask).sum()\n        loss.backward()\n        optimizer.step()\n\n    for p in model.parameters():\n        p.requires_grad_(True)\n\n    return torch.sigmoid(edge_mask).detach().cpu().numpy()  # importance per input slice, in [0, 1]\n\n# ---- Self-test: graphcam gradient-weighting on synthetic node embeddings ----\n# IMPORTANT: logits must be computed FROM node_embeds through a differentiable op (exactly like\n# the real model, where both come out of the same forward pass) -- two independent torch.randn()\n# tensors have no computational edge between them, so backward() would never populate\n# node_embeds.grad. A tiny nn.Linear stands in for the real cls_heads here.\ntry:\n    _test_nodes = torch.randn(1, CFG[\"num_levels\"], len(CFG[\"conditions\"]), CFG[\"graph_hidden\"], requires_grad=True)\n    _test_head = nn.Linear(CFG[\"graph_hidden\"], CFG[\"num_classes\"])\n    _test_logits = _test_head(_test_nodes)  # (1, L, C, num_classes), differentiably tied to _test_nodes\n    _importance = graphcam(_test_nodes, _test_logits, target_class_idx=2)\n    assert _importance.shape == (1, CFG[\"num_levels\"], len(CFG[\"conditions\"])), _importance.shape\n    assert callable(gnn_explainer_lite)\n    cell_status(\"CELL 15 — GraphCAM & GNNExplainer-lite\", extra=f\"graphcam importance shape {tuple(_importance.shape)} OK; gnn_explainer_lite defined (needs a real model to run)\")\n    del _test_nodes, _test_head, _test_logits, _importance\nexcept Exception as e:\n    cell_status(\"CELL 15 — GraphCAM & GNNExplainer-lite\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:46:09.952456Z","iopub.execute_input":"2026-08-02T05:46:09.953375Z","iopub.status.idle":"2026-08-02T05:46:09.979474Z","shell.execute_reply.started":"2026-08-02T05:46:09.953343Z","shell.execute_reply":"2026-08-02T05:46:09.978862Z"}},"outputs":[],"execution_count":null},{"id":"2526c931-1245-46b8-a921-7557236ab18d","cell_type":"code","source":"# =====================================================================\n# CELL 16 — Attention / node / edge importance readout + MC-Dropout uncertainty\n# =====================================================================\ndef extract_attention_importance(attn_maps):\n    '''attn_maps comes straight out of HierarchicalGraphTransformer.forward(); GATv2Conv already\n    returns per-edge attention coefficients, so no extra computation is needed here — we just\n    aggregate them into node-level and edge-level importance scores.'''\n    results = {}\n    for name, (edge_index, alpha) in {\n        k: v for k, v in attn_maps.items() if k in (\"slice_attn\", \"level_attn\", \"neuro_attn\")\n    }.items():\n        alpha_mean = alpha.mean(dim=-1)  # average over attention heads -> (num_edges,)\n        results[f\"{name}_edge_importance\"] = alpha_mean.detach().cpu().numpy()\n        # Node importance: sum of incoming edge weights\n        num_nodes = int(edge_index.max().item()) + 1 if edge_index.numel() > 0 else 0\n        node_importance = torch.zeros(num_nodes)\n        for e in range(edge_index.shape[1]):\n            node_importance[edge_index[1, e]] += alpha_mean[e].item()\n        results[f\"{name}_node_importance\"] = node_importance.numpy()\n    return results\n\n\n@torch.no_grad()\ndef mc_dropout_uncertainty(model, slice_stack, slice_counts, n_passes=20):\n    '''Enable dropout at inference time and repeat the forward pass n_passes times.\n    The predictive mean approximates the model's best estimate; the variance approximates\n    epistemic uncertainty (Gal & Ghahramani, 2016).'''\n    def enable_dropout(m):\n        for module in m.modules():\n            if isinstance(module, nn.Dropout):\n                module.train()\n\n    model.eval()\n    enable_dropout(model)\n\n    probs = []\n    for _ in range(n_passes):\n        out = model(slice_stack, slice_counts)\n        probs.append(F.softmax(out[\"logits\"], dim=-1).unsqueeze(0))\n    probs = torch.cat(probs, dim=0)  # (n_passes, B, L, C, K)\n\n    mean_probs = probs.mean(0)\n    variance = probs.var(0)  # per-class predictive variance -> higher = less certain\n    predictive_entropy = -(mean_probs * (mean_probs.clamp_min(1e-8)).log()).sum(-1)\n    return {\"mean_probs\": mean_probs, \"variance\": variance, \"entropy\": predictive_entropy}\n\n# ---- Self-test: extract_attention_importance on synthetic GATv2-style output -\ntry:\n    _n = 4\n    _ei = torch.tensor([[0, 1, 2, 3], [1, 2, 3, 0]], dtype=torch.long)\n    _alpha = torch.rand(4, CFG[\"graph_heads\"])\n    _fake_attn_maps = {\n        \"slice_attn\": (_ei, _alpha), \"level_attn\": (_ei, _alpha), \"neuro_attn\": (_ei, _alpha),\n        # level_to_neuro_attn and neuro_to_level_attn removed (CrossGraphAttention removed)\n    }\n    _importance = extract_attention_importance(_fake_attn_maps)\n    assert \"slice_attn_edge_importance\" in _importance and \"slice_attn_node_importance\" in _importance\n    assert callable(mc_dropout_uncertainty)\n    cell_status(\"CELL 16 — Attention/node/edge importance & MC-Dropout\", extra=\"extract_attention_importance OK; mc_dropout_uncertainty defined (needs a real model to run)\")\n    del _ei, _alpha, _fake_attn_maps, _importance\nexcept Exception as e:\n    cell_status(\"CELL 16 — Attention/node/edge importance & MC-Dropout\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:46:13.502147Z","iopub.execute_input":"2026-08-02T05:46:13.502618Z","iopub.status.idle":"2026-08-02T05:46:13.515045Z","shell.execute_reply.started":"2026-08-02T05:46:13.502590Z","shell.execute_reply":"2026-08-02T05:46:13.514391Z"}},"outputs":[],"execution_count":null},{"id":"030cfa40-52e5-46d6-a9e8-3c75ff56bfa3","cell_type":"markdown","source":"## Stage 4 — Clinical Consistency Checking (simplified, honest version)\n\nThe original architecture calls for a full **neuro-symbolic reasoning engine** built on an\n**anatomical knowledge graph**. A clinically-validated knowledge graph of this kind does not exist\nin the RSNA dataset and is out of scope to author from scratch for a thesis project — doing so\nproperly would require a curated ontology reviewed by radiologists.\n\nWhat *is* implementable and defensible is a transparent **rule-based consistency layer**: a small\nset of anatomically-motivated rules (encoded as plain Python, not a black box) that flag\npredictions likely to be wrong, plus standard **confidence calibration**. This is explicitly a\nsimplified stand-in, not a claim of clinical-grade reasoning.\n","metadata":{}},{"id":"de684bab-4369-4d9d-bfac-3871242783cc","cell_type":"code","source":"# =====================================================================\n# CELL 17 — Stage 4: Neuro-Symbolic Explainable Clinical Decision Support\n# System (NS-ECDSS). Six modules, each operating on REAL outputs from Stages\n# 1-3 (not synthetic/hardcoded numbers): predicted severities, calibrated\n# confidences, GraphCAM node importance, GATv2 edge attention, and\n# per-view (Sagittal T1 / Sagittal T2-STIR / Axial T2) predictions.\n# =====================================================================\nCONDITION_IDX = {c: i for i, c in enumerate(CFG[\"conditions\"])}\nSEVERITY_NAMES = [\"Normal/Mild\", \"Moderate\", \"Severe\"]\n\n# A mapping from the 5 competition conditions to the anatomical sub-structure each one grades,\n# used to label nodes in the Anatomical Knowledge Graph below.\nCONDITION_TO_SUBSTRUCTURE = {\n    \"spinal_canal_stenosis\": \"canal\",\n    \"left_neural_foraminal_narrowing\": \"left_foramina\",\n    \"right_neural_foraminal_narrowing\": \"right_foramina\",\n    \"left_subarticular_stenosis\": \"left_facet\",\n    \"right_subarticular_stenosis\": \"right_facet\",\n}\n\n\n# ---------------------------------------------------------------------------\n# Module 1 — Anatomical Knowledge Graph\n# ---------------------------------------------------------------------------\ndef build_anatomical_knowledge_graph():\n    '''Builds a static (not learned) graph encoding lumbar spine anatomy: the 5 vertebral\n    levels in their real adjacency order, and for each level a sub-node per graded structure\n    (canal, left/right foramina, left/right facet joints). Canal shares a physical boundary with\n    all four other structures at the same level, so it gets an explicit \"shared_anatomy\" edge to\n    each of them -- this is the edge Module 2 leans on to sanity-check canal-vs-foraminal\n    predictions. This graph is fixed medical-anatomy knowledge, not something learned from data,\n    which is the actual meaning of \"knowledge graph\" here (as opposed to the learned Level/Neuro\n    graphs in Stage 2, which encode statistical co-occurrence rather than anatomy).'''\n    G = nx.Graph()\n    levels = CFG[\"levels\"]\n    for lvl in levels:\n        G.add_node(lvl, kind=\"level\")\n    for i in range(len(levels) - 1):\n        G.add_edge(levels[i], levels[i + 1], kind=\"level_adjacency\")\n    for lvl in levels:\n        for cond, subname in CONDITION_TO_SUBSTRUCTURE.items():\n            node_id = f\"{lvl}::{subname}\"\n            G.add_node(node_id, kind=\"condition\", level=lvl, condition=cond)\n            G.add_edge(lvl, node_id, kind=\"level_condition\")\n    for lvl in levels:\n        canal = f\"{lvl}::canal\"\n        for subname in [\"left_foramina\", \"right_foramina\", \"left_facet\", \"right_facet\"]:\n            G.add_edge(canal, f\"{lvl}::{subname}\", kind=\"shared_anatomy\")\n    return G\n\n\n# ---------------------------------------------------------------------------\n# Module 2 — Clinical Rule Verification\n# ---------------------------------------------------------------------------\ndef clinical_rule_verification(severity_grid, importance_grid, low_evidence_thresh=0.15, high_evidence_thresh=0.75):\n    '''Cross-checks each predicted severity against the model's OWN GraphCAM node importance for\n    that exact (level, condition). A \"Severe\" call backed by low importance is a sign the\n    prediction is not well grounded in what the graph transformer actually attended to; a\n    \"Normal/Mild\" call with unusually high importance is the mirror case, worth a second look.\n    importance_grid values are min-max normalized to [0, 1] across the whole study.'''\n    flags = []\n    for lvl in CFG[\"levels\"]:\n        for cond in CFG[\"conditions\"]:\n            sev = severity_grid[lvl][cond]\n            imp = importance_grid[lvl][cond]\n            if sev == 2 and imp < low_evidence_thresh:\n                flags.append({\n                    \"module\": \"clinical_rule_verification\", \"type\": \"low_evidence_severe\",\n                    \"level\": lvl, \"condition\": cond,\n                    \"message\": f\"{cond.replace('_',' ')} at {lvl.upper().replace('_','/')} graded Severe, \"\n                               f\"but the model's own node importance here is low ({imp:.2f}) — \"\n                               f\"prediction reliability reduced.\",\n                })\n            elif sev == 0 and imp > high_evidence_thresh:\n                flags.append({\n                    \"module\": \"clinical_rule_verification\", \"type\": \"high_evidence_normal\",\n                    \"level\": lvl, \"condition\": cond,\n                    \"message\": f\"{cond.replace('_',' ')} at {lvl.upper().replace('_','/')} graded Normal/Mild, \"\n                               f\"but node importance here is unusually high ({imp:.2f}) — worth reviewing.\",\n                })\n    return flags\n\n\n# ---------------------------------------------------------------------------\n# Module 3 — Multi-view Consistency Checking\n# ---------------------------------------------------------------------------\ndef multiview_consistency_check(per_view_severity_grids):\n    '''per_view_severity_grids: dict[view_name] -> severity_grid (or None if that view had no\n    series for this study). Flags any (level, condition) where two views disagree by 2 full\n    severity classes (e.g. one view says Normal/Mild, another says Severe) -- the RSNA dataset\n    provides Sagittal T1, Sagittal T2/STIR, and Axial T2 precisely because a single view can be\n    ambiguous, so real disagreement between views is clinically meaningful, not noise to average\n    away.'''\n    flags = []\n    views = {v: g for v, g in per_view_severity_grids.items() if g is not None}\n    if len(views) < 2:\n        return flags\n    for lvl in CFG[\"levels\"]:\n        for cond in CFG[\"conditions\"]:\n            values = {v: g[lvl][cond] for v, g in views.items()}\n            spread = max(values.values()) - min(values.values())\n            if spread >= 2:\n                flags.append({\n                    \"module\": \"multiview_consistency\", \"type\": \"multiview_disagreement\",\n                    \"level\": lvl, \"condition\": cond,\n                    \"message\": f\"{cond.replace('_',' ')} at {lvl.upper().replace('_','/')}: views disagree \"\n                               f\"({', '.join(f'{k}={SEVERITY_NAMES[v]}' for k, v in values.items())}) — \"\n                               f\"confidence should be reduced.\",\n                })\n    return flags\n\n\n# ---------------------------------------------------------------------------\n# Module 4 — Severity Consistency Checking (across adjacent vertebral levels)\n# ---------------------------------------------------------------------------\ndef severity_consistency_check(level_severities_for_condition, condition_name):\n    '''level_severities_for_condition: list of 5 severity ints (0/1/2), one per level, for a\n    single condition, in level order. Degeneration severity normally varies smoothly along the\n    spine; an isolated Severe level sandwiched between two Normal/Mild neighbours is anatomically\n    atypical (adjacent segments usually share at least some load-related degeneration) and is\n    flagged for manual review rather than silently trusted.'''\n    flags = []\n    for i in range(1, len(level_severities_for_condition) - 1):\n        prev_s, cur_s, next_s = level_severities_for_condition[i - 1:i + 2]\n        if cur_s == 2 and prev_s == 0 and next_s == 0:\n            flags.append({\n                \"module\": \"severity_consistency\", \"type\": \"isolated_severe_level\",\n                \"level\": CFG[\"levels\"][i], \"condition\": condition_name,\n                \"message\": f\"{condition_name.replace('_',' ')} at {CFG['levels'][i].upper().replace('_','/')} \"\n                           f\"graded Severe with Normal/Mild neighbours on both sides — isolated finding, \"\n                           f\"recommend manual review.\",\n            })\n    return flags\n\n\n# ---------------------------------------------------------------------------\n# Module 5 — Explainability Consistency Checking\n# ---------------------------------------------------------------------------\ndef explainability_consistency_check(severity_grid, importance_grid):\n    '''For each condition, checks whether the level with the HIGHEST predicted severity is also\n    the level the model's own explanation (GraphCAM node importance) points to most strongly. If\n    the model says \"L4/L5 is the worst level\" but its own explanation is actually strongest at\n    L2/L3, that is a genuine conflict between the prediction and the model's explanation of\n    itself -- exactly the kind of thing a radiologist would want surfaced rather than hidden.'''\n    flags = []\n    for cond in CFG[\"conditions\"]:\n        sev_by_level = {lvl: severity_grid[lvl][cond] for lvl in CFG[\"levels\"]}\n        imp_by_level = {lvl: importance_grid[lvl][cond] for lvl in CFG[\"levels\"]}\n        max_sev_level = max(sev_by_level, key=sev_by_level.get)\n        max_imp_level = max(imp_by_level, key=imp_by_level.get)\n        if sev_by_level[max_sev_level] >= 1 and max_sev_level != max_imp_level:\n            flags.append({\n                \"module\": \"explainability_consistency\", \"type\": \"explanation_conflict\",\n                \"condition\": cond, \"severity_level\": max_sev_level, \"explanation_level\": max_imp_level,\n                \"message\": f\"{cond.replace('_',' ')}: highest severity is predicted at \"\n                           f\"{max_sev_level.upper().replace('_','/')}, but the model's explanation \"\n                           f\"(GraphCAM node importance) is strongest at \"\n                           f\"{max_imp_level.upper().replace('_','/')} — explanation conflict.\",\n            })\n    return flags\n\n\n# ---------------------------------------------------------------------------\n# Module 6a — Confidence Calibration (temperature scaling, defined earlier in\n# the pipeline as TemperatureScaler) + a transparent rule-based penalty on top\n# ---------------------------------------------------------------------------\nclass TemperatureScaler(nn.Module):\n    '''Post-hoc confidence calibration (Guo et al., 2017): learns a single scalar T on a held-out\n    validation set so that softmax(logits / T) better reflects true accuracy.'''\n    def __init__(self):\n        super().__init__()\n        self.temperature = nn.Parameter(torch.ones(1) * 1.5)\n\n    def forward(self, logits):\n        return logits / self.temperature\n\n    def fit(self, logits, targets, lr=0.01, steps=200):\n        optimizer = torch.optim.LBFGS([self.temperature], lr=lr, max_iter=steps)\n        def closure():\n            optimizer.zero_grad()\n            loss = F.cross_entropy(self(logits), targets)\n            loss.backward()\n            return loss\n        optimizer.step(closure)\n        return self.temperature.item()\n\n\ndef apply_rule_based_penalty(calibrated_confidence, level, condition, all_flags, penalty_per_flag=0.08):\n    '''Temperature scaling calibrates confidence against a held-out set on average, but it cannot\n    know that THIS particular (level, condition) was just flagged as inconsistent by Modules 2-5.\n    This applies a small, transparent, additive penalty per relevant flag -- e.g. \"99% confident,\n    but 2 clinical-rule flags fired for this exact finding -> 99% * (1 - 2*0.08) ≈ 84%\" -- so the\n    reported confidence reflects both statistical calibration AND the reasoning engine's own\n    doubts about this specific finding.'''\n    n_relevant = sum(1 for f in all_flags if f.get(\"level\") == level and f.get(\"condition\") == condition)\n    penalty = min(0.9, penalty_per_flag * n_relevant)\n    return max(0.0, calibrated_confidence * (1 - penalty))\n\n\n# ---------------------------------------------------------------------------\n# Module 6b — Clinical Evidence Extraction\n# ---------------------------------------------------------------------------\ndef extract_clinical_evidence(severity_grid, confidence_grid, importance_grid, provenance, kg):\n    '''Aggregates everything above into the structured evidence object Stage 5 needs: the single\n    most clinically significant finding (highest severity x confidence), which anatomical\n    connections in the knowledge graph are most relevant to it, and which DICOM series /\n    instance numbers actually support it (traced back through `provenance`, which\n    run_full_pipeline records while cropping ROIs -- NOT invented after the fact).'''\n    best = None\n    for lvl in CFG[\"levels\"]:\n        for cond in CFG[\"conditions\"]:\n            score = severity_grid[lvl][cond] * confidence_grid[lvl][cond]\n            if best is None or score > best[2]:\n                best = (lvl, cond, score)\n    most_important_level, most_important_pathology, _ = best\n\n    influential_connections = []\n    if kg.has_node(most_important_level):\n        for neighbor in kg.neighbors(most_important_level):\n            if kg.nodes[neighbor].get(\"kind\") == \"level\":\n                influential_connections.append(f\"{most_important_level.upper().replace('_','/')} -> \"\n                                                f\"{neighbor.upper().replace('_','/')}\")\n\n    level_idx = CFG[\"levels\"].index(most_important_level)\n    supporting = provenance.get(level_idx, [])\n    supporting_series = sorted({p[\"series_id\"] for p in supporting})\n    supporting_instances = sorted({p[\"instance_number_approx\"] for p in supporting if p[\"instance_number_approx\"] is not None})\n\n    return {\n        \"most_important_level\": most_important_level,\n        \"most_important_pathology\": most_important_pathology,\n        \"most_influential_connections\": influential_connections,\n        \"supporting_series_ids\": supporting_series,\n        \"supporting_instance_numbers\": supporting_instances,  # approximate -- see provenance note in run_full_pipeline\n        \"importance_grid\": importance_grid,\n    }\n\n\n# ---------------------------------------------------------------------------\n# Evidence Graph (bonus): a small, prediction-specific subgraph of the\n# Anatomical Knowledge Graph, weighted by this study's actual importance\n# scores -- lets a radiologist visually trace which structures drove the\n# decision, instead of only reading numbers.\n# ---------------------------------------------------------------------------\ndef build_evidence_graph(kg, severity_grid, importance_grid, top_k_levels=3):\n    G = nx.Graph()\n    level_scores = {lvl: max(severity_grid[lvl][c] * importance_grid[lvl][c] for c in CFG[\"conditions\"])\n                     for lvl in CFG[\"levels\"]}\n    top_levels = sorted(level_scores, key=level_scores.get, reverse=True)[:top_k_levels]\n    for lvl in top_levels:\n        G.add_node(lvl, weight=float(level_scores[lvl]) + 0.1, kind=\"level\")\n        for cond in CFG[\"conditions\"]:\n            imp = importance_grid[lvl][cond]\n            if imp > 0.1:\n                node_id = f\"{lvl}::{cond}\"\n                G.add_node(node_id, weight=float(imp), kind=\"condition\")\n                G.add_edge(lvl, node_id, weight=float(imp))\n    for i, l1 in enumerate(top_levels):\n        for l2 in top_levels[i + 1:]:\n            if kg.has_edge(l1, l2):\n                G.add_edge(l1, l2, weight=0.3, kind=\"adjacency\")\n    return G\n\n\ndef plot_evidence_graph(G, out_path, title=\"Evidence Graph — most influential anatomical structures\"):\n    import matplotlib.pyplot as plt\n    if G.number_of_nodes() == 0:\n        return None\n    pos = nx.spring_layout(G, seed=42)\n    node_sizes = [max(300, G.nodes[n].get(\"weight\", 0.3) * 1500) for n in G.nodes]\n    node_colors = [\"#4C72B0\" if G.nodes[n].get(\"kind\") == \"level\" else \"#DD8452\" for n in G.nodes]\n    labels = {n: (n.upper().replace(\"_\", \"/\") if G.nodes[n].get(\"kind\") == \"level\"\n                  else n.split(\"::\")[-1].replace(\"_\", \" \")) for n in G.nodes}\n    plt.figure(figsize=(8, 6))\n    nx.draw(G, pos, labels=labels, node_size=node_sizes, node_color=node_colors,\n            font_size=8, edge_color=\"#999999\", width=[G[u][v].get(\"weight\", 0.3) * 3 for u, v in G.edges])\n    plt.title(title)\n    plt.tight_layout()\n    plt.savefig(out_path, dpi=150)\n    plt.close()\n    return out_path\n\n\n# ---------------------------------------------------------------------------\n# Orchestrator: runs all 6 modules + evidence extraction in order\n# ---------------------------------------------------------------------------\ndef run_stage4_ns_ecdss(severity_grid, confidence_grid, importance_grid, per_view_severity_grids, provenance, kg):\n    all_flags = []\n    all_flags += clinical_rule_verification(severity_grid, importance_grid)\n    all_flags += multiview_consistency_check(per_view_severity_grids)\n    for cond in CFG[\"conditions\"]:\n        all_flags += severity_consistency_check([severity_grid[lvl][cond] for lvl in CFG[\"levels\"]], cond)\n    all_flags += explainability_consistency_check(severity_grid, importance_grid)\n\n    calibrated_confidence_grid = {lvl: {} for lvl in CFG[\"levels\"]}\n    for lvl in CFG[\"levels\"]:\n        for cond in CFG[\"conditions\"]:\n            calibrated_confidence_grid[lvl][cond] = apply_rule_based_penalty(\n                confidence_grid[lvl][cond], lvl, cond, all_flags)\n\n    evidence = extract_clinical_evidence(severity_grid, calibrated_confidence_grid, importance_grid, provenance, kg)\n\n    return {\n        \"flags\": all_flags,\n        \"calibrated_confidence_grid\": calibrated_confidence_grid,\n        \"evidence\": evidence,\n    }\n\n# ---- Self-test: every module runs on synthetic data and returns well-formed output ----\ntry:\n    _n_lvl, _n_cond = CFG[\"num_levels\"], len(CFG[\"conditions\"])\n    _kg = build_anatomical_knowledge_graph()\n    assert _kg.number_of_nodes() == _n_lvl + _n_lvl * _n_cond\n    assert _kg.number_of_edges() > 0\n\n    _sev = {lvl: {c: int(np.random.randint(0, 3)) for c in CFG[\"conditions\"]} for lvl in CFG[\"levels\"]}\n    _conf = {lvl: {c: float(np.random.rand()) for c in CFG[\"conditions\"]} for lvl in CFG[\"levels\"]}\n    _imp = {lvl: {c: float(np.random.rand()) for c in CFG[\"conditions\"]} for lvl in CFG[\"levels\"]}\n    _prov = {i: [{\"series_id\": 111, \"instance_number_approx\": 10 + i}] for i in range(_n_lvl)}\n    _views = {\"Sagittal T1\": _sev, \"Axial T2\": {lvl: {c: int(np.random.randint(0, 3)) for c in CFG[\"conditions\"]} for lvl in CFG[\"levels\"]}}\n\n    _result = run_stage4_ns_ecdss(_sev, _conf, _imp, _views, _prov, _kg)\n    assert \"flags\" in _result and \"calibrated_confidence_grid\" in _result and \"evidence\" in _result\n    assert isinstance(_result[\"flags\"], list)\n\n    _ts = TemperatureScaler()\n    _t_val = _ts.fit(torch.randn(20, CFG[\"num_classes\"]), torch.randint(0, CFG[\"num_classes\"], (20,)), steps=10)\n    assert np.isfinite(_t_val)\n\n    _eg = build_evidence_graph(_kg, _sev, _imp)\n    _png_path = plot_evidence_graph(_eg, os.path.join(CFG[\"temp_dir\"], \"test_evidence_graph.png\"))\n\n    cell_status(\"CELL 17 — Stage 4 NS-ECDSS (6 modules + evidence graph)\",\n                extra=f\"{len(_result['flags'])} flags on synthetic data, KG has {_kg.number_of_nodes()} nodes, evidence graph plotted OK\")\n    del _kg, _sev, _conf, _imp, _prov, _views, _result, _ts, _t_val, _eg, _png_path\nexcept Exception as e:\n    cell_status(\"CELL 17 — Stage 4 NS-ECDSS (6 modules + evidence graph)\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:46:18.689527Z","iopub.execute_input":"2026-08-02T05:46:18.690270Z","iopub.status.idle":"2026-08-02T05:46:18.947134Z","shell.execute_reply.started":"2026-08-02T05:46:18.690232Z","shell.execute_reply":"2026-08-02T05:46:18.946552Z"}},"outputs":[],"execution_count":null},{"id":"7846d70e-38ae-4b75-ad7b-11c47d4b9308","cell_type":"markdown","source":"## Stage 5 — Structured Clinical Report Generator (template-based)\n\nRather than claiming a fine-tuned \"Medical LLM\" (which would require licensed clinical training\ndata we don't have), Stage 5 fills a **structured template** directly from Stage 2 predictions,\nStage 3 explanations, and Stage 4 consistency checks. Optionally, a general-purpose LLM can be\ncalled to turn the structured fields into flowing prose — but the *facts* in the report always come\nfrom the deterministic template, never from free LLM generation, so nothing is hallucinated.\n","metadata":{}},{"id":"5fb0f8d3-a992-4318-a84c-4741819304dd","cell_type":"code","source":"# =====================================================================\n# CELL 18 — Stage 5: Explainable Clinical Report Generator\n# =====================================================================\n# The report is built entirely from the structured facts computed in Stages 1-4 (severity,\n# calibrated confidence, evidence, consistency flags). An LLM is OPTIONAL and, if used, is only\n# allowed to rewrite the \"Model Explanation\" sentence for readability -- it is never given the\n# freedom to invent a grade, a slice number, or a consistency result, so nothing in the report\n# can be hallucinated by the LLM step. This is a deliberate design choice, not a shortcut: the\n# actual reasoning happens in Stage 4 (NS-ECDSS), and Stage 5's job is presentation only.\n\ndef _format_finding_block(study_id, level, condition, severity_grid, calibrated_confidence_grid,\n                            evidence, flags, view_used=\"Sagittal T2/STIR\"):\n    '''Formats one finding in the exact field layout requested: Level / Condition / Severity /\n    Confidence / Supporting Sequence / Supporting Slices / Most Important Vertebra / Most\n    Influential Anatomical Connection / Model Explanation / Clinical Consistency / Recommendation.'''\n    sev = severity_grid[level][condition]\n    conf = calibrated_confidence_grid[level][condition]\n    relevant_flags = [f for f in flags if f.get(\"level\") == level and f.get(\"condition\") == condition]\n\n    consistency_lines = []\n    checked_modules = [\"multiview_consistency\", \"severity_consistency\", \"explainability_consistency\", \"clinical_rule_verification\"]\n    module_labels = {\n        \"multiview_consistency\": \"Multi-view agreement\",\n        \"severity_consistency\": \"Adjacent-level consistency\",\n        \"explainability_consistency\": \"Explainability consistency\",\n        \"clinical_rule_verification\": \"Clinical rule verification\",\n    }\n    for m in checked_modules:\n        failed = any(f[\"module\"] == m for f in relevant_flags)\n        mark = \"\\u2717\" if failed else \"\\u2713\"\n        consistency_lines.append(f\"{mark} {module_labels[m]}\")\n\n    explanation = (\n        f\"The prediction at {level.upper().replace('_','/')} for {condition.replace('_',' ')} is \"\n        f\"supported by GraphCAM node importance of {evidence['importance_grid'][level][condition]:.2f} \"\n        f\"(0-1 scale) at this level. \"\n    )\n    if evidence[\"most_influential_connections\"]:\n        explanation += f\"The most influential anatomical connection was {evidence['most_influential_connections'][0]}. \"\n    if relevant_flags:\n        explanation += f\"NOTE: {len(relevant_flags)} consistency flag(s) were raised for this finding (see below) -- interpret with caution.\"\n    else:\n        explanation += \"No consistency flags were raised for this finding.\"\n\n    recommendation = (\n        \"Further radiologist review is advised due to the Severe grade at this level.\"\n        if sev == 2 else\n        \"Routine follow-up per standard protocol; no urgent flag raised by the automated review.\"\n        if sev == 0 else\n        \"Radiologist review recommended given the Moderate grade and any flags noted above.\"\n    )\n\n    lines = [\n        \"-\" * 50,\n        f\"Affected Level       {level.upper().replace('_','/')}\",\n        \"-\" * 50,\n        f\"Condition            {condition.replace('_',' ').title()}\",\n        \"-\" * 50,\n        f\"Severity             {SEVERITY_NAMES[sev]}\",\n        \"-\" * 50,\n        f\"Prediction Confidence {conf:.1%}\",\n        \"-\" * 50,\n        f\"Supporting MRI Sequence  {view_used}\",\n        \"-\" * 50,\n        f\"Supporting Slices (approx. instance #s)  {', '.join(str(s) for s in evidence['supporting_instance_numbers']) or 'n/a'}\",\n        \"-\" * 50,\n        f\"Most Important Vertebra  {evidence['most_important_level'].upper().replace('_','/')}\",\n        \"-\" * 50,\n        f\"Most Influential Anatomical Connection  {evidence['most_influential_connections'][0] if evidence['most_influential_connections'] else 'n/a'}\",\n        \"-\" * 50,\n        f\"Model Explanation\",\n        f\"  {explanation}\",\n        \"-\" * 50,\n        f\"Clinical Consistency\",\n    ] + [f\"  {line}\" for line in consistency_lines] + [\n        \"-\" * 50,\n        f\"Recommendation       {recommendation}\",\n        \"-\" * 50,\n    ]\n    return \"\\n\".join(lines)\n\n\ndef generate_explainable_report(study_id, severity_grid, confidence_grid, uncertainty_grid,\n                                  ns_ecdss_result, view_used=\"Sagittal T2/STIR\", max_findings=3):\n    '''Builds the full report: a detailed block (see _format_finding_block) for the most\n    clinically significant finding plus up to `max_findings`-1 other Moderate/Severe findings,\n    then a compact summary table of every level/condition, then the consistency-flag log and\n    disclaimer.'''\n    calibrated_confidence_grid = ns_ecdss_result[\"calibrated_confidence_grid\"]\n    evidence = ns_ecdss_result[\"evidence\"]\n    flags = ns_ecdss_result[\"flags\"]\n\n    lines = [\"=\" * 50, \"LUMBAR SPINE MRI \\u2014 EXPLAINABLE CLINICAL REPORT\", f\"Study ID: {study_id}\", \"=\" * 50, \"\"]\n\n    # Rank findings by severity * confidence, most significant first\n    ranked = []\n    for lvl in CFG[\"levels\"]:\n        for cond in CFG[\"conditions\"]:\n            score = severity_grid[lvl][cond] * calibrated_confidence_grid[lvl][cond]\n            ranked.append((score, lvl, cond))\n    ranked.sort(reverse=True)\n    notable = [(lvl, cond) for score, lvl, cond in ranked if severity_grid[lvl][cond] >= 1][:max_findings]\n    if not notable:\n        notable = [(ranked[0][1], ranked[0][2])]\n\n    lines.append(f\"DETAILED FINDINGS ({len(notable)} most clinically significant)\")\n    lines.append(\"\")\n    for lvl, cond in notable:\n        lines.append(_format_finding_block(study_id, lvl, cond, severity_grid, calibrated_confidence_grid,\n                                             evidence, flags, view_used))\n        lines.append(\"\")\n\n    lines.append(\"SUMMARY TABLE (all levels / conditions)\")\n    for level in CFG[\"levels\"]:\n        lines.append(f\"  {level.upper().replace('_', '/')}:\")\n        for cond in CFG[\"conditions\"]:\n            sev = SEVERITY_NAMES[severity_grid[level][cond]]\n            conf = calibrated_confidence_grid[level][cond]\n            unc = uncertainty_grid[level][cond]\n            lines.append(f\"    - {cond.replace('_', ' ').title()}: {sev} \"\n                         f\"(confidence {conf:.0%}, predictive entropy {unc:.2f})\")\n    lines.append(\"\")\n\n    lines.append(f\"NS-ECDSS CONSISTENCY LOG ({len(flags)} flag(s) across all modules)\")\n    if flags:\n        for f in flags:\n            lines.append(f\"  [{f['module']}] {f['message']}\")\n    else:\n        lines.append(\"  No anatomical, multi-view, or explainability inconsistencies flagged.\")\n    lines.append(\"\")\n\n    lines.append(\"DISCLAIMER: This report was generated automatically by a research model (NS-ECDSS) and has \"\n                  \"NOT been reviewed by a radiologist. It is provided for thesis / research demonstration \"\n                  \"purposes only and must not be used for clinical decision-making.\")\n    return \"\\n\".join(lines)\n\n\ndef maybe_narrate_with_llm(structured_report_text, use_llm=False, llm_call_fn=None):\n    '''Optional: pass a callable llm_call_fn(prompt) -> str (e.g. wrapping the Anthropic or OpenAI\n    API) to rewrite ONLY the \"Model Explanation\" prose for readability. The structured facts\n    (grades, confidences, flags, slice numbers) are always generated first by Stage 4/5 as ground\n    truth and passed IN to the LLM -- the LLM is instructed to touch only wording, never add,\n    remove, or change a number or a flag, so nothing here can be hallucinated. Left unwired by\n    default (no API key is configured in this notebook); wire llm_call_fn yourself if you want to\n    use it, e.g.:\n        from anthropic import Anthropic\n        client = Anthropic(api_key=\"...\")\n        def llm_call_fn(prompt):\n            msg = client.messages.create(model=\"claude-sonnet-4-6\", max_tokens=1000,\n                                          messages=[{\"role\": \"user\", \"content\": prompt}])\n            return msg.content[0].text\n    '''\n    if not use_llm or llm_call_fn is None:\n        return structured_report_text\n    prompt = (\n        \"Rewrite the following structured radiology findings as clear prose for a thesis \"\n        \"appendix. Do not add, remove, or change any severity grade, confidence value, slice \"\n        \"number, or consistency flag — only improve readability.\\n\\n\" + structured_report_text\n    )\n    return llm_call_fn(prompt)\n\n# ---- Self-test: generate a report end-to-end from fully synthetic data ------\ntry:\n    _n_lvl, _n_cond = CFG[\"num_levels\"], len(CFG[\"conditions\"])\n    _sev = {lvl: {c: int(np.random.randint(0, 3)) for c in CFG[\"conditions\"]} for lvl in CFG[\"levels\"]}\n    _conf = {lvl: {c: float(np.random.rand()) for c in CFG[\"conditions\"]} for lvl in CFG[\"levels\"]}\n    _unc = {lvl: {c: float(np.random.rand()) for c in CFG[\"conditions\"]} for lvl in CFG[\"levels\"]}\n    _imp = {lvl: {c: float(np.random.rand()) for c in CFG[\"conditions\"]} for lvl in CFG[\"levels\"]}\n    _kg = build_anatomical_knowledge_graph()\n    _prov = {i: [{\"series_id\": 111, \"instance_number_approx\": 10 + i}] for i in range(_n_lvl)}\n    _ns_result = run_stage4_ns_ecdss(_sev, _conf, _imp, {\"Sagittal T1\": _sev}, _prov, _kg)\n\n    _report = generate_explainable_report(\"SYNTHETIC_TEST_STUDY\", _sev, _conf, _unc, _ns_result)\n    assert \"SYNTHETIC_TEST_STUDY\" in _report and \"DISCLAIMER\" in _report and \"NS-ECDSS\" in _report\n    _narrated = maybe_narrate_with_llm(_report, use_llm=False)\n    assert _narrated == _report\n    cell_status(\"CELL 18 — Stage 5 explainable report generator\",\n                extra=f\"generated {len(_report.splitlines())}-line report with {len(_ns_result['flags'])} flags on synthetic data\")\n    del _sev, _conf, _unc, _imp, _kg, _prov, _ns_result, _report, _narrated\nexcept Exception as e:\n    cell_status(\"CELL 18 — Stage 5 explainable report generator\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:46:34.338573Z","iopub.execute_input":"2026-08-02T05:46:34.338974Z","iopub.status.idle":"2026-08-02T05:46:34.362176Z","shell.execute_reply.started":"2026-08-02T05:46:34.338946Z","shell.execute_reply":"2026-08-02T05:46:34.361479Z"}},"outputs":[],"execution_count":null},{"id":"8566b997-1291-4eaf-aa33-99ea59976b45","cell_type":"markdown","source":"## End-to-end inference pipeline + export utilities","metadata":{}},{"id":"8e3d2ca1-d2c5-4337-9836-4cfd82485731","cell_type":"code","source":"# =====================================================================\n# CELL 19 — Full inference pipeline tying Stages 1-5 together for one study\n# =====================================================================\n# Refactored into two helpers so the SAME code path serves both the primary (all-views) pass\n# used for the main prediction + GraphCAM explanation, and the per-view-only passes used for\n# Stage 4's multi-view consistency check.\n\n@torch.no_grad()\ndef _collect_rois_for_series_list(study_id, series_ids, stage1_model, data_root=CFG[\"data_root\"]):\n    '''Runs Stage 1 on each given series, pools ROI crops per level, and records PROVENANCE --\n    which series and (approximate) original DICOM instance number each ROI crop actually came\n    from. This provenance is what Stage 4/5 report as \"Supporting Slices\" -- it is read off the\n    real crop, not invented after the fact. The approximation: Stage 1 predicts a normalized\n    depth in the RESAMPLED volume; we map that back to the nearest original instance number via\n    linear interpolation, since resample_depth() itself is a linear index remap.'''\n    all_rois = {li: [] for li in range(CFG[\"num_levels\"])}\n    provenance = {li: [] for li in range(CFG[\"num_levels\"])}\n    for series_id in series_ids:\n        volume, instance_numbers, orig_hw = load_dicom_series(study_id, series_id, data_root)\n        if volume is None or len(instance_numbers) == 0:\n            continue\n        volume_r = resample_depth(volume, CFG[\"vol_size_3d\"][0])\n        vol_t = torch.from_numpy(volume_r).unsqueeze(0).unsqueeze(0).float().to(DEVICE)\n        assert_on_gpu(vol_t, \"Stage1 inference input volume\")\n        out1 = stage1_model(vol_t)\n        depth_pred = out1[\"depth_pred\"][0].cpu().numpy()          # (L,)\n        yx_pred = out1[\"coord_pred\"][0].cpu().numpy()              # (L, 2)\n        coord_pred = np.concatenate([depth_pred[:, None], yx_pred], axis=1)  # (L, 3)\n        inst_logits = out1[\"instance_logits\"][0].cpu()\n        rois = extract_rois(volume_r, coord_pred, inst_logits)\n        for li, (crop, present) in rois.items():\n            if not present:\n                continue\n            all_rois[li].append(roi_to_slices(crop))\n            d_resampled = float(np.clip(coord_pred[li, 0] * (CFG[\"vol_size_3d\"][0] - 1), 0, CFG[\"vol_size_3d\"][0] - 1))\n            d_orig_idx = int(round(d_resampled / max(1, CFG[\"vol_size_3d\"][0] - 1) * (len(instance_numbers) - 1))) \\\n                if len(instance_numbers) > 1 else 0\n            provenance[li].append({\"series_id\": series_id, \"instance_number_approx\": instance_numbers[d_orig_idx]})\n    return all_rois, provenance\n\n\ndef _severity_inference(all_rois, stage2_model, temp_scaler, compute_importance=False):\n    '''Runs Stage 2 + calibration on already-collected ROI crops. When compute_importance=True,\n    this call must NOT be wrapped in torch.no_grad() by the caller -- GraphCAM needs a real\n    backward pass through node_embeds to compute node importance, matching graphcam() from the\n    Stage 3 cell exactly (predicted-class logit, summed over all level/condition nodes in one\n    backward pass for efficiency, then min-max normalized to [0, 1] for use in Stage 4's rules).'''\n    slice_tensors, slice_counts = [], []\n    for li in range(CFG[\"num_levels\"]):\n        stacked = np.concatenate(all_rois[li], axis=0) if all_rois[li] else \\\n            np.zeros((1, CFG[\"img_size_2d\"], CFG[\"img_size_2d\"]), dtype=np.float32)\n        slice_counts.append(stacked.shape[0])\n        slice_tensors.append(torch.from_numpy(stacked).unsqueeze(1))\n    slice_stack = torch.cat(slice_tensors, dim=0).float().to(DEVICE)\n    assert_on_gpu(slice_stack, \"Stage2 inference input slice_stack\")\n\n    out2 = stage2_model(slice_stack, slice_counts)\n    calibrated_logits = temp_scaler(out2[\"logits\"])\n    probs = F.softmax(calibrated_logits, dim=-1)[0]  # (L, C, K)\n\n    severity_grid, confidence_grid = {}, {}\n    for li, lvl in enumerate(CFG[\"levels\"]):\n        severity_grid[lvl], confidence_grid[lvl] = {}, {}\n        for ci, cond in enumerate(CFG[\"conditions\"]):\n            cls = int(probs[li, ci].argmax().item())\n            severity_grid[lvl][cond] = cls\n            confidence_grid[lvl][cond] = float(probs[li, ci].max().item())\n\n    importance_grid = None\n    uncertainty_grid = None\n    if compute_importance:\n        node_embeds = out2[\"node_embeds\"]\n        node_embeds.retain_grad()\n        pred_classes = torch.tensor(\n            [[severity_grid[lvl][cond] for cond in CFG[\"conditions\"]] for lvl in CFG[\"levels\"]],\n            device=calibrated_logits.device)\n        gathered = calibrated_logits[0].gather(-1, pred_classes.unsqueeze(-1)).squeeze(-1)  # (L, C)\n        stage2_model.zero_grad(set_to_none=True)\n        gathered.sum().backward(retain_graph=False)\n        grad = node_embeds.grad\n        raw_importance = F.relu((grad * node_embeds).sum(-1))[0]  # (L, C)\n        raw_importance = (raw_importance - raw_importance.min()) / (raw_importance.max() - raw_importance.min() + 1e-8)\n        importance_grid = {}\n        for li, lvl in enumerate(CFG[\"levels\"]):\n            importance_grid[lvl] = {}\n            for ci, cond in enumerate(CFG[\"conditions\"]):\n                importance_grid[lvl][cond] = float(raw_importance[li, ci].item())\n\n        with torch.no_grad():\n            unc = mc_dropout_uncertainty(stage2_model, slice_stack, slice_counts, n_passes=10)\n        uncertainty_grid = {}\n        for li, lvl in enumerate(CFG[\"levels\"]):\n            uncertainty_grid[lvl] = {cond: float(unc[\"entropy\"][0, li, ci].item())\n                                       for ci, cond in enumerate(CFG[\"conditions\"])}\n\n    return severity_grid, confidence_grid, uncertainty_grid, importance_grid\n\n\nVIEW_NAMES = [\"Sagittal T1\", \"Sagittal T2/STIR\", \"Axial T2\"]\n\n\ndef run_full_pipeline(study_id, series_ids, stage1_model, stage2_model, temp_scaler,\n                       df_series=None, data_root=CFG[\"data_root\"]):\n    stage1_model.eval(); stage2_model.eval()\n\n    # ---- Primary pass: all views combined -> main prediction + GraphCAM importance ----\n    all_rois, provenance = _collect_rois_for_series_list(study_id, series_ids, stage1_model, data_root)\n    severity_grid, confidence_grid, uncertainty_grid, importance_grid = _severity_inference(\n        all_rois, stage2_model, temp_scaler, compute_importance=True)\n\n    # ---- Per-view passes (Module 3: Multi-view Consistency Checking) ----\n    # Requires df_series (train_series_descriptions.csv) to know which series is which view; if\n    # not supplied, multi-view checking is skipped (Stage 4 handles per_view_severity_grids=None\n    # views gracefully) rather than failing the whole pipeline.\n    per_view_severity_grids = {}\n    if df_series is not None:\n        for view in VIEW_NAMES:\n            view_series_ids = df_series[\n                (df_series[\"study_id\"] == study_id) & (df_series[\"series_description\"] == view)\n            ][\"series_id\"].tolist()\n            if not view_series_ids:\n                per_view_severity_grids[view] = None\n                continue\n            view_rois, _ = _collect_rois_for_series_list(study_id, view_series_ids, stage1_model, data_root)\n            with torch.no_grad():\n                v_sev, _, _, _ = _severity_inference(view_rois, stage2_model, temp_scaler, compute_importance=False)\n            per_view_severity_grids[view] = v_sev\n\n    # ---- Stage 4: NS-ECDSS (all 6 modules + evidence extraction) ----\n    kg = build_anatomical_knowledge_graph()\n    ns_result = run_stage4_ns_ecdss(severity_grid, confidence_grid, importance_grid,\n                                      per_view_severity_grids, provenance, kg)\n\n    # ---- Evidence Graph ----\n    evidence_graph = build_evidence_graph(kg, severity_grid, importance_grid)\n    evidence_graph_path = os.path.join(CFG[\"work_dir\"], f\"evidence_graph_{study_id}.png\")\n    plot_evidence_graph(evidence_graph, evidence_graph_path,\n                         title=f\"Evidence Graph — study {study_id}\")\n\n    # ---- Stage 5: explainable report ----\n    primary_view = next((v for v in VIEW_NAMES if per_view_severity_grids.get(v) is not None), \"Sagittal T2/STIR\")\n    report = generate_explainable_report(study_id, severity_grid, confidence_grid, uncertainty_grid,\n                                           ns_result, view_used=primary_view)\n\n    return {\n        \"severity_grid\": severity_grid,\n        \"confidence_grid\": confidence_grid,\n        \"calibrated_confidence_grid\": ns_result[\"calibrated_confidence_grid\"],\n        \"uncertainty_grid\": uncertainty_grid,\n        \"importance_grid\": importance_grid,\n        \"per_view_severity_grids\": per_view_severity_grids,\n        \"ns_ecdss_flags\": ns_result[\"flags\"],\n        \"evidence\": ns_result[\"evidence\"],\n        \"evidence_graph_path\": evidence_graph_path,\n        \"report\": report,\n    }\n\ncell_status(\"CELL 19 — Full inference pipeline (run_full_pipeline)\",\n            extra=\"function defined (requires trained Stage 1 + Stage 2 models and real DICOM data to run — exercised in the Driver cell)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:46:44.406821Z","iopub.execute_input":"2026-08-02T05:46:44.407076Z","iopub.status.idle":"2026-08-02T05:46:44.429308Z","shell.execute_reply.started":"2026-08-02T05:46:44.407056Z","shell.execute_reply":"2026-08-02T05:46:44.428614Z"}},"outputs":[],"execution_count":null},{"id":"089b003a-a028-472b-b845-8d5472fc9d1c","cell_type":"code","source":"# =====================================================================\n# CELL 20 — Export trained models, results, plots, logs, and a project zip\n# =====================================================================\nimport matplotlib.pyplot as plt\n\ndef plot_history(history, title, out_path):\n    plt.figure(figsize=(6, 4))\n    plt.plot(history[\"train_loss\"], label=\"train\")\n    plt.plot(history[\"val_loss\"], label=\"val\")\n    plt.xlabel(\"epoch\"); plt.ylabel(\"loss\"); plt.title(title); plt.legend()\n    plt.tight_layout()\n    plt.savefig(out_path, dpi=150)\n    plt.close()\n\n\ndef export_everything(cfg=CFG, results=None, history_stage1=None, history_stage2=None):\n    out_dir = cfg[\"work_dir\"]\n    if history_stage1:\n        plot_history(history_stage1, \"Stage 1 — Localization loss\", os.path.join(out_dir, \"stage1_loss.png\"))\n    if history_stage2:\n        plot_history(history_stage2, \"Stage 2 — Severity loss\", os.path.join(out_dir, \"stage2_loss.png\"))\n    if results is not None:\n        with open(os.path.join(out_dir, \"results.json\"), \"w\") as f:\n            json.dump(results, f, indent=2, default=str)\n\n    zip_path = os.path.join(out_dir, \"project_export.zip\")\n    with zipfile.ZipFile(zip_path, \"w\", zipfile.ZIP_DEFLATED) as zf:\n        for fname in [\"stage1_best.pt\", \"stage2_best.pt\", \"stage1_loss.png\", \"stage2_loss.png\", \"results.json\"]:\n            fpath = os.path.join(out_dir, fname)\n            if os.path.exists(fpath):\n                zf.write(fpath, arcname=fname)\n        # Include the NS-ECDSS evidence graph image if this run produced one\n        _evidence_path = (results or {}).get(\"evidence_graph_path\")\n        if _evidence_path and os.path.exists(_evidence_path):\n            zf.write(_evidence_path, arcname=os.path.basename(_evidence_path))\n    print(\"Exported project to\", zip_path)\n    return zip_path\n\ntry:\n    assert callable(plot_history) and callable(export_everything)\n    cell_status(\"CELL 20 — Export utilities\", extra=\"plot_history & export_everything defined\")\nexcept Exception as e:\n    cell_status(\"CELL 20 — Export utilities\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:46:49.653881Z","iopub.execute_input":"2026-08-02T05:46:49.654135Z","iopub.status.idle":"2026-08-02T05:46:49.663682Z","shell.execute_reply.started":"2026-08-02T05:46:49.654116Z","shell.execute_reply":"2026-08-02T05:46:49.662855Z"}},"outputs":[],"execution_count":null},{"id":"f411ec68-b9e0-4675-92cd-e905417cbcfe","cell_type":"markdown","source":"## Driver cell (example usage)\n\nWire everything together once the DataLoaders are built from `train.csv`, `train_series_descriptions.csv`,\nand `train_label_coordinates.csv`. Left as an explicit call so each stage can be run/debugged independently\ninside Kaggle before doing a full end-to-end run.\n","metadata":{}},{"id":"e451af47-313e-48be-9adf-7eabaa44c031","cell_type":"code","source":"# =====================================================================\n# CELL 20b — Architecture Diagram (auto-displayed once the full pipeline\n# is built, before any training starts)\n# =====================================================================\ndef draw_architecture_diagram(save_path=None):\n    import matplotlib.pyplot as plt\n    from matplotlib.patches import FancyBboxPatch\n\n    stages = [\n        (\"MRI Volume\", \"#eeeeee\"),\n        (\"STAGE 1\\nAnatomical Localization Network (3D ConvNeXt)\\nInstance / Depth / Coordinate Prediction\", \"#cfe2f3\"),\n        (\"ROI Extraction\", \"#eeeeee\"),\n        (\"STAGE 2 — EfficientNetV2-L\\nROI Embeddings\", \"#d9ead3\"),\n        (\"Slice Graph   |   Level Graph   |   Neuro Graph\", \"#d9ead3\"),\n        (\"Cross Graph Attention\", \"#d9ead3\"),\n        (\"Hierarchical Graph Transformer\", \"#d9ead3\"),\n        (\"Graph Readout -> Severity Classification\", \"#d9ead3\"),\n        (\"STAGE 3 — Explainable AI\\nGradCAM++ / GraphCAM / GNNExplainer / Integrated Gradients\\nAttention, Node & Edge Importance / Uncertainty\", \"#fff2cc\"),\n        (\"STAGE 4 — Neuro-Symbolic Clinical\\nReasoning Engine (NS-ECDSS)\", \"#f4cccc\"),\n        (\"STAGE 5 — Explainable Clinical\\nReport Generator\", \"#d0e0e3\"),\n        (\"Structured Radiology Report\", \"#eeeeee\"),\n    ]\n    box_w, box_h, gap = 8.6, 1.0, 0.45\n    n = len(stages)\n    fig, ax = plt.subplots(figsize=(7.5, n * 0.62))\n\n    cursor_y = n * (box_h + gap)\n    spans = []\n    for label, color in stages:\n        y0 = cursor_y - box_h\n        rect = FancyBboxPatch((1, y0), box_w, box_h, boxstyle=\"round,pad=0.07\",\n                               facecolor=color, edgecolor=\"black\", linewidth=1.3)\n        ax.add_patch(rect)\n        ax.text(1 + box_w / 2, y0 + box_h / 2, label, ha=\"center\", va=\"center\",\n                fontsize=8.5, fontweight=\"bold\", linespacing=1.4)\n        spans.append((y0, y0 + box_h))\n        cursor_y -= (box_h + gap)\n\n    for i in range(n - 1):\n        ax.annotate(\"\", xy=(1 + box_w / 2, spans[i + 1][1]), xytext=(1 + box_w / 2, spans[i][0]),\n                     arrowprops=dict(arrowstyle=\"-|>\", lw=1.6, color=\"#444444\"))\n\n    ax.set_xlim(0, box_w + 2)\n    ax.set_ylim(-0.5, n * (box_h + gap) + 0.5)\n    ax.axis(\"off\")\n    ax.set_title(\"Lumbar Spine NS-ECDSS — Full Pipeline Architecture\", fontsize=13, fontweight=\"bold\", pad=16)\n    plt.tight_layout()\n    if save_path:\n        plt.savefig(save_path, dpi=160, bbox_inches=\"tight\")\n    plt.show()\n    return save_path\n\n\ntry:\n    _diagram_path = draw_architecture_diagram(os.path.join(CFG[\"work_dir\"], \"architecture_diagram.png\"))\n    cell_status(\"CELL 20b — Architecture Diagram\", extra=f\"saved to {_diagram_path}\")\nexcept Exception as e:\n    cell_status(\"CELL 20b — Architecture Diagram\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:46:53.076692Z","iopub.execute_input":"2026-08-02T05:46:53.077244Z","iopub.status.idle":"2026-08-02T05:46:53.624202Z","shell.execute_reply.started":"2026-08-02T05:46:53.077217Z","shell.execute_reply":"2026-08-02T05:46:53.623426Z"}},"outputs":[],"execution_count":null},{"id":"2c314c9f-fa41-4303-b0fc-5b741a57e9de","cell_type":"code","source":"# =====================================================================\n# CELL 21 — Stage 1 Training (Anatomical Localization Network) ONLY\n# =====================================================================\n# Split from the old combined driver cell: this cell trains Stage 1 exclusively. Stage 2 lives\n# in CELL 22, which loads the checkpoint saved here instead of retraining Stage 1.\nif torch.cuda.is_available():\n    print(f\"GPU memory allocated at Stage 1 start: {torch.cuda.memory_allocated() / 1e9:.2f} GB \"\n          f\"(if this is already several GB, restart the kernel before continuing)\")\nreport_gpu_status(prefix=\"[Driver] \")\n\ndf_train = pd.read_csv(os.path.join(CFG[\"data_root\"], \"train.csv\"))\ndf_series = pd.read_csv(os.path.join(CFG[\"data_root\"], \"train_series_descriptions.csv\"))\ndf_coords = pd.read_csv(os.path.join(CFG[\"data_root\"], \"train_label_coordinates.csv\"))\nprint(f\"Loaded {len(df_train)} studies, {len(df_series)} series rows, {len(df_coords)} labelled coordinates\")\n\nprint(\"\\n--- Running: Train Stage 1 (Anatomical Localization Network) ---\")\nloc_dataset = LocalizationDataset(df_coords, df_series)\nn_train = max(1, int(0.8 * len(loc_dataset)))\ntrain_ds1, val_ds1 = torch.utils.data.random_split(loc_dataset, [n_train, len(loc_dataset) - n_train]) \\\n    if len(loc_dataset) > 1 else (loc_dataset, loc_dataset)\n_loader_kwargs1 = make_fast_loader_kwargs(CFG[\"num_workers\"])\ntrain_loader1 = DataLoader(train_ds1, batch_size=CFG[\"batch_size_stage1\"], shuffle=True,\n                            num_workers=CFG[\"num_workers\"], **_loader_kwargs1)\nval_loader1 = DataLoader(val_ds1, batch_size=CFG[\"batch_size_stage1\"], shuffle=False,\n                          num_workers=CFG[\"num_workers\"], **_loader_kwargs1)\n\nstage1_model = AnatomicalLocalizationNet()\nhistory1 = None\ntry:\n    history1 = train_stage1(stage1_model, train_loader1, val_loader1, epochs=CFG[\"epochs_stage1\"])\nexcept Exception as e:\n    cell_status(\"CELL 21 — Train Stage 1 (training loop)\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n\nif history1 is not None and history1[\"train_acc\"]:\n    cell_status(\"CELL 21 — Train Stage 1 (training loop)\",\n                extra=f\"final train_acc={history1['train_acc'][-1]:.2%}, val_acc={history1['val_acc'][-1]:.2%}\")\nelse:\n    cell_status(\"CELL 21 — Train Stage 1 (training loop)\", ok=False, extra=\"history returned empty\")\n\ngc.collect()\ntorch.cuda.empty_cache()\nif torch.cuda.is_available():\n    print(f\"GPU memory allocated after Stage 1 + cleanup: {torch.cuda.memory_allocated() / 1e9:.2f} GB\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T05:47:01.690620Z","iopub.execute_input":"2026-08-02T05:47:01.691338Z","iopub.status.idle":"2026-08-02T06:35:01.413102Z","shell.execute_reply.started":"2026-08-02T05:47:01.691309Z","shell.execute_reply":"2026-08-02T06:35:01.410633Z"}},"outputs":[],"execution_count":null},{"id":"77bb8181-537d-4a7a-8c02-340afeb3a467","cell_type":"code","source":"# =====================================================================\n# CELL 21b -- Stage 1 Evaluation\n# Precision / Recall / F1-score / Confusion Matrix\n# (standalone cell between Stage 1 training and Stage 2 training)\n# =====================================================================\n# This cell can be run:\n#   A) immediately after Cell 21 finishes (stage1_model is already in memory), OR\n#   B) independently at any time -- it will load stage1_best.pt from disk and\n#      rebuild val_loader1 from the same fixed random seed, so results are\n#      reproducible and always refer to the same held-out validation split.\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import (\n    classification_report, confusion_matrix, accuracy_score,\n    ConfusionMatrixDisplay,\n)\n\nprint(\"=\" * 60)\nprint(\"STAGE 1 EVALUATION (validation set -- best checkpoint)\")\nprint(\"=\" * 60)\n\n# ---- 1) Ensure we have a model and a val loader -----------------\n_eval_model = None\nif \"stage1_model\" in globals() and isinstance(stage1_model, nn.Module):\n    print(\"Using stage1_model already in session memory.\")\n    _eval_model = stage1_model\nelse:\n    print(\"stage1_model not in memory -- loading from disk.\")\n    _eval_model = AnatomicalLocalizationNet()\n\n_best_ckpt = os.path.join(CFG[\"work_dir\"], \"stage1_best.pt\")\nif os.path.exists(_best_ckpt):\n    _eval_model.load_state_dict(torch.load(_best_ckpt, map_location=DEVICE))\n    _eval_model.to(DEVICE)\n    print(f\"Loaded best Stage 1 checkpoint: {_best_ckpt}\")\nelse:\n    raise FileNotFoundError(\n        f\"Best Stage 1 checkpoint not found at {_best_ckpt}. \"\n        f\"Run Cell 21 (Stage 1 training) first.\"\n    )\n\n# ---- 2) Rebuild val_loader1 if not in memory --------------------\nif \"val_loader1\" not in globals():\n    print(\"Rebuilding val_loader1 from dataset (same 80/20 split with fixed seed).\")\n    if \"df_coords\" not in globals():\n        _df_coords = pd.read_csv(os.path.join(CFG[\"data_root\"], \"train_label_coordinates.csv\"))\n        _df_series = pd.read_csv(os.path.join(CFG[\"data_root\"], \"train_series_descriptions.csv\"))\n    else:\n        _df_coords, _df_series = df_coords, df_series\n    _loc_dataset = LocalizationDataset(_df_coords, _df_series)\n    _n_train = max(1, int(0.8 * len(_loc_dataset)))\n    torch.manual_seed(42)   # same seed as the training split\n    _, _val_ds1 = torch.utils.data.random_split(_loc_dataset, [_n_train, len(_loc_dataset) - _n_train])\n    _val_loader1 = DataLoader(_val_ds1, batch_size=CFG[\"batch_size_stage1\"], shuffle=False,\n                               num_workers=CFG[\"num_workers\"], **make_fast_loader_kwargs(CFG[\"num_workers\"]))\nelse:\n    _val_loader1 = val_loader1\n    print(\"Using val_loader1 already in session memory.\")\n\n# ---- 3) Collect predictions -------------------------------------\n_eval_model.eval()\nall_preds, all_targets = [], []\nwith torch.no_grad():\n    for _batch in _val_loader1:\n        _vol = _batch[\"volume\"].to(DEVICE, non_blocking=True)\n        _presence = _batch[\"presence\"].to(DEVICE, non_blocking=True)\n        _out = _eval_model(_vol)\n        _preds = (torch.sigmoid(_out[\"instance_logits\"]) > 0.5).long().view(-1).cpu()\n        _targets = _presence.long().view(-1).cpu()\n        all_preds.extend(_preds.tolist())\n        all_targets.extend(_targets.tolist())\n\n# ---- 4) Print metrics -------------------------------------------\nprint(f\"\\nTotal validation samples (level-presence cells): {len(all_targets)}\")\nprint(f\"Overall Accuracy : {accuracy_score(all_targets, all_preds):.4f}\")\nprint(\"\\nClassification Report:\")\nprint(classification_report(\n    all_targets, all_preds,\n    labels=[0, 1],\n    target_names=[\"Absent (level not visible)\", \"Present (level visible)\"],\n    zero_division=0,\n))\n\n_cm = confusion_matrix(all_targets, all_preds, labels=[0, 1])\nprint(\"Confusion Matrix (rows = True class, cols = Predicted class):\")\nprint(f\"  {'':>28} {'Pred Absent':>12} {'Pred Present':>13}\")\nprint(f\"  {'True Absent':>28} {_cm[0,0]:>12d} {_cm[0,1]:>13d}\")\nprint(f\"  {'True Present':>28} {_cm[1,0]:>12d} {_cm[1,1]:>13d}\")\n\n# ---- 5) Plot confusion matrix -----------------------------------\n_cm_fig, _cm_ax = plt.subplots(figsize=(4, 3))\nConfusionMatrixDisplay(\n    confusion_matrix=_cm,\n    display_labels=[\"Absent\", \"Present\"],\n).plot(ax=_cm_ax, colorbar=False, cmap=\"Blues\")\n_cm_ax.set_title(\"Stage 1 -- Instance Presence\\nConfusion Matrix (validation set, best checkpoint)\")\n_cm_path = os.path.join(CFG[\"work_dir\"], \"stage1_confusion_matrix.png\")\n_cm_fig.tight_layout()\n_cm_fig.savefig(_cm_path, dpi=150)\nplt.show()\nprint(f\"\\nSaved: {_cm_path}\")\ncell_status(\"CELL 21b -- Stage 1 Evaluation\",\n            extra=f\"accuracy={accuracy_score(all_targets, all_preds):.4f}, {len(all_targets)} val cells evaluated\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T06:37:33.368418Z","iopub.execute_input":"2026-08-02T06:37:33.369015Z","iopub.status.idle":"2026-08-02T06:37:42.454362Z","shell.execute_reply.started":"2026-08-02T06:37:33.368976Z","shell.execute_reply":"2026-08-02T06:37:42.453791Z"}},"outputs":[],"execution_count":null},{"id":"d2d8ab0e-49a5-414a-bdd7-6343deba4311","cell_type":"code","source":"# =====================================================================\n# CELL 21c — Stage 1 Localization Verification (inline display)\n# For N_SAMPLES random val-set studies: predicted disc-level points\n# vs ground-truth annotations, shown directly in the notebook output.\n# =====================================================================\nimport matplotlib.pyplot as plt\n\nN_SAMPLES = 12   # change to any number 10-15\n\n# ---- Load model ----\n_viz_model = None\nif \"stage1_model\" in globals() and isinstance(stage1_model, nn.Module):\n    _viz_model = stage1_model\nelse:\n    _viz_model = AnatomicalLocalizationNet()\n\n_ckpt = os.path.join(CFG[\"work_dir\"], \"stage1_best.pt\")\nif not os.path.exists(_ckpt):\n    raise FileNotFoundError(f\"stage1_best.pt not found at {_ckpt}. Run Cell 21 first.\")\n_viz_model.load_state_dict(torch.load(_ckpt, map_location=DEVICE))\n_viz_model.to(DEVICE)\n_viz_model.eval()\nprint(f\"Loaded: {_ckpt}\")\n\n# ---- Load CSVs if needed ----\nif \"df_train\" not in globals():\n    df_train  = pd.read_csv(os.path.join(CFG[\"data_root\"], \"train.csv\"))\n    df_series = pd.read_csv(os.path.join(CFG[\"data_root\"], \"train_series_descriptions.csv\"))\n    df_coords = pd.read_csv(os.path.join(CFG[\"data_root\"], \"train_label_coordinates.csv\"))\n\n# ---- Sample from val split ----\nif \"val_ds1\" in globals():\n    _loc_ds = val_ds1.dataset\n    _indices = val_ds1.indices\nelse:\n    _loc_ds = LocalizationDataset(df_coords, df_series)\n    _indices = list(range(len(_loc_ds)))\n\ntorch.manual_seed(42)\n_chosen = torch.randperm(len(_indices))[:N_SAMPLES].tolist()\n_chosen_groups = [_loc_ds.groups[_indices[k]] for k in _chosen]\n\n_COLOURS   = [\"#e74c3c\",\"#e67e22\",\"#2ecc71\",\"#3498db\",\"#9b59b6\"]\n_LEVELS    = CFG[\"levels\"]\n_lv2idx    = {lvl: i for i, lvl in enumerate(_LEVELS)}\n\nprint(f\"Visualizing {N_SAMPLES} studies...\\n\")\n\nfor _si, ((study_id, series_id), _rows) in enumerate(_chosen_groups):\n\n    _volume, _instance_nums, _orig_hw = load_dicom_series(\n        study_id, series_id, CFG[\"data_root\"])\n    if _volume is None:\n        print(f\"  [{_si+1}] study {study_id}: no DICOMs — skipped\"); continue\n    _oh, _ow = _orig_hw\n    _vr = resample_depth(_volume, CFG[\"vol_size_3d\"][0])   # (D, 128, 128)\n    D, H, W = _vr.shape\n\n    _vol_t = torch.from_numpy(_vr).unsqueeze(0).unsqueeze(0).float().to(DEVICE)\n    with torch.no_grad():\n        _out = _viz_model(_vol_t)\n    _dp  = _out[\"depth_pred\"][0].cpu().numpy()    # (5,)\n    _yxp = _out[\"coord_pred\"][0].cpu().numpy()    # (5, 2)\n    _sig = torch.sigmoid(_out[\"instance_logits\"][0]).cpu().numpy()  # (5,)\n\n    # Ground-truth per level\n    _gt = {}\n    for _, _r in _rows.iterrows():\n        _lvl = str(_r[\"level\"]).lower().replace(\"/\",\"_\")\n        if _lvl not in _lv2idx: continue\n        _li = _lv2idx[_lvl]\n        _gt_y = float(_r[\"y\"]) / _oh * H\n        _gt_x = float(_r[\"x\"]) / _ow * W\n        _in   = int(_r[\"instance_number\"])\n        _d    = (_instance_nums.index(_in) / max(1,len(_instance_nums)-1)\n                 * (D-1)) if _in in _instance_nums else _dp[_li]*(D-1)\n        _gt[_li] = (int(_d), _gt_y, _gt_x)\n\n    # One row per level: [MRI with predicted X] [MRI with GT circle]\n    _fig, _axes = plt.subplots(len(_LEVELS), 2, figsize=(8, len(_LEVELS)*2.5))\n    _fig.suptitle(\n        f\"Study {study_id} · Series {series_id}  [{_si+1}/{N_SAMPLES}]\",\n        fontsize=10, fontweight=\"bold\")\n\n    for _li, _lvl in enumerate(_LEVELS):\n        _pd  = int(np.clip(_dp[_li]*(D-1), 0, D-1))\n        _py  = _yxp[_li,0]*H\n        _px  = _yxp[_li,1]*W\n        _c   = _COLOURS[_li]\n        _con = _sig[_li]\n        _prs = \"present\" if _con>0.5 else \"absent\"\n\n        ax0 = _axes[_li][0]\n        ax0.imshow(_vr[_pd], cmap=\"gray\", aspect=\"auto\")\n        ax0.scatter([_px],[_py], c=_c, s=100, marker=\"x\", linewidths=2.5)\n        ax0.set_title(\n            f\"{_lvl.upper().replace('_','/')}  Pred  slice={_pd}  \"\n            f\"conf={_con:.0%}  ({_prs})\", fontsize=7)\n        ax0.axis(\"off\")\n\n        ax1 = _axes[_li][1]\n        ax1.imshow(_vr[_pd], cmap=\"gray\", aspect=\"auto\")\n        ax1.scatter([_px],[_py], c=_c, s=60, marker=\"x\", linewidths=2, alpha=0.5)\n        if _li in _gt:\n            _gd, _gy, _gx = _gt[_li]\n            ax1.scatter([_gx],[_gy], c=\"cyan\", s=120, marker=\"o\",\n                        facecolors=\"none\", linewidths=2)\n            _dist = np.sqrt((_px-_gx)**2+(_py-_gy)**2)\n            _ok   = \"CLOSE ✓\" if _dist<15 else \"FAR ✗\"\n            ax1.set_title(\n                f\"{_lvl.upper().replace('_','/')}  GT  gt_slice={_gd}  \"\n                f\"dist={_dist:.1f}px  {_ok}\", fontsize=7)\n        else:\n            ax1.set_title(f\"{_lvl.upper().replace('_','/')}  GT  (no annotation)\", fontsize=7)\n        ax1.axis(\"off\")\n\n    _fig.tight_layout()\n    plt.show()   # ← inline display, no PDF\n    print(f\"  [{_si+1}/{N_SAMPLES}] study {study_id} ✓\")\n\ncell_status(\"CELL 21c — Stage 1 Localization Verification\",\n            extra=f\"{N_SAMPLES} studies shown inline\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T06:37:59.421148Z","iopub.execute_input":"2026-08-02T06:37:59.421786Z","iopub.status.idle":"2026-08-02T06:38:11.915754Z","shell.execute_reply.started":"2026-08-02T06:37:59.421759Z","shell.execute_reply":"2026-08-02T06:38:11.915003Z"}},"outputs":[],"execution_count":null},{"id":"3d4e3d2a-4a67-4626-a5ef-433ae0142e43","cell_type":"code","source":"# =====================================================================\n# CELL 22 — Stage 2 Training (Graph-based Severity Prediction) ONLY\n# =====================================================================\n# Automatically reuses stage1_model from CELL 21 if it is already in memory; otherwise loads the\n# saved checkpoint from disk. Stage 1 is NEVER retrained here.\nreport_gpu_status(prefix=\"[Driver] \")\n\nstage1_ckpt_path = os.path.join(CFG[\"work_dir\"], \"stage1_best.pt\")\nif \"stage1_model\" in globals() and isinstance(stage1_model, nn.Module):\n    print(\"Reusing Stage 1 model already trained in this session (no reload, no retraining).\")\nelif os.path.exists(stage1_ckpt_path):\n    print(f\"Loading previously trained Stage 1 model from {stage1_ckpt_path} (Stage 1 will NOT be retrained).\")\n    stage1_model = AnatomicalLocalizationNet()\n    stage1_model.load_state_dict(torch.load(stage1_ckpt_path, map_location=DEVICE))\n    stage1_model.to(DEVICE)\nelse:\n    raise RuntimeError(\n        f\"No trained Stage 1 model found in memory or on disk (expected {stage1_ckpt_path}). \"\n        f\"Run CELL 21 (Stage 1 training) first.\"\n    )\n\nif \"df_train\" not in globals():\n    df_train = pd.read_csv(os.path.join(CFG[\"data_root\"], \"train.csv\"))\n    df_series = pd.read_csv(os.path.join(CFG[\"data_root\"], \"train_series_descriptions.csv\"))\n    df_coords = pd.read_csv(os.path.join(CFG[\"data_root\"], \"train_label_coordinates.csv\"))\n    print(f\"Loaded {len(df_train)} studies, {len(df_series)} series rows, {len(df_coords)} labelled coordinates\")\n\nprint(\"\\n--- Running: Train Stage 2 (Graph-based Severity Prediction) ---\")\ntorch.cuda.empty_cache()\nsev_dataset = SeverityDataset(df_train, df_coords, max_slices_per_level=8)\n\n# PRE-WARM the disk cache: iterate over every study in the main process BEFORE handing control\n# to the DataLoader workers. Without this, every worker independently decodes DICOMs and writes\n# .npz files during the first training epoch, and all those intermediate numpy arrays pile up in\n# host RAM simultaneously (because Python GC can't evict them fast enough between batches), which\n# is the root cause of the host-RAM growth from 14 GB -> 30 GB -> OOM crash seen in debug logs.\n# After this loop every .npz file exists on disk; workers in epoch 1 will be cache hits (fast\n# numpy load) instead of cache misses (slow DICOM decode + numpy write), so RAM stays flat.\n_cache_dir = os.path.join(CFG[\"temp_dir\"], \"severity_roi_cache\")\n_already_cached = len([f for f in os.listdir(_cache_dir) if f.endswith(\".npz\")]) if os.path.isdir(_cache_dir) else 0\n_total_studies = len(df_train)\nif _already_cached >= _total_studies:\n    print(f\"Disk cache already complete ({_already_cached}/{_total_studies} studies). Skipping pre-warm.\")\nelse:\n    print(f\"Pre-warming disk cache ({_already_cached}/{_total_studies} already cached). \"\n          f\"This runs once and may take a few minutes -- epochs will be fast afterward.\")\n    import gc as _gc\n    for _pw_idx, _study_id in enumerate(df_train[\"study_id\"]):\n        sev_dataset.__getitem__(sev_dataset.study_ids.index(_study_id))\n        if _pw_idx % 100 == 0:\n            print(f\"  Pre-warm: {_pw_idx}/{_total_studies} studies cached \"\n                  f\"(host RAM: {host_ram_gb():.1f} GB)\")\n            _gc.collect()\n    _gc.collect()\n    print(f\"Pre-warm complete. {len([f for f in os.listdir(_cache_dir) if f.endswith('.npz')])} cache files written.\")\n\nn_train2 = max(1, int(0.8 * len(sev_dataset)))\ntrain_ds2, val_ds2 = torch.utils.data.random_split(sev_dataset, [n_train2, len(sev_dataset) - n_train2]) \\\n    if len(sev_dataset) > 1 else (sev_dataset, sev_dataset)\n# Stage 2's DataLoader workers do MUCH heavier per-item work than Stage 1's (DICOM decode +\n# crop + disk-cache write across up to 5 levels x multiple series per study), so the default\n# prefetch_factor=4 means up to num_workers*4 studies' worth of that work can be in flight in\n# host RAM at once -- while the GPU is busy in backward(), workers keep prefetching ahead,\n# compounding host RAM right when the silent \"tried to allocate more memory\" kernel death has\n# been happening. Reduced here specifically for Stage 2 (Stage 1's lighter workers keep the\n# shared default from make_fast_loader_kwargs).\n# Stage 2 DataLoaders use num_workers=0 (main-process loading).\n# ROOT CAUSE OF OOM: each DataLoader worker receives a full fork-time copy of the main\n# process heap (~14 GB at that point). With 3 workers + prefetch=2, forking happens 6 times\n# simultaneously right when SeverityGraphModel is also being constructed -- pushing RAM from\n# 14 GB to >30 GB. Since disk cache is pre-warmed, np.load() takes ~1-2 ms per study, so\n# num_workers=0 is fast enough and eliminates the RAM spike entirely.\n_loader_kwargs2 = {\"pin_memory\": torch.cuda.is_available()}\ntrain_loader2 = DataLoader(train_ds2, batch_size=CFG[\"batch_size_stage2\"], shuffle=True,\n                            num_workers=0, collate_fn=severity_collate_fn, **_loader_kwargs2)\nstage2_val_loader = DataLoader(val_ds2, batch_size=CFG[\"batch_size_stage2\"], shuffle=False,\n                                num_workers=0, collate_fn=severity_collate_fn, **_loader_kwargs2)\n\nimport gc as _gc2\n# Release large DataFrames from global scope before allocating the model.\n# sev_dataset already holds its own internal references, so __getitem__ still works.\n# We are only freeing the redundant global-scope copies that otherwise contribute to the\n# 14 GB baseline RAM at the moment SeverityGraphModel is constructed.\nfor _df_gbl in [\"df_coords\"]:\n    if _df_gbl in globals():\n        del globals()[_df_gbl]\n_gc2.collect()\nprint(f\"Host RAM before model construction: {host_ram_gb():.1f} GB  (after releasing df_coords global)\")\ntorch.cuda.empty_cache()\n\nstage2_model = SeverityGraphModel(backbone_name=\"tf_efficientnetv2_l\")\nprint(f\"Host RAM after model construction:  {host_ram_gb():.1f} GB\")\ntry:\n    history2 = train_stage2(stage2_model, train_loader2, stage2_val_loader, epochs=CFG[\"epochs_stage2\"])\n    cell_status(\"CELL 22 — Train Stage 2\",\n                extra=f\"final train_acc={history2['train_acc'][-1]:.2%}, val_acc={history2['val_acc'][-1]:.2%}\")\nexcept Exception as e:\n    cell_status(\"CELL 22 — Train Stage 2\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n\ngc.collect()\ntorch.cuda.empty_cache()\nif torch.cuda.is_available():\n    print(f\"GPU memory allocated after Stage 2 + cleanup: {torch.cuda.memory_allocated() / 1e9:.2f} GB\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T06:42:12.153185Z","iopub.execute_input":"2026-08-02T06:42:12.153863Z","iopub.status.idle":"2026-08-02T06:42:12.170246Z","shell.execute_reply.started":"2026-08-02T06:42:12.153831Z","shell.execute_reply":"2026-08-02T06:42:12.169241Z"}},"outputs":[],"execution_count":null},{"id":"0e188b7d-101b-4dc3-8627-ee9e6c3e0e62","cell_type":"code","source":"# =====================================================================\n# CELL 23 — Inference (Stage 1-5) on one example study + Export\n# =====================================================================\n@torch.no_grad()\ndef _collect_val_logits_targets(model, val_loader):\n    '''Gathers raw (uncalibrated) logits and targets for every valid (level, condition) cell in\n    the held-out Stage 2 validation set, for fitting TemperatureScaler on real data.'''\n    model.eval()\n    all_logits, all_targets = [], []\n    for batch in val_loader:\n        slice_stack = batch[\"slice_stack\"].to(DEVICE)\n        targets = batch[\"targets\"].to(DEVICE)\n        valid_mask = batch[\"valid_mask\"].to(DEVICE).bool()\n        out = model(slice_stack, batch[\"slice_counts\"])\n        all_logits.append(out[\"logits\"][valid_mask])\n        all_targets.append(targets[valid_mask])\n    if not all_logits:\n        return None, None\n    return torch.cat(all_logits, dim=0), torch.cat(all_targets, dim=0)\n\n\ntemp_scaler = TemperatureScaler().to(DEVICE)\nval_logits, val_targets = _collect_val_logits_targets(stage2_model, stage2_val_loader)\nif val_logits is not None and len(val_targets) > 0:\n    t_value = temp_scaler.fit(val_logits, val_targets)\n    print(f\"Temperature scaler fit on {len(val_targets)} validation cells -> T={t_value:.3f}\")\nelse:\n    print(\"No validation cells available to fit temperature scaler; using default T=1.5\")\n\ntry:\n    example_study_id = int(df_train[\"study_id\"].iloc[0])\n    example_series_ids = df_series[df_series[\"study_id\"] == example_study_id][\"series_id\"].tolist()\n    print(f\"Running full Stage 1-5 pipeline on study {example_study_id} ({len(example_series_ids)} series)...\")\n    result = run_full_pipeline(example_study_id, example_series_ids, stage1_model, stage2_model,\n                                temp_scaler, df_series=df_series)\n    print(\"\\n\" + result[\"report\"])\n    print(f\"\\nEvidence graph saved to: {result['evidence_graph_path']}\")\n    cell_status(\"CELL 23 — Inference\", extra=f\"study {example_study_id} processed\")\nexcept Exception as e:\n    cell_status(\"CELL 23 — Inference\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n\ntry:\n    zip_path = export_everything(results=result, history_stage1=history1, history_stage2=history2)\n    cell_status(\"CELL 23 — Export\", extra=f\"saved to {zip_path}\")\nexcept Exception as e:\n    cell_status(\"CELL 23 — Export\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"513bf234-eb02-44fe-a1a3-710865e6da17","cell_type":"code","source":"# =====================================================================\n# CELL 24 — Stage 2 evaluation metrics: accuracy, precision, recall, F1, confusion matrix\n# =====================================================================\nfrom sklearn.metrics import accuracy_score, precision_recall_fscore_support, confusion_matrix, classification_report\nimport matplotlib.pyplot as plt\n\n@torch.no_grad()\ndef evaluate_stage2(model, data_loader):\n    '''Runs the trained Stage 2 model over every batch in data_loader (STRICT_GPU_ONLY-enforced,\n    same as training) and collects (prediction, target) pairs for every (level, condition) cell\n    that has a valid label -- pooled overall, and split out per condition, so you can see whether\n    e.g. subarticular stenosis is harder than canal stenosis specifically.'''\n    model.eval()\n    all_preds, all_targets, all_conditions = [], [], []\n    for batch in data_loader:\n        slice_stack = batch[\"slice_stack\"].to(DEVICE)\n        assert_on_gpu(slice_stack, \"eval slice_stack\")\n        targets = batch[\"targets\"].to(DEVICE)\n        valid_mask = batch[\"valid_mask\"].to(DEVICE)\n\n        out = model(slice_stack, batch[\"slice_counts\"])\n        preds = out[\"logits\"].argmax(dim=-1)  # (B, L, C)\n\n        mask = valid_mask.cpu().numpy().astype(bool)\n        preds_np = preds.cpu().numpy()\n        targets_np = targets.cpu().numpy()\n        B, L, C = mask.shape\n        for b in range(B):\n            for l in range(L):\n                for c in range(C):\n                    if mask[b, l, c]:\n                        all_preds.append(int(preds_np[b, l, c]))\n                        all_targets.append(int(targets_np[b, l, c]))\n                        all_conditions.append(CFG[\"conditions\"][c])\n    return np.array(all_preds), np.array(all_targets), np.array(all_conditions)\n\n\ntry:\n    all_preds, all_targets, all_conditions = evaluate_stage2(stage2_model, stage2_val_loader)\n\n    if len(all_targets) == 0:\n        print(\"No labelled validation samples were found -- cannot compute metrics. \"\n              \"Make sure the Driver cell has finished a full training run first.\")\n        cell_status(\"CELL 24 — Stage 2 evaluation metrics\", ok=False, extra=\"no labelled validation samples available\")\n    else:\n        print(\"=\" * 60)\n        print(f\"OVERALL Stage 2 metrics ({len(all_targets)} labelled level/condition cells)\")\n        print(\"=\" * 60)\n        acc = accuracy_score(all_targets, all_preds)\n        precision, recall, f1, support = precision_recall_fscore_support(\n            all_targets, all_preds, labels=[0, 1, 2], average=None, zero_division=0)\n        macro_p, macro_r, macro_f1, _ = precision_recall_fscore_support(\n            all_targets, all_preds, labels=[0, 1, 2], average=\"macro\", zero_division=0)\n\n        print(f\"Accuracy: {acc:.4f}\\n\")\n        print(f\"{'Class':<15}{'Precision':>10}{'Recall':>10}{'F1':>10}{'Support':>10}\")\n        for i, name in enumerate(SEVERITY_NAMES):\n            print(f\"{name:<15}{precision[i]:>10.3f}{recall[i]:>10.3f}{f1[i]:>10.3f}{support[i]:>10d}\")\n        print(f\"{'Macro avg':<15}{macro_p:>10.3f}{macro_r:>10.3f}{macro_f1:>10.3f}{len(all_targets):>10d}\")\n        print()\n        print(classification_report(all_targets, all_preds, labels=[0, 1, 2],\n                                     target_names=SEVERITY_NAMES, zero_division=0))\n\n        cm = confusion_matrix(all_targets, all_preds, labels=[0, 1, 2])\n        fig, ax = plt.subplots(figsize=(5, 4))\n        ax.imshow(cm, cmap=\"Blues\")\n        ax.set_xticks(range(3)); ax.set_xticklabels(SEVERITY_NAMES, rotation=45, ha=\"right\")\n        ax.set_yticks(range(3)); ax.set_yticklabels(SEVERITY_NAMES)\n        ax.set_xlabel(\"Predicted\"); ax.set_ylabel(\"True\")\n        ax.set_title(\"Overall confusion matrix (all conditions pooled)\")\n        for i in range(3):\n            for j in range(3):\n                ax.text(j, i, str(cm[i, j]), ha=\"center\", va=\"center\",\n                        color=\"white\" if cm[i, j] > cm.max() / 2 else \"black\")\n        plt.tight_layout()\n        _cm_path = os.path.join(CFG[\"work_dir\"], \"confusion_matrix_overall.png\")\n        plt.savefig(_cm_path, dpi=150)\n        plt.show()\n        print(f\"Saved: {_cm_path}\")\n\n        print(\"\\n\" + \"=\" * 60)\n        print(\"PER-CONDITION metrics\")\n        print(\"=\" * 60)\n        for cond in CFG[\"conditions\"]:\n            cond_mask = all_conditions == cond\n            if cond_mask.sum() == 0:\n                continue\n            cond_targets = all_targets[cond_mask]\n            cond_preds = all_preds[cond_mask]\n            cond_acc = accuracy_score(cond_targets, cond_preds)\n            _, _, cond_f1, _ = precision_recall_fscore_support(\n                cond_targets, cond_preds, labels=[0, 1, 2], average=\"macro\", zero_division=0)\n            print(f\"\\n{cond.replace('_',' ').title()}  (n={cond_mask.sum()})\")\n            print(f\"  Accuracy: {cond_acc:.4f}  |  Macro F1: {cond_f1:.4f}\")\n\n        cell_status(\"CELL 24 — Stage 2 evaluation metrics\",\n                    extra=f\"accuracy={acc:.4f}, macro F1={macro_f1:.4f} on {len(all_targets)} val cells\")\nexcept NameError as e:\n    print(f\"Missing variable: {e}. This cell needs stage2_model and stage2_val_loader from the \"\n          f\"Driver cell -- run that cell first.\")\n    cell_status(\"CELL 24 — Stage 2 evaluation metrics\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"4d80ea82-d52d-4e3e-85d6-739198106c28","cell_type":"code","source":"# =====================================================================\n# CELL 25 — Demonstration on 3 unseen (held-out validation) MRI studies\n# =====================================================================\ndef demonstrate_on_unseen_studies(stage1_model, stage2_model, temp_scaler, df_series,\n                                    val_study_ids, n_demo=3, data_root=CFG[\"data_root\"]):\n    demo_ids = val_study_ids[:n_demo]\n    for study_id in demo_ids:\n        try:\n            print(\"\\n\" + \"=\" * 70)\n            print(f\"DEMONSTRATION — Study {study_id} (held-out / unseen during training)\")\n            print(\"=\" * 70)\n            series_ids = df_series[df_series[\"study_id\"] == study_id][\"series_id\"].tolist()\n            if not series_ids:\n                print(\"No series found for this study, skipping.\")\n                continue\n\n            # ---- 1) Original MRI: representative slice --------------------\n            volume, instance_numbers, orig_hw = load_dicom_series(study_id, series_ids[0], data_root)\n            if volume is not None:\n                mid = volume.shape[0] // 2\n                plt.figure(figsize=(4, 4))\n                plt.imshow(volume[mid], cmap=\"gray\")\n                plt.title(f\"Study {study_id} — representative MRI slice\")\n                plt.axis(\"off\")\n                plt.show()\n\n            # ---- 2) Full Stage 1-5 pipeline --------------------------------\n            result = run_full_pipeline(study_id, series_ids, stage1_model, stage2_model,\n                                        temp_scaler, df_series=df_series)\n\n            # ---- 3) ROI localization overlay (Stage 1 predicted points) ---\n            if volume is not None:\n                volume_r = resample_depth(volume, CFG[\"vol_size_3d\"][0])\n                vol_t = torch.from_numpy(volume_r).unsqueeze(0).unsqueeze(0).float().to(DEVICE)\n                with torch.no_grad():\n                    out1 = stage1_model(vol_t)\n                depth_pred = out1[\"depth_pred\"][0].cpu().numpy()\n                yx_pred = out1[\"coord_pred\"][0].cpu().numpy()\n                mid_r = int(np.clip(depth_pred.mean() * (volume_r.shape[0] - 1), 0, volume_r.shape[0] - 1))\n                plt.figure(figsize=(4, 4))\n                plt.imshow(volume_r[mid_r], cmap=\"gray\")\n                for li, lvl in enumerate(CFG[\"levels\"]):\n                    y = yx_pred[li, 0] * volume_r.shape[1]\n                    x = yx_pred[li, 1] * volume_r.shape[2]\n                    plt.scatter([x], [y], label=lvl.upper().replace(\"_\", \"/\"), s=40)\n                plt.legend(fontsize=6, loc=\"upper right\")\n                plt.title(\"ROI Localization — Stage 1 predicted anatomical points\")\n                plt.axis(\"off\")\n                plt.show()\n\n            # ---- 4) Severity classification table --------------------------\n            print(\"\\nSeverity Classification (per level, per condition):\")\n            for lvl in CFG[\"levels\"]:\n                for cond in CFG[\"conditions\"]:\n                    sev = SEVERITY_NAMES[result[\"severity_grid\"][lvl][cond]]\n                    conf = result[\"calibrated_confidence_grid\"][lvl][cond]\n                    print(f\"  {lvl.upper().replace('_','/'):<8} {cond.replace('_',' '):<35} : \"\n                          f\"{sev:<12} (confidence {conf:.1%})\")\n\n            # ---- 5) Explainability: Evidence Graph -------------------------\n            print(f\"\\nEvidence graph saved to: {result['evidence_graph_path']}\")\n            try:\n                _img = plt.imread(result[\"evidence_graph_path\"])\n                plt.figure(figsize=(5, 4))\n                plt.imshow(_img)\n                plt.axis(\"off\")\n                plt.title(\"Explainability — Evidence Graph\")\n                plt.show()\n            except Exception as e:\n                print(f\"  (could not display evidence graph inline: {e})\")\n\n            # ---- 6) Structured clinical report ------------------------------\n            print(\"\\n\" + result[\"report\"])\n\n        except Exception as e:\n            print(f\"  [WARN] Demonstration failed for study {study_id}: {type(e).__name__}: {e}\")\n            continue\n\n\ntry:\n    _val_subset = stage2_val_loader.dataset          # torch.utils.data.Subset\n    _val_study_ids = [_val_subset.dataset.study_ids[i] for i in _val_subset.indices]\n    demonstrate_on_unseen_studies(stage1_model, stage2_model, temp_scaler, df_series, _val_study_ids, n_demo=3)\n    cell_status(\"CELL 25 — Demonstration on unseen studies\", extra=f\"ran on {min(3, len(_val_study_ids))} held-out studies\")\nexcept Exception as e:\n    cell_status(\"CELL 25 — Demonstration on unseen studies\", ok=False, extra=f\"{type(e).__name__}: {e}\")\n    raise\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}