{"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":"bddb3764-6e60-4923-a926-fb9c693bc1af","cell_type":"markdown","source":"## RSNA Knee — EfficientNet-B3 2.5D Attention MIL\n\nAn efficient baseline solution, optimized to balance AUC and runtime according to the competition's Efficiency Score criterion. Uses an ImageNet-pretrained EfficientNet-B3 as the backbone to extract features per 2D slice, combined with attention pooling to aggregate spatial information across MRI slices — without the heavy compute cost of a true 3D CNN. Prioritizes the most informative sequences (sagittal PD/T2) to limit the number of views processed, keeping runtime safely within the 9-hour limit.\n\nBefore running: enable GPU and make sure a pretrained EfficientNet-B3 checkpoint (from timm) is attached offline as a Kaggle Dataset, since internet access is disabled during grading. The notebook automatically reads train.csv, test.csv, and series metadata to build the pipeline — DICOM decoding → slice caching → training → inference. The final output file is submission.csv.","metadata":{}},{"id":"d5c79cae-9701-41f0-ac2a-2ba4edab6a22","cell_type":"markdown","source":"## 1. Configuration, imports, and reproducibility","metadata":{}},{"id":"1355afea-a78f-407b-b209-73f809f28339","cell_type":"code","source":"from __future__ import annotations\n\nimport os\nimport re\nimport time\nimport random\nimport hashlib\nimport unicodedata\nimport warnings\nimport tempfile\nimport shutil\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import roc_auc_score\nfrom timm.utils import ModelEmaV2\nwarnings.filterwarnings(\"ignore\")\n\nROOT = Path(\"/kaggle/input/competitions/rsna-knee-abnormality-detection\")\nTRAIN_DICOM_DIR = ROOT / \"train_series\"\nTEST_DICOM_DIR = ROOT / \"test_series\"\n\nOUTPUT_DIR = Path(\"/kaggle/working\")\nOUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n\n# Decoded-slice cache: DICOM decode + physical crop + resize is the same for a\n# given study on every epoch and every fold, so it is done once and reused\n# instead of being recomputed ~30x (N_FOLDS x EPOCHS) over a full run.\n#\n# IMPORTANT: this cache is NOT written under OUTPUT_DIR (/kaggle/working).\n# /kaggle/working has a hard ~20GB output quota, and a full-corpus slice cache\n# (~4400 studies x 2 views x 12 slices x 224x224) can approach or exceed that on\n# its own, silently filling the disk and then crashing an unrelated later write\n# (e.g. torch.save of a model checkpoint). /kaggle/temp is Kaggle-provided scratch\n# space that is *not* counted toward the output quota and is wiped when the\n# session ends — exactly what a same-run cache needs. Falls back to the OS temp\n# dir if /kaggle/temp is not present (e.g. running this notebook outside Kaggle).\n_scratch_root = Path(\"/kaggle/temp\") if Path(\"/kaggle/temp\").exists() else Path(tempfile.gettempdir())\nCACHE_DIR = _scratch_root / \"slice_cache\"\nCACHE_DIR.mkdir(parents=True, exist_ok=True)\n\n# Purge any stale cache left behind by a previous (e.g. crashed) run in this\n# same session — otherwise its size counts against the fresh budget below even\n# though this run has not written anything yet.\nfor _f in CACHE_DIR.glob(\"*.npy\"):\n    try:\n        _f.unlink()\n    except Exception:\n        pass\n\n# Cached slices are quantized to uint8 (bag values are already normalized to\n# [0, 1] by decode_slice) instead of stored as float32 — a straight 4x size\n# reduction for a negligible precision loss the model will not notice.\nENABLE_SLICE_CACHE = True\n_cache_write_warned = False  # emit the \"cache disabled\" warning at most once\n_cache_bytes_written = 0\n\n# Do NOT assume which mount has spare capacity (/kaggle/working, /kaggle/temp,\n# and /tmp can share the same underlying writable disk on some Kaggle instance\n# types, so moving the cache path alone is not guaranteed to free anything up).\n# Instead, measure REAL free space where CACHE_DIR actually lives right now, and\n# cap total cache growth to that minus a safety margin, so later writes — model\n# checkpoints, submission.csv — always have guaranteed room regardless of the\n# platform's real disk layout.\nSAFETY_MARGIN_BYTES = 3 * 1024 ** 3  # keep 3GB free at all times\n_free_at_start = shutil.disk_usage(CACHE_DIR).free\nSLICE_CACHE_BYTE_BUDGET = max(0, _free_at_start - SAFETY_MARGIN_BYTES)\nprint(f\"[cache] {_free_at_start / 1e9:.1f}GB free at {CACHE_DIR}; \"\n      f\"cache budget={SLICE_CACHE_BYTE_BUDGET / 1e9:.1f}GB \"\n      f\"(keeping {SAFETY_MARGIN_BYTES / 1e9:.0f}GB headroom)\")\n\n# Series-index cache: header-only DICOM reads (SeriesDescription, geometry tags)\n# for every study, cached to disk so a kernel restart does not redo it.\nINDEX_CACHE_DIR = OUTPUT_DIR / \"study_index_cache\"\nINDEX_CACHE_DIR.mkdir(parents=True, exist_ok=True)\nINDEX_WORKERS = 16  # header reads are I/O-bound, so threads (not processes) help here\n\nTARGETS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n    \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\",\n    \"Synovitis\", \"Baker\\'s\", \"Contusion\", \"Fracture\",\n]\nN_TARGETS = len(TARGETS)\n\nSEED = 2026\nIMG_SIZE = 256  # up from 224 — small lesions (chondral fissures, tiny tears)\n                # get more pixels to be visible in; runtime headroom allows it\n                # (last full run used ~17% of the 8h budget)\nCROP_MM = 130.0        # physical crop before resize; keeps pixel pitch comparable\n                       # across studies with very different fields of view\nN_SLICES = 16          # up from 12 — more coverage per sequence, cheap given\n                       # the large unused time budget\nN_FOLDS = 5             # raised from 3 now that the pipeline has run clean\n                        # end-to-end (no NaN/OOM/disk crashes) — a full 5-fold\n                        # run still leaves most of the 8h budget unused\nBATCH_SIZE = 4        # studies per batch. Was bumped to 16 to feed 2 GPUs via\n                      # DataParallel (~8/GPU) — but DataParallel was later\n                      # disabled (it stalls on Kaggle's dual-T4, which has no\n                      # NVLink; see ENABLE_DATAPARALLEL below), leaving the\n                      # full batch of 16 landing on ONE GPU. At the current\n                      # 3 views x 16 slices x 256px, that's 768 images per\n                      # forward pass — ~10x the 96-image load of the last\n                      # config known to run without OOM. Back to 4.\nEVAL_BATCH_SIZE = 4   # kept in step with BATCH_SIZE for the same reason\nEPOCHS = 6\nLR_HEAD = 3e-4\nLR_BACKBONE = 3e-5    # ~100x below LR_HEAD: the encoder is nudged, not retrained\nWEIGHT_DECAY = 1e-3\n\n# After the main EPOCHS finish, fine-tune a couple more epochs on ONLY the gold\n# studies that fell in this fold's train split (~45-47 of the 58, the rest are in\n# this fold's val). The gold studies are already up-weighted (GOLD_WEIGHT_BOOST)\n# during main training, but they are still a small fraction of every batch — this\n# stage lets the model take a few steps informed by nothing else, at a much lower\n# LR so it nudges toward the gold labels rather than overfitting the ~46 studies.\n# Lowered from 2: most of a refine stage's benefit shows up in its first epoch;\n# a second epoch on the same ~46 gold-in-train studies adds a full extra\n# train+val pass per fold for a marginal, unmeasurable-at-this-n gain (see the\n# OOF gold AUC sigma note below main()). Raise back to 2 if runtime allows.\nEXACT_REFINE_EPOCHS = 1\nLR_REFINE_HEAD = LR_HEAD * 0.1\nLR_REFINE_BACKBONE = LR_BACKBONE * 0.1\n\n# Weak-label sample weight already ranges 0.15-1.0 by extraction confidence\n# (see build_labels), and gold rows get weight=1.0 — i.e. today a gold label\n# counts the SAME as the most confident weak label, even though it's the only\n# label source we actually trust. Boosting it makes the 58 gold studies pull\n# harder in both compute_pos_weight and the loss itself.\nGOLD_WEIGHT_BOOST = 3.0\n\n# Which targets use max-pooling across TTA variants instead of mean-pooling.\n# Fracture and Contusion (bone bruising) are often small and localized -- easy\n# for one scale/flip variant to catch clearly while others dilute it toward\n# ambiguous. Averaging suppresses a strong signal seen in only one variant; for\n# a small-and-localized finding, \"any variant flagged it clearly\" is the more\n# useful summary than \"the variants agree on average\". The other ten targets\n# (ligament tears, OA grading, effusion, ...) are more spatially diffuse and\n# mean-pooling's noise-averaging is the better fit there, so they keep it.\nTTA_TARGET_POOL = {\"Fracture\": \"max\", \"Contusion\": \"max\"}\n\n# Floor for per-target weak-label trust (see build_labels/compute_target_\n# reliability). trust = FLOOR + (1-FLOOR) * clip((reliability-0.5)*2, 0, 1),\n# so a target with gold-measured reliability at chance (AUC 0.5, e.g. Synovitis\n# was 0.56) gets pulled all the way down toward FLOOR instead of stopping at a\n# higher plateau. Lowered from 0.3 to 0.1 to let the weak/reliable targets\n# separate more sharply; kept above 0 so an unlucky small-sample reliability\n# estimate can't zero a target's weak-label gradient out entirely.\nTARGET_TRUST_FLOOR = 0.10\n\n# Exponential moving average of the weights, updated every training step. Usually\n# generalizes slightly better than the raw last-step weights because it averages\n# out step-to-step noise — cheap (one extra weighted-average op per step) and\n# does not change how anything is computed during training itself.\nUSE_EMA = True\nEMA_DECAY = 0.999\n\n# Focal-loss style modulation ((1 - p_t) ** FOCAL_GAMMA) applied on top of the\n# existing pos_weight/sample-weight scheme, to lean more on hard/rare examples\n# (Fracture, Baker's, ... have very few positives). 0 disables it (plain\n# weighted BCE).\nFOCAL_GAMMA = 1.5\nNUM_WORKERS = 4\nPERSISTENT_WORKERS = NUM_WORKERS > 0  # keep worker processes alive across epochs\n                                       # instead of respawning them every time\n                                       # (each spawn re-imports pydicom/cv2/timm)\nTIME_BUDGET_SECONDS = 8.0 * 3600\n\n# Multi-scale test-time augmentation: in addition to the existing flip, run\n# inference at a couple of extra scales (slight zoom-in/out via center-crop or\n# pad-then-resize) and average. Test-set size is small, so the extra passes\n# are cheap; this mainly smooths out scale-sensitivity in the attention pooling.\nTTA_SCALES = (1.0, 0.9, 1.1)\n\n# Anatomical prior for the per-target attention over views (Sagittal/Coronal/\n# Axial): which plane each finding is normally read on. exp(0.55) ~= 1.73x —\n# a soft nudge on where the softmax starts, not a hard restriction; the model\n# is free to learn away from it. Findings with no strong single-view preference\n# (Contusion, Fracture — visible on any fluid-sensitive sequence) are left at 0.\nVIEW_PRIOR_STRENGTH = 0.55\nTARGET_VIEW_PRIOR = {\n    \"ACL\":              {\"Sagittal\": VIEW_PRIOR_STRENGTH},\n    \"MCL\":               {\"Coronal\": VIEW_PRIOR_STRENGTH},\n    \"Medial Meniscus\":   {\"Sagittal\": VIEW_PRIOR_STRENGTH, \"Coronal\": VIEW_PRIOR_STRENGTH},\n    \"Lateral Meniscus\":  {\"Sagittal\": VIEW_PRIOR_STRENGTH, \"Coronal\": VIEW_PRIOR_STRENGTH},\n    \"Medial OA\":         {\"Coronal\": VIEW_PRIOR_STRENGTH},\n    \"Lateral OA\":        {\"Coronal\": VIEW_PRIOR_STRENGTH},\n    \"PF OA\":             {\"Axial\": VIEW_PRIOR_STRENGTH},\n    \"Effusion\":          {\"Sagittal\": VIEW_PRIOR_STRENGTH},\n    \"Synovitis\":         {\"Sagittal\": VIEW_PRIOR_STRENGTH, \"Axial\": VIEW_PRIOR_STRENGTH},\n    \"Baker's\":           {\"Sagittal\": VIEW_PRIOR_STRENGTH},\n    \"Contusion\":         {},\n    \"Fracture\":          {},\n}\n\n# Unfreezing the whole backbone made the epoch-2 backward pass OOM on a 96-image\n# mega-batch, so v2 kept the backbone almost entirely frozen (only 2 of 7 MBConv\n# stages). That dodged the OOM but caps what the head can be trained to see, since\n# ImageNet features don't describe MRI tissue contrast well. This opens 4 stages\n# instead, and leans on gradient checkpointing (below) rather than freezing to keep\n# memory affordable — checkpointing recomputes activations during backward instead\n# of storing them, trading compute for memory.\nUNFREEZE_LAST_STAGES = 4  # up from 2 — the run completed with no OOM at 2\n                          # stages, so this opens more of the backbone to adapt\n                          # to MRI contrast (EfficientNet-B3 has 7 stages total,\n                          # so 4 still leaves the earliest, most generic ones frozen)\nUSE_GRAD_CHECKPOINTING = True\n\nSEQUENCE_PRIORITY = [\n    (re.compile(r\"\\bsag\\b.*(pd|t2|fs)|\\b(pd|t2|fs)\\b.*\\bsag\\b\", re.I), \"sagittal_fluid\", \"Sagittal\"),\n    (re.compile(r\"\\bcor\\b.*(pd|t2|fs)|\\b(pd|t2|fs)\\b.*\\bcor\\b\", re.I), \"coronal_fluid\", \"Coronal\"),\n    # 3rd view: axial fluid-sensitive. Sagittal/coronal alone under-represent\n    # patellofemoral OA and some effusion/synovitis presentations that are\n    # clearer in the axial plane.\n    (re.compile(r\"\\bax\\b.*(pd|t2|fs)|\\b(pd|t2|fs)\\b.*\\bax\\b\", re.I), \"axial_fluid\", \"Axial\"),\n]\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nN_GPUS = torch.cuda.device_count()\n\nif torch.cuda.is_available():\n    # Input shapes are fixed every step (drop_last=True, fixed IMG_SIZE/N_SLICES),\n    # so cuDNN's autotuner can safely benchmark conv algorithms once and reuse the\n    # fastest one for the rest of the run — free speed, no effect on results.\n    torch.backends.cudnn.benchmark = True\n    # Allow TF32 for matmuls on Ampere+ GPUs; free speed, negligible precision\n    # cost for this task (autocast already uses fp16/bf16 for the heavy ops).\n    torch.set_float32_matmul_precision(\"high\")\n# USE_MULTI_GPU = N_GPUS > 1  # wraps the model in nn.DataParallel during training\n                            # when Kaggle gives this session 2 GPUs (e.g. 2x T4)\n\n\ndef set_seed(seed: int = SEED) -> None:\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n\n\nset_seed()\nT0 = time.time()\n\n\ndef log(msg: str) -> None:\n    print(f\"[{time.time() - T0:8.1f}s] {msg}\")\n\n\n# log(f\"device={DEVICE}; targets={N_TARGETS}; gpus={N_GPUS} \"\n#     f\"(multi-GPU training={'on' if USE_MULTI_GPU else 'off'})\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T18:11:51.610165Z","iopub.execute_input":"2026-08-07T18:11:51.610453Z","iopub.status.idle":"2026-08-07T18:12:07.242409Z","shell.execute_reply.started":"2026-08-07T18:11:51.610423Z","shell.execute_reply":"2026-08-07T18:12:07.24157Z"}},"outputs":[],"execution_count":null},{"id":"841f022b-f741-4ed0-8b2e-0e9cc890cada","cell_type":"markdown","source":"## 2. Report-derived weak labels\n\n`train.csv` has 4407 rows but only 58 carry the twelve structured labels; the rest\nare `NaN` there. Every row does carry a free-text `Report`. Training on the raw\ncolumns means almost every batch has a `NaN` target, which poisons the loss and then\nthe weights from batch one (this is exactly what produced `train_loss=nan` on epoch\n1 in v1).\n\nThis cell reads each report clause by clause and decides, for each of the twelve\nfindings, whether the clause asserts it, negates it, or hedges it — across the\nlanguages actually present in the corpus (the sample reports seen so far are\nSpanish, but a few other Latin-script languages are included cheaply since they cost\nnothing when absent and other RSNA-style multi-site corpora do mix languages).\nA silent lexicon (no rule fires) is safer than a wrong one: it emits low confidence\nrather than a guessed label, and that confidence becomes the sample weight used in\ntraining. The 58 structured rows, where present, override the extracted score and\nget full weight.","metadata":{}},{"id":"568d12f9-4268-459a-be1b-9b95463b0d6f","cell_type":"code","source":"_PRE = str.maketrans({\n    \"\\u0131\": \"i\", \"\\u0130\": \"i\", \"\\u00df\": \"ss\", \"\\u0111\": \"d\", \"\\u0110\": \"d\",\n    \"\\u00f8\": \"o\", \"\\u00d8\": \"o\", \"\\u00e6\": \"ae\", \"\\u00c6\": \"ae\",\n})\n\n\ndef normalize(text: str) -> str:\n    \"\"\"Fold case, diacritics and separators.\"\"\"\n    if not isinstance(text, str):\n        return \"\"\n    text = text.translate(_PRE).lower()\n    text = unicodedata.normalize(\"NFKD\", text)\n    text = \"\".join(ch for ch in text if not unicodedata.combining(ch))\n    text = re.sub(r\"[_\\-/\\\\\\\\]+\", \" \", text)\n    text = re.sub(r\"[ \\t]+\", \" \", text)\n    return text\n\n\n_SENT_SPLIT = re.compile(r\"(?<=[.;!?])\\s+|\\n+\")\n\n\ndef clauses(text: str):\n    \"\"\"Split into clauses; attach a trailing `header:` fragment to the value after it,\n    so `Fracturas:` followed by `Ninguna.` reads as one statement rather than two.\"\"\"\n    norm = normalize(text)\n    raw = [c.strip() for c in _SENT_SPLIT.split(norm) if c and c.strip()]\n    merged = []\n    for i, c in enumerate(raw):\n        if c.endswith(\":\") and len(c.split()) <= 14 and i + 1 < len(raw):\n            merged.append(c + \" \" + raw[i + 1])\n        merged.append(c)\n    out = []\n    for c in merged:\n        out.append(c)\n        if len(c.split()) > 25:\n            out.extend(p.strip() for p in c.split(\",\") if len(p.split()) > 2)\n    return out\n\n\ndef _rx(*alts: str) -> re.Pattern:\n    return re.compile(\"|\".join(alts))\n\n\nNEGATION = _rx(\n    r\"\\bno\\b\", r\"\\bsin\\b\", r\"\\bno hay\\b\", r\"\\bausencia\\b\", r\"\\bausentes?\\b\",\n    r\"\\bnot\\b\", r\"\\bwithout\\b\", r\"\\bno evidence\\b\", r\"\\bunremarkable\\b\",\n    r\"\\bpas de\\b\", r\"\\bsans\\b\", r\"\\baucune?\\b\",\n)\n\nNORMALITY = _rx(\n    r\"\\bnormal\", r\"\\bintact\\b\", r\"\\bpreserved\\b\", r\"limites normales\",\n    r\"\\bconservad\", r\"\\bintegr\", r\"\\bnormales\\b\", r\"within normal limits\",\n)\n\nUNCERTAIN = _rx(\n    r\"\\bpossible\\b\", r\"\\bprobable\\b\", r\"\\bsuspicious\\b\", r\"\\bposible\\b\",\n    r\"sin criterios categoricos\", r\"\\bdudos\", r\"cannot (be )?exclude\", r\"\\bmay\\b\",\n)\n\nTEAR = _rx(\n    r\"\\btear\", r\"\\btorn\\b\", r\"\\brupture\", r\"discontinuit\",\n    r\"\\brotura\\b\", r\"\\broturas\\b\", r\"\\bruptura\", r\"\\bdesgarro\", r\"\\broto\\b\",\n    r\"\\bdechirure\", r\"\\bdechire\",\n)\n\nDEGEN = _rx(\n    r\"degenerat\", r\"\\bmucoid\\b\", r\"\\bmyxoid\\b\", r\"\\bfray\", r\"\\bfissur\",\n)\n\nINJURY = _rx(\n    r\"\\binjur\", r\"\\bsprain\", r\"\\blesion\", r\"\\bedema\\b\", r\"\\boedema\\b\",\n    r\"aumento de senal\", r\"alteracion de senal\", r\"cambio de senal\",\n    r\"\\bhigh signal\\b\", r\"\\bhiperintens\", r\"\\bhyperintens\", r\"\\bpartial\\b\", r\"\\bparcial\",\n)\n\nANAT = {\n    \"ACL\": _rx(\n        r\"anterior cruciate\", r\"\\bacl\\b\", r\"cruzado anterior\", r\"\\blca\\b\",\n        r\"croise anterieur\", r\"cruciate ligaments\", r\"ligamentos cruzados\",\n    ),\n    \"MCL\": _rx(\n        r\"medial collateral\", r\"\\bmcl\\b\", r\"colateral medial\", r\"colateral interno\",\n        r\"\\blcm\\b\", r\"collateral medial\", r\"\\bcolaterales\\b\", r\"ligamentos colaterales\",\n    ),\n    \"Medial Meniscus\": _rx(\n        r\"medial meniscus\", r\"medial menisc\", r\"menisco medial\", r\"menisco interno\",\n        r\"menisque medial\",\n    ),\n    \"Lateral Meniscus\": _rx(\n        r\"lateral meniscus\", r\"lateral menisc\", r\"menisco lateral\", r\"menisco externo\",\n        r\"menisque lateral\",\n    ),\n}\n\nOA_EVIDENCE = _rx(\n    r\"osteoarthrit\", r\"\\barthros\", r\"\\bgonarthros\", r\"chondropath\", r\"chondromalac\",\n    r\"condropat\", r\"condromalac\", r\"cartilage loss\", r\"chondral (loss|defect|ulcer|thinning)\",\n    r\"osteophyt\", r\"osteofit\", r\"joint space narrowing\", r\"pinzamiento articular\",\n    r\"ulcera[s]? condral\", r\"cartilago[^.]{0,25}(perdida|adelgaz)\",\n)\n\nCOMPARTMENT = {\n    \"Medial OA\": _rx(\n        r\"medial (femorotibial|tibiofemoral|compartment)\",\n        r\"compartimento femorotibial medial\", r\"femorotibial interno\",\n        r\"medial (femoral|tibial) (condyle|plateau)\", r\"condilo femoral medial\",\n    ),\n    \"Lateral OA\": _rx(\n        r\"lateral (femorotibial|tibiofemoral|compartment)\",\n        r\"compartimento femorotibial lateral\", r\"femorotibial externo\",\n        r\"lateral (femoral|tibial) (condyle|plateau)\", r\"condilo femoral lateral\",\n    ),\n    \"PF OA\": _rx(\n        r\"patellofemoral\", r\"femoropatellar\", r\"femoropatelar\", r\"patelofemoral\",\n        r\"retropatellar\", r\"retrorotulian\", r\"\\btrochlea\", r\"\\btroclea\",\n        r\"\\bpatella\\b\", r\"\\bpatellar\\b\", r\"\\brotulian\", r\"\\brotula\\b\",\n    ),\n}\n\nDIRECT = {\n    \"Effusion\": _rx(\n        r\"\\beffusion\", r\"joint fluid\", r\"derrame articular\", r\"\\bderrame\\b\",\n        r\"liquido articular\", r\"epanchement\",\n    ),\n    \"Synovitis\": _rx(\n        r\"synovit\", r\"sinovit\", r\"synovial (thickening|proliferation|hypertroph)\",\n    ),\n    \"Baker\\'s\": _rx(\n        r\"baker\", r\"popliteal cyst\", r\"quiste popliteo\", r\"quistes popliteos\",\n        r\"kyste poplite\",\n    ),\n    \"Contusion\": _rx(\n        r\"\\bcontusion\", r\"bone bruise\", r\"bone marrow (o?edema|contusion)\",\n        r\"contusion osea\", r\"edema oseo\", r\"edema de medula osea\",\n    ),\n    \"Fracture\": _rx(\n        r\"\\bfractur\", r\"\\bfract\\b\", r\"\\bfractura\", r\"\\bfracturas\\b\",\n        r\"\\bfraktur\", r\"\\bbreuk\\b\", r\"insufficiency fracture\", r\"stress fracture\",\n        r\"avulsion fracture\", r\"subchondral fracture\",\n    ),\n}\n\nDECOY = {\n    \"Fracture\": _rx(r\"no fracture\", r\"microfractur\"),\n    \"Baker\\'s\": _rx(r\"meniscal cyst\", r\"quiste meniscal\", r\"ganglion\"),\n}\n\nPAIRED = {\"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\"}\nOA_TARGETS = {\"Medial OA\", \"Lateral OA\", \"PF OA\"}\n\nGLOBAL_OA = _rx(\n    r\"tri ?compartment\", r\"all three compartment\", r\"global(ised)? (oa|osteoarthrit)\",\n    r\"\\bgonarthros\", r\"osteoarthritis of the knee\", r\"artrosis (de |)(la )?rodilla\",\n    r\"knee osteoarthrit\", r\"degenerative joint disease\", r\"\\bdjd\\b\",\n)\n\nSEV_LOW = _rx(\n    r\"\\bsmall\\b\", r\"\\bminimal\\b\", r\"\\btrace\\b\", r\"\\bmild\\b\", r\"\\bslight\\b\",\n    r\"\\bleve\\b\", r\"\\bminim\", r\"\\bpeque\", r\"\\bligero\\b\", r\"\\bescaso\\b\", r\"\\bdiscreto\\b\",\n)\n\nSEV_HIGH = _rx(\n    r\"\\blarge\\b\", r\"\\bmarked\\b\", r\"\\bmassive\\b\", r\"\\bsevere\\b\", r\"\\bextensive\\b\",\n    r\"\\bmoderate\\b\", r\"\\bmoderad\", r\"\\bimportante\\b\", r\"\\bsevera?\\b\", r\"\\bmarcad\",\n)\n\n\ndef _polarity(clause: str) -> str:\n    if UNCERTAIN.search(clause):\n        return \"uncertain\"\n    if NEGATION.search(clause):\n        return \"negative\"\n    if NORMALITY.search(clause):\n        if TEAR.search(clause) or re.search(r\"\\bgrade [34]\\b\", clause):\n            return \"positive\"\n        return \"negative\"\n    return \"positive\"\n\n\ndef _severity(clause: str) -> float:\n    high = SEV_HIGH.search(clause) is not None\n    low = SEV_LOW.search(clause) is not None\n    if high and not low:\n        return 1.0\n    if low and not high:\n        return 0.45\n    return 0.75\n\n\ndef _score_clauses(cls, anat_rx, path_rx=None, decoy_rx=None):\n    n_pos = n_neg = n_unc = 0\n    best = 0.0\n    for c in cls:\n        m = anat_rx.search(c)\n        if not m:\n            continue\n        if decoy_rx is not None and decoy_rx.search(c):\n            continue\n        if path_rx is not None and not path_rx.search(c):\n            if NORMALITY.search(c) and not NEGATION.search(c):\n                n_neg += 1\n            continue\n        pol = _polarity(c)\n        if pol == \"positive\":\n            n_pos += 1\n            best = max(best, _severity(c))\n        elif pol == \"negative\":\n            n_neg += 1\n        else:\n            n_unc += 1\n            best = max(best, 0.30)\n\n    if n_pos or n_unc:\n        score = min(0.95, 0.50 + 0.42 * best + 0.03 * min(n_pos, 3))\n        conf = min(1.0, 0.55 + 0.15 * n_pos)\n    elif n_neg:\n        score = max(0.04, 0.20 - 0.04 * n_neg)\n        conf = min(0.9, 0.45 + 0.12 * n_neg)\n    else:\n        score, conf = 0.28, 0.05\n    return score, conf, n_pos, n_neg\n\n\ndef extract_report_labels(report: str) -> dict:\n    \"\"\"Twelve (score, confidence) pairs read from one free-text report.\"\"\"\n    cls = clauses(report)\n    out = {}\n    path_paired = _rx(TEAR.pattern, DEGEN.pattern, INJURY.pattern)\n\n    for tgt in TARGETS:\n        if tgt in PAIRED:\n            s, c, npos, nneg = _score_clauses(cls, ANAT[tgt], path_paired)\n        elif tgt in OA_TARGETS:\n            s, c, npos, nneg = _score_clauses(cls, COMPARTMENT[tgt], OA_EVIDENCE)\n        else:\n            s, c, npos, nneg = _score_clauses(cls, DIRECT[tgt], None, DECOY.get(tgt))\n        out[tgt] = s\n        out[tgt + \"__conf\"] = c\n        out[tgt + \"__npos\"] = npos\n        out[tgt + \"__nneg\"] = nneg\n\n    g_hits = [c for c in cls if GLOBAL_OA.search(c) and _polarity(c) == \"positive\"]\n    if g_hits:\n        gscore = 0.50 + 0.42 * max(_severity(c) for c in g_hits)\n        for tgt in OA_TARGETS:\n            if out[tgt + \"__npos\"] == 0 and out[tgt + \"__nneg\"] == 0:\n                out[tgt] = max(out[tgt], gscore * 0.92)\n                out[tgt + \"__conf\"] = max(out[tgt + \"__conf\"], 0.4)\n\n    if out[\"Synovitis__npos\"] == 0 and out[\"Synovitis__nneg\"] == 0:\n        out[\"Synovitis\"] = max(out[\"Synovitis\"], 0.28 + 0.45 * (out[\"Effusion\"] - 0.28))\n\n    return out\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T18:12:07.243951Z","iopub.execute_input":"2026-08-07T18:12:07.244462Z","iopub.status.idle":"2026-08-07T18:12:07.273183Z","shell.execute_reply.started":"2026-08-07T18:12:07.244439Z","shell.execute_reply":"2026-08-07T18:12:07.272629Z"}},"outputs":[],"execution_count":null},{"id":"874e208a-f397-4a86-a79d-d031fa545b00","cell_type":"markdown","source":"## 3. Series discovery, geometry-based slice ordering, and sequence selection\n\nA second, in-memory cache sits in front of the per-study disk cache (see the\ndecode cell below): if the whole corpus's decoded bags fit comfortably in\navailable RAM, they are kept there instead, skipping per-item disk I/O\nentirely. This is an opportunistic upgrade, decided automatically at runtime\nby reading `/proc/meminfo` — if it doesn't fit with a safety margin, the\npipeline falls back to exactly the disk-cache behavior that already works,\nwith no other change.\n\nView selection now prefers the ground-truth `Anatomical_Plane` / `Fluid_Sensitive`\ncolumns from `train_series.csv` / `test_series.csv` (curator-verified) over\nguessing from free-text `SeriesDescription`. A real example from this dataset —\n`\"DP SPIR CS_SAG\"` — is a sagittal proton-density fat-saturated sequence that the\nregex below fails to recognise at all: `\"CS_SAG\"` has no word boundary before\n`\"sag\"` (an underscore is a word character), and `\"DP\"` is the reverse of the\npattern's `\"pd\"`. The regex is kept as a fallback for studies missing a lookup row\n(and for the test set, whose `test_series.csv` is a placeholder until grading\ntime per the dataset description) — not removed, since it still catches real\ncases the lookup doesn't cover.\n\nSame 2-view slot scheme as v1 (sagittal fluid-sensitive, coronal fluid-sensitive),\nbut two corrections:\n\n- **Slices are ordered by physical position, not filename.** A DICOM file name is a\n  SOP Instance UID — assigned to be unique, not to be ordered. Sorting by it gives a\n  sequence essentially uncorrelated with anatomy, which silently breaks \"adjacent\n  slices as 2.5D context\" and \"sample the middle of the stack\". The true order comes\n  from projecting `ImagePositionPatient` onto the slice normal\n  (`ImageOrientationPatient` cross product).\n- **A physical crop is applied before resizing**, using `PixelSpacing`, so a 1-3 mm\n  finding does not get destroyed by resampling a field of view that varies several-fold\n  across the corpus down to a fixed pixel grid.","metadata":{}},{"id":"4479ad44-9f05-4b3f-a86a-a0389d68b1bc","cell_type":"code","source":"import pickle\nfrom concurrent.futures import ThreadPoolExecutor, as_completed\n\n\ndef _order_slices(files: list[Path]) -> list[Path]:\n    \"\"\"Sort a series' files along the through-plane axis using DICOM geometry.\n\n    Falls back to InstanceNumber, then to the existing (filename) order, if the\n    geometry tags are missing on some or all slices.\n    \"\"\"\n    keyed = []\n    any_geom = False\n    for f in files:\n        k = None\n        try:\n            ds = pydicom.dcmread(f, stop_before_pixels=True, force=True,\n                                  specific_tags=[\"ImagePositionPatient\", \"ImageOrientationPatient\",\n                                                 \"InstanceNumber\"])\n            iop = np.asarray(ds.ImageOrientationPatient, dtype=float)\n            ipp = np.asarray(ds.ImagePositionPatient, dtype=float)\n            normal = np.cross(iop[:3], iop[3:])\n            k = float(np.dot(ipp, normal))\n            any_geom = True\n        except Exception:\n            try:\n                k = float(ds.InstanceNumber)\n            except Exception:\n                k = None\n        keyed.append((k, f))\n    if not any_geom or any(k is None for k, _ in keyed):\n        # Missing geometry on some slices: keep the safer of the two orders rather\n        # than mixing a geometry-sorted prefix with an arbitrary tail.\n        valid = [(k, f) for k, f in keyed if k is not None]\n        if len(valid) >= max(2, int(0.8 * len(keyed))):\n            return [f for _, f in sorted(valid, key=lambda t: t[0])] + \\\n                   [f for k, f in keyed if k is None]\n        return files\n    return [f for _, f in sorted(keyed, key=lambda t: t[0])]\n\n\ndef load_series_lookup(csv_path: Path) -> dict[str, dict[str, dict]]:\n    \"\"\"StudyInstanceUID -> SeriesInstanceUID -> {\"plane\": ..., \"fluid\": ...},\n    read from train_series.csv / test_series.csv. These give the imaging plane\n    and fluid-sensitivity directly (curator-verified) rather than needing to be\n    guessed from SeriesDescription free text.\n    \"\"\"\n    if not csv_path.exists():\n        return {}\n    df = pd.read_csv(csv_path)\n    lookup: dict[str, dict[str, dict]] = {}\n    for row in df.itertuples(index=False):\n        study_lu = lookup.setdefault(row.StudyInstanceUID, {})\n        fluid = bool(row.Fluid_Sensitive) if pd.notna(row.Fluid_Sensitive) else None\n        plane = str(row.Anatomical_Plane) if pd.notna(row.Anatomical_Plane) else None\n        study_lu[row.SeriesInstanceUID] = {\"plane\": plane, \"fluid\": fluid}\n    return lookup\n\n\ndef list_series(study_dir: Path, series_lookup: dict | None = None) -> list[dict]:\n    records = []\n    if not study_dir.exists():\n        return records\n    study_lookup = series_lookup.get(study_dir.name, {}) if series_lookup else {}\n    for series_dir in sorted(p for p in study_dir.iterdir() if p.is_dir()):\n        dcm_files = sorted(series_dir.glob(\"*.dcm\"))\n        if not dcm_files:\n            continue\n        description = \"\"\n        laterality = None\n        px_spacing = None\n        try:\n            ds = pydicom.dcmread(dcm_files[0], stop_before_pixels=True, force=True)\n            description = str(getattr(ds, \"SeriesDescription\", \"\") or \"\")\n            lat_raw = str(getattr(ds, \"Laterality\", \"\") or \"\").strip().upper()\n            laterality = lat_raw[0] if lat_raw and lat_raw[0] in (\"L\", \"R\") else None\n            ps = getattr(ds, \"PixelSpacing\", None)\n            if ps is not None and len(ps) > 0:\n                px_spacing = float(ps[0])\n        except Exception:\n            pass\n        gt = study_lookup.get(series_dir.name)\n        records.append({\n            \"series_dir\": series_dir,\n            \"n_files\": len(dcm_files),\n            \"files\": dcm_files,\n            \"description\": description,\n            \"laterality\": laterality,\n            \"px_spacing\": px_spacing,\n            \"gt_plane\": gt[\"plane\"] if gt else None,\n            \"gt_fluid\": gt[\"fluid\"] if gt else None,\n        })\n    return records\n\n\ndef select_sequences(series_records: list[dict],\n                      max_views: int = len(SEQUENCE_PRIORITY)) -> list[dict]:\n    chosen = []\n    used_dirs = set()\n\n    def try_pick(matches_fn):\n        for rec in series_records:\n            if rec[\"series_dir\"] in used_dirs:\n                continue\n            if matches_fn(rec):\n                return rec\n        return None\n\n    for pattern, _tag, plane in SEQUENCE_PRIORITY:\n        # Ground truth first (see cell above for why); regex only as a fallback\n        # for studies the lookup doesn't cover.\n        rec = try_pick(lambda r, plane=plane: r.get(\"gt_plane\") == plane and r.get(\"gt_fluid\") is True)\n        if rec is None:\n            rec = try_pick(lambda r, pattern=pattern: pattern.search(r[\"description\"]))\n        if rec is not None:\n            rec = dict(rec)\n            rec[\"plane\"] = plane\n            chosen.append(rec)\n            used_dirs.add(rec[\"series_dir\"])\n        if len(chosen) >= max_views:\n            break\n\n    if not chosen and series_records:\n        fallback = max(series_records, key=lambda r: r[\"n_files\"])\n        fallback = dict(fallback)\n        fallback[\"plane\"] = \"Sagittal\"\n        chosen.append(fallback)\n\n    for rec in chosen:\n        rec[\"files\"] = _order_slices(rec[\"files\"])\n\n    return chosen[:max_views]\n\n\ndef _index_one_study(dicom_root: Path, study_id: str,\n                      series_lookup: dict | None = None) -> tuple[str, list[dict]]:\n    series = list_series(dicom_root / study_id, series_lookup)\n    return study_id, select_sequences(series)\n\n\ndef build_study_index(dicom_root: Path, study_ids: list[str],\n                       max_workers: int = INDEX_WORKERS,\n                       cache_name: str | None = None,\n                       series_lookup: dict | None = None) -> dict[str, list[dict]]:\n    \"\"\"Header-only indexing (SeriesDescription, geometry tags) is I/O-bound, not\n    CPU-bound, so a thread pool (not a process pool) speeds it up without paying\n    multiprocessing/pickling overhead. The full index is also cached to disk under\n    `cache_name`, so re-running the notebook after a kernel restart does not repeat\n    ~40 minutes of header reads.\n    \"\"\"\n    cache_path = INDEX_CACHE_DIR / f\"{cache_name}.pkl\" if cache_name else None\n    if cache_path is not None and cache_path.exists():\n        try:\n            with open(cache_path, \"rb\") as f:\n                index = pickle.load(f)\n            log(f\"loaded cached index '{cache_name}' ({len(index)} studies)\")\n            return index\n        except Exception:\n            pass  # corrupt cache file, fall through and rebuild it\n\n    index = {}\n    done = 0\n    with ThreadPoolExecutor(max_workers=max_workers) as ex:\n        futures = [ex.submit(_index_one_study, dicom_root, sid, series_lookup) for sid in study_ids]\n        for fut in as_completed(futures):\n            study_id, seqs = fut.result()\n            index[study_id] = seqs\n            done += 1\n            if done % 500 == 0:\n                log(f\"  indexed {done}/{len(study_ids)} studies\")\n\n    if cache_path is not None:\n        try:\n            with open(cache_path, \"wb\") as f:\n                pickle.dump(index, f)\n        except Exception:\n            pass  # e.g. read-only or full disk — training still proceeds normally\n\n    return index\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T18:12:07.274164Z","iopub.execute_input":"2026-08-07T18:12:07.274597Z","iopub.status.idle":"2026-08-07T18:12:07.295713Z","shell.execute_reply.started":"2026-08-07T18:12:07.274557Z","shell.execute_reply":"2026-08-07T18:12:07.295025Z"}},"outputs":[],"execution_count":null},{"id":"fedb1462-5e01-437f-a11f-e751e90585fe","cell_type":"markdown","source":"## 4. DICOM decoding, physical crop, laterality normalization, and slice sampling","metadata":{}},{"id":"046f05ab-300e-4329-9f52-9d73fc75103b","cell_type":"code","source":"import cv2\nfrom concurrent.futures import ThreadPoolExecutor\n\n\ndef decode_slice(path: Path, crop_mm: float = CROP_MM, out_size: int = IMG_SIZE) -> np.ndarray:\n    ds = pydicom.dcmread(path, force=True)\n    arr = ds.pixel_array.astype(np.float32)\n\n    slope = float(getattr(ds, \"RescaleSlope\", 1.0) or 1.0)\n    intercept = float(getattr(ds, \"RescaleIntercept\", 0.0) or 0.0)\n    arr = arr * slope + intercept\n\n    arr = np.nan_to_num(arr, nan=0.0, posinf=0.0, neginf=0.0)\n\n    # Padding pixels (scanner-reported blank border) should not enter the intensity\n    # percentiles below, or a large padded border silently compresses the real range.\n    valid = np.ones_like(arr, dtype=bool)\n    if hasattr(ds, \"PixelPaddingValue\"):\n        try:\n            padding = int(ds.PixelPaddingValue)\n            valid = arr != (padding * slope + intercept)\n        except Exception:\n            pass\n\n    if str(getattr(ds, \"PhotometricInterpretation\", \"\")).upper() == \"MONOCHROME1\":\n        if valid.any():\n            arr[valid] = arr[valid].max() - arr[valid] + arr[valid].min()\n\n    # Physical crop: PixelSpacing varies several-fold across this corpus, so a fixed\n    # pixel-count crop covers a different amount of anatomy per study. Cropping to a\n    # constant physical extent first keeps the resample ratio comparable.\n    px = None\n    ps = getattr(ds, \"PixelSpacing\", None)\n    if ps is not None and len(ps) > 0:\n        try:\n            px = float(ps[0])\n        except Exception:\n            px = None\n    if px and np.isfinite(px) and px > 0:\n        want = int(round(crop_mm / px))\n        h, w = arr.shape\n        if 16 < want < min(h, w):\n            cy, cx = h // 2, w // 2\n            half = want // 2\n            arr = arr[max(0, cy - half):cy + half, max(0, cx - half):cx + half]\n            valid = valid[max(0, cy - half):cy + half, max(0, cx - half):cx + half]\n\n    finite = arr[valid] if valid.any() else arr\n    finite = finite[np.isfinite(finite)]\n    if finite.size == 0:\n        return np.zeros((out_size, out_size), dtype=np.float32)\n    lo, hi = np.percentile(finite, [1, 99])\n    if not np.isfinite(hi - lo) or hi <= lo:\n        lo, hi = float(arr.min()), (float(arr.max()) if arr.max() > arr.min() else lo + 1.0)\n    arr = np.clip((arr - lo) / max(hi - lo, 1e-6), 0, 1)\n\n    arr = cv2.resize(arr, (out_size, out_size), interpolation=cv2.INTER_AREA)\n    return arr.astype(np.float32)\n\n\ndef normalize_laterality(bag: np.ndarray, planes: list[str], laterality: str | None) -> np.ndarray:\n    \"\"\"Map every knee onto a single (left-knee) convention.\n\n    Medial/lateral labels are defined relative to the body midline, so which side of\n    the image they fall on depends on which knee was scanned. Coronal views mirror\n    under a horizontal flip; sagittal stacks do not mirror slice-by-slice — the\n    medial-lateral axis runs *through* the stack, so the slice order is reversed\n    instead. `Laterality` is absent for some studies; those are left alone, since a\n    wrong flip is worse than none.\n    \"\"\"\n    if laterality != \"R\":\n        return bag\n    for v, plane in enumerate(planes):\n        if plane in (\"Coronal\", \"Axial\"):\n            # Both planes keep the patient's left-right axis in-plane (horizontal\n            # in the image), so the same mirroring convention as Coronal applies.\n            bag[v] = bag[v, :, :, ::-1]\n        else:\n            bag[v] = bag[v, ::-1]\n    return bag\n\n\ndef _bag_cache_path(study_id: str) -> Path:\n    # '/' shows up in some StudyInstanceUID-adjacent identifiers in other RSNA\n    # sets; not expected here, but sanitize defensively before using as a filename.\n    safe_id = study_id.replace('/', '_')\n    # Shape-affecting config baked into the filename: if IMG_SIZE, N_SLICES, or\n    # the number of views changes between runs in the same session, old cache\n    # entries are simply orphaned (never matched) instead of being loaded with\n    # the wrong shape.\n    tag = f\"{IMG_SIZE}x{N_SLICES}x{len(SEQUENCE_PRIORITY)}\"\n    return CACHE_DIR / f\"{safe_id}_{tag}.npy\"\n\n\ndef sample_indices(n_files: int, n_sample: int) -> np.ndarray:\n    if n_files <= 1:\n        return np.zeros(n_sample, dtype=int)\n    lo = int(round(0.15 * (n_files - 1)))\n    hi = int(round(0.85 * (n_files - 1)))\n    if hi <= lo:\n        hi = n_files - 1\n    idx = np.round(np.linspace(lo, hi, n_sample)).astype(int)\n    return np.clip(idx, 0, n_files - 1)\n\n\ndef load_study_bag(series_list: list[dict], n_slices: int = N_SLICES) -> np.ndarray:\n    max_views = len(SEQUENCE_PRIORITY)\n    bag = np.zeros((max_views, n_slices, IMG_SIZE, IMG_SIZE), dtype=np.float32)\n    planes = ([\"Sagittal\", \"Coronal\", \"Axial\"] + [\"Sagittal\"] * max_views)[:max_views]\n    laterality = None\n    for v, rec in enumerate(series_list[:max_views]):\n        files = rec[\"files\"]\n        planes[v] = rec.get(\"plane\", planes[v])\n        laterality = laterality or rec.get(\"laterality\")\n        idxs = sample_indices(len(files), n_slices)\n        for s, idx in enumerate(idxs):\n            try:\n                bag[v, s] = decode_slice(files[int(idx)])\n            except Exception:\n                pass\n    bag = normalize_laterality(bag, planes, laterality)\n    return bag\n\n\n# In-memory cache: if the whole corpus fits in available RAM (checked at\n# runtime by build_ram_cache below), decoded bags live here instead of on\n# disk, and every load_study_bag_cached call becomes a pure array index --\n# no file I/O, no per-item numpy deserialization. A plain numpy array's data\n# buffer is not touched by Python's per-object refcounting the way a large\n# dict of Path objects would be, so when DataLoader workers fork after this\n# is built, the OS keeps it as shared read-only memory rather than\n# duplicating it per worker -- unlike the earlier _cache_bytes_written race,\n# there is no NUM_WORKERS multiplier to worry about here.\n_RAM_CACHE: dict[str, int] | None = None\n_RAM_CACHE_ARRAY: np.ndarray | None = None\nRAM_CACHE_SAFETY_MARGIN_GB = 6.0  # headroom for the model, CUDA pinned buffers,\n                                  # OS, and everything else sharing this\n                                  # session's RAM -- guessing wrong here means\n                                  # an OOM-killed kernel, worse than just\n                                  # staying on the disk cache, so this errs\n                                  # conservative.\nRAM_CACHE_THREADS = 16\n\n\ndef _available_ram_bytes() -> int | None:\n    try:\n        with open(\"/proc/meminfo\") as f:\n            for line in f:\n                if line.startswith(\"MemAvailable:\"):\n                    return int(line.split()[1]) * 1024\n    except Exception:\n        pass\n    return None\n\n\ndef build_ram_cache(all_ids: list[str], combined_index: dict,\n                     n_slices: int = N_SLICES) -> bool:\n    \"\"\"Attempt to decode every study once into one big in-memory uint8 array.\n\n    Returns True if the RAM cache was built and is now in use, False if it was\n    skipped (in which case load_study_bag_cached keeps working exactly as\n    before, via the per-study disk cache) -- this is an optimization attempt,\n    not a requirement.\n    \"\"\"\n    global _RAM_CACHE, _RAM_CACHE_ARRAY\n\n    n_views = len(SEQUENCE_PRIORITY)\n    per_study_bytes = n_views * n_slices * IMG_SIZE * IMG_SIZE  # uint8\n    total_gb = per_study_bytes * len(all_ids) / 1024 ** 3\n\n    available = _available_ram_bytes()\n    if available is None:\n        log(\"RAM cache: could not read /proc/meminfo, skipping (staying on disk cache)\")\n        return False\n\n    available_gb = available / 1024 ** 3\n    log(f\"RAM cache: need ~{total_gb:.1f}GB for {len(all_ids)} studies \"\n        f\"({available_gb:.1f}GB currently available)\")\n    if total_gb + RAM_CACHE_SAFETY_MARGIN_GB > available_gb:\n        log(f\"RAM cache: does not fit within a {RAM_CACHE_SAFETY_MARGIN_GB:.0f}GB \"\n            f\"safety margin -- staying on the per-study disk cache instead\")\n        return False\n\n    cache_array = np.zeros((len(all_ids), n_views, n_slices, IMG_SIZE, IMG_SIZE), dtype=np.uint8)\n    id_to_row = {sid: i for i, sid in enumerate(all_ids)}\n\n    def _job(i):\n        sid = all_ids[i]\n        bag = load_study_bag(combined_index.get(sid, []), n_slices=n_slices)\n        return i, np.clip(bag * 255.0, 0, 255).astype(np.uint8)\n\n    done = 0\n    with ThreadPoolExecutor(max_workers=RAM_CACHE_THREADS) as pool:\n        for i, quantized in pool.map(_job, range(len(all_ids))):\n            cache_array[i] = quantized\n            done += 1\n            if done % 500 == 0:\n                log(f\"  RAM cache: built {done}/{len(all_ids)}\")\n\n    _RAM_CACHE = id_to_row\n    _RAM_CACHE_ARRAY = cache_array\n    log(f\"RAM cache: ready, {cache_array.nbytes / 1024 ** 3:.1f}GB resident, \"\n        f\"disk cache no longer used for these studies\")\n    return True\n\n\ndef load_study_bag_cached(study_id: str, series_list: list[dict],\n                           n_slices: int = N_SLICES) -> np.ndarray:\n    \"\"\"Decode-once, reuse-forever wrapper around `load_study_bag`.\n\n    Checks the in-memory cache first (pure array index, no I/O at all); falls\n    through to the per-study disk cache below when the RAM cache wasn't built\n    or doesn't cover this study.\n\n    Without this, every one of the ~30 (N_FOLDS x EPOCHS) passes over a study\n    re-reads and re-decodes the same DICOM files from scratch — by far the largest\n    cost in the whole pipeline. The physical crop, laterality normalization, and\n    resize are deterministic given the same series selection, so the decoded bag\n    is cached to disk and only random *augmentation* (flip/gamma/noise) is still\n    applied fresh in the Dataset on every access.\n\n    BUGFIX: with num_workers > 0 the actual decode+write happens inside forked\n    DataLoader worker processes, not the main process. `_cache_bytes_written` is\n    a plain Python int, not shared memory (no multiprocessing.Value/Manager), so\n    each worker inherits its own private copy (=0 at fork time) and increments it\n    independently — every worker then writes up to the *full* SLICE_CACHE_BYTE_BUDGET\n    on its own clock, with no visibility into what the other workers wrote. With\n    NUM_WORKERS=4 that means actual bytes written to /kaggle/temp can approach\n    4x SLICE_CACHE_BYTE_BUDGET, defeating the free-space safety margin and risking\n    a full disk mid-run (which then breaks unrelated writes like checkpoints).\n    Fix: split the budget evenly across workers using DataLoader's worker_info,\n    so each worker enforces only its own share and the sum across all workers\n    stays within SLICE_CACHE_BYTE_BUDGET regardless of NUM_WORKERS.\n    \"\"\"\n    if _RAM_CACHE is not None and study_id in _RAM_CACHE:\n        return _RAM_CACHE_ARRAY[_RAM_CACHE[study_id]].astype(np.float32) / 255.0\n\n    global ENABLE_SLICE_CACHE, _cache_write_warned, _cache_bytes_written\n\n    cache_path = _bag_cache_path(study_id)\n    if ENABLE_SLICE_CACHE and cache_path.exists():\n        try:\n            quantized = np.load(cache_path)\n            return quantized.astype(np.float32) / 255.0\n        except Exception:\n            pass  # corrupt cache entry, fall through and rebuild it\n\n    bag = load_study_bag(series_list, n_slices=n_slices)\n\n    # Each DataLoader worker is a separate forked process with its own private\n    # copy of _cache_bytes_written, so the budget below must be this worker's\n    # *share* of SLICE_CACHE_BYTE_BUDGET, not the full budget — otherwise every\n    # worker independently writes up to the full budget and the real total on\n    # disk can reach ~NUM_WORKERS x SLICE_CACHE_BYTE_BUDGET.\n    worker_info = torch.utils.data.get_worker_info()\n    n_workers = worker_info.num_workers if worker_info is not None else 1\n    _worker_cache_budget = SLICE_CACHE_BYTE_BUDGET // max(1, n_workers)\n\n    # Stop growing the cache once the real measured budget is reached, instead\n    # of assuming free space based on which directory this is. Already-cached\n    # entries keep being served; only NEW entries beyond the budget are skipped.\n    if ENABLE_SLICE_CACHE and _cache_bytes_written < _worker_cache_budget:\n        try:\n            quantized = np.clip(bag * 255.0, 0, 255).astype(np.uint8)\n            np.save(cache_path, quantized)\n            _cache_bytes_written += quantized.nbytes\n        except OSError as e:\n            # Still guard against an unexpectedly wrong free-space reading (e.g.\n            # another process on the same shared host writing at the same time).\n            if not _cache_write_warned:\n                log(f\"WARNING: slice cache write failed ({e}); disabling the \"\n                    f\"slice cache for the rest of this run (falling back to \"\n                    f\"always-decode, which is slower but still correct).\")\n                _cache_write_warned = True\n            ENABLE_SLICE_CACHE = False\n    return bag\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T18:12:07.297234Z","iopub.execute_input":"2026-08-07T18:12:07.297576Z","iopub.status.idle":"2026-08-07T18:12:07.557238Z","shell.execute_reply.started":"2026-08-07T18:12:07.297555Z","shell.execute_reply":"2026-08-07T18:12:07.556686Z"}},"outputs":[],"execution_count":null},{"id":"bc0fdf86-eb22-44e7-a1c5-3be6a88ef4bf","cell_type":"markdown","source":"## 5. Dataset and DataLoader\n\nEach item now carries a continuous target vector `y` (weak label, or gold label\nwhere available) **and** a per-target sample weight `w`, instead of the raw `NaN`-\nriddled label columns. `w` is 1.0 for the 58 gold-labeled studies and\n`confidence`-scaled for the rest, so a report that never mentions a finding pulls\nweakly on that output instead of asserting a hard negative.","metadata":{}},{"id":"1b53e628-a0d5-42ad-85c7-da360a33639a","cell_type":"code","source":"class KneeMRIDataset(Dataset):\n    def __init__(self, study_ids, study_index, y=None, w=None, train=False):\n        self.study_ids = list(study_ids)\n        self.study_index = study_index\n        self.y = y   # dict[study_id] -> np.ndarray (N_TARGETS,)\n        self.w = w   # dict[study_id] -> np.ndarray (N_TARGETS,)\n        self.train = train\n\n    def __len__(self):\n        return len(self.study_ids)\n\n    def __getitem__(self, i):\n        study_id = self.study_ids[i]\n        series_list = self.study_index.get(study_id, [])\n        bag = load_study_bag_cached(study_id, series_list)  # (views, slices, H, W)\n        n_views = bag.shape[0]\n\n        # Which bag position holds real pixels vs zero-padding, and which\n        # anatomical plane occupies each position for THIS study specifically —\n        # select_sequences compacts around missing views (e.g. if sagittal has no\n        # matching series, position 0 becomes coronal), so this is read per-study\n        # from series_list rather than assumed fixed.\n        planes = [rec.get(\"plane\", \"Sagittal\") for rec in series_list[:n_views]]\n        planes += [\"Sagittal\"] * (n_views - len(planes))  # unused padding slots;\n                                                            # mask=0 makes the label moot\n        mask = np.zeros(n_views, dtype=np.float32)\n        mask[:len(series_list)] = 1.0\n\n        prior_bias = np.zeros((N_TARGETS, n_views), dtype=np.float32)\n        for v, (plane, present) in enumerate(zip(planes, mask)):\n            if not present:\n                continue\n            for t_idx, t_name in enumerate(TARGETS):\n                prior_bias[t_idx, v] = TARGET_VIEW_PRIOR.get(t_name, {}).get(plane, 0.0)\n\n        if self.train:\n            bag = self._augment(bag)\n\n        bag = torch.from_numpy(bag).unsqueeze(2)  # (views, slices, 1, H, W)\n        bag = bag.repeat(1, 1, 3, 1, 1)            # 3-channel for ImageNet backbone\n        mask_t = torch.from_numpy(mask)\n        prior_t = torch.from_numpy(prior_bias)\n\n        if self.y is not None:\n            y = torch.from_numpy(self.y[study_id].astype(np.float32))\n            w = torch.from_numpy(self.w[study_id].astype(np.float32))\n            return bag, mask_t, prior_t, y, w, study_id\n        return bag, mask_t, prior_t, study_id\n\n    @staticmethod\n    def _augment(bag: np.ndarray) -> np.ndarray:\n        if random.random() < 0.5:\n            # Vertical, not horizontal: load_study_bag_cached already normalizes\n            # every knee onto one left/right convention (normalize_laterality), and\n            # a horizontal flip here would silently undo that on half the batches.\n            bag = bag[:, :, ::-1, :].copy()\n        gamma = random.uniform(0.9, 1.1)\n        bag = np.clip(bag, 0, 1) ** gamma\n        if random.random() < 0.3:\n            bag = np.clip(bag + np.random.normal(0, 0.02, bag.shape), 0, 1).astype(np.float32)\n        return bag\n\n\ndef collate_bags(batch):\n    if len(batch[0]) == 6:\n        bags, masks, priors, ys, ws, ids = zip(*batch)\n        return (torch.stack(bags), torch.stack(masks), torch.stack(priors),\n                torch.stack(ys), torch.stack(ws), list(ids))\n    bags, masks, priors, ids = zip(*batch)\n    return torch.stack(bags), torch.stack(masks), torch.stack(priors), list(ids)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T18:12:07.558154Z","iopub.execute_input":"2026-08-07T18:12:07.558453Z","iopub.status.idle":"2026-08-07T18:12:07.567042Z","shell.execute_reply.started":"2026-08-07T18:12:07.558425Z","shell.execute_reply":"2026-08-07T18:12:07.56614Z"}},"outputs":[],"execution_count":null},{"id":"3db9b326-ace8-4f7c-a303-59f4716a669b","cell_type":"markdown","source":"## 6. Model — EfficientNet-B3 encoder + attention pooling MIL head\n\nUnchanged from v1 except for `set_backbone_trainable`, which now only exposes the\nlast `UNFREEZE_LAST_STAGES` MBConv stages instead of the whole backbone — this is\nthe fix for the OOM. timm's EfficientNet stores its stages as `backbone.blocks`, a\n`Sequential` of stage `Sequential`s; freezing all but the last few keeps most of the\nnetwork out of autograd's graph, so a 96-image mega-batch backward pass stays\naffordable.","metadata":{}},{"id":"92a91ef9-cc1b-43f9-a89a-eeaad20bfce2","cell_type":"code","source":"import timm\n\n\ndef find_timm_checkpoint(keyword: str = \"efficientnet\") -> Path | None:\n    \"\"\"Locate a mounted Kaggle 'Models' checkpoint file whose path contains\n    `keyword` (e.g. the TIMM > EfficientNet > tf-efficientnet-b3 Kaggle Model).\n\n    Kaggle Models mount under a slug-derived path that isn't fixed in advance\n    (varies by framework/variation/version chosen when attaching), so this\n    searches rather than hardcodes a path — same approach as find_dinov2.\n    \"\"\"\n    base = Path(\"/kaggle/input\")\n    if not base.is_dir():\n        return None\n    hits = []\n    for root, dirs, files in os.walk(base):\n        dirs[:] = [d for d in dirs if d not in (\"train_series\", \"test_series\")]\n        if keyword.lower() not in root.lower():\n            continue\n        for f in files:\n            if f.endswith((\".pth\", \".pt\", \".safetensors\", \".bin\")):\n                hits.append(Path(root) / f)\n    return hits[0] if hits else None\n\n\nTIMM_CHECKPOINT_PATH = find_timm_checkpoint(\"efficientnet\")\n\n\nclass AttentionPool(nn.Module):\n    \"\"\"Gated attention pooling (Ilse et al., 2018) over the instance axis.\"\"\"\n\n    def __init__(self, in_dim: int, hidden_dim: int = 128):\n        super().__init__()\n        self.V = nn.Linear(in_dim, hidden_dim)\n        self.U = nn.Linear(in_dim, hidden_dim)\n        self.w = nn.Linear(hidden_dim, 1)\n\n    def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:\n        a = torch.tanh(self.V(x)) * torch.sigmoid(self.U(x))\n        scores = self.w(a).squeeze(-1)\n        weights = torch.softmax(scores, dim=-1)\n        pooled = (weights.unsqueeze(-1) * x).sum(dim=1)\n        return pooled, weights\n\n\nclass TargetSlotHead(nn.Module):\n    \"\"\"Per-target attention over the view axis (Sagittal/Coronal/Axial), masked\n    to views actually present in a study and seeded with an anatomical prior\n    (e.g. ACL read predominantly on the sagittal sequence — see\n    TARGET_VIEW_PRIOR). The prior only shifts where the softmax starts; it never\n    excludes a view outright, and training is free to override it.\n\n    Replaces a single shared view_pool + head: with those, all 12 targets read\n    off the SAME pooled representation, so a finding visible mainly on one plane\n    (PF OA on axial, MCL on coronal, ...) had its evidence diluted by whichever\n    views happened to dominate the shared pooling. Each target now gets its own\n    query and can weight the views independently.\n    \"\"\"\n\n    def __init__(self, dim, n_views, n_targets, hidden=256, p=0.3):\n        super().__init__()\n        self.proj = nn.Sequential(nn.LayerNorm(dim), nn.Linear(dim, hidden), nn.GELU())\n        self.view_emb = nn.Parameter(torch.randn(n_views, hidden) * 0.02)\n        self.query = nn.Parameter(torch.randn(n_targets, hidden) * 0.02)\n        self.drop = nn.Dropout(p)\n        self.out = nn.Linear(hidden, 1)\n\n    def forward(self, x: torch.Tensor, mask: torch.Tensor, prior_bias: torch.Tensor) -> torch.Tensor:\n        # x: (batch, n_views, dim); mask: (batch, n_views); prior_bias: (batch, n_targets, n_views)\n        h = self.proj(x) + self.view_emb.unsqueeze(0)         # (batch, n_views, hidden)\n        att = torch.einsum(\"bvh,th->btv\", h, self.query) + prior_bias  # (batch, n_targets, n_views)\n\n        # A study whose series lookup found literally nothing (should be rare) would\n        # otherwise mask out every view and softmax over an all -inf row, producing\n        # NaN. Fall back to unmasked attention just for those rows so one bad study\n        # can't poison a whole batch's loss.\n        has_any = mask.sum(dim=-1, keepdim=True) > 0\n        safe_mask = torch.where(has_any, mask, torch.ones_like(mask))\n        att = att.masked_fill(safe_mask.unsqueeze(1) < 0.5, float(\"-inf\"))\n\n        weights = torch.softmax(att, dim=-1)                   # (batch, n_targets, n_views)\n        pooled = torch.einsum(\"btv,bvh->bth\", weights, h)      # (batch, n_targets, hidden)\n        pooled = self.drop(pooled)\n        logits = self.out(pooled).squeeze(-1)                  # (batch, n_targets)\n        return logits\n\n\nclass KneeMILModel(nn.Module):\n    def __init__(self, n_targets: int = N_TARGETS, backbone_name: str = \"tf_efficientnet_b3\"):\n        super().__init__()\n        # tf_efficientnet_b3 (not efficientnet_b3): matches the \"tf-efficientnet-b3\"\n        # variation of the Kaggle TIMM > EfficientNet Model — TF-ported weights use\n        # different padding/BN-epsilon conventions than timm's native training\n        # recipe, so the model name has to match the checkpoint's architecture or\n        # loading silently mismatches keys (the old manual-path version reported\n        # \"unexpected=2\" every time, a symptom of exactly this kind of mismatch).\n        if TIMM_CHECKPOINT_PATH is None:\n            raise FileNotFoundError(\n                \"TIMM EfficientNet checkpoint not found under /kaggle/input — attach \"\n                \"the 'TIMM > EfficientNet > tf-efficientnet-b3 > PyTorch' Kaggle Model \"\n                \"before running.\"\n            )\n        # checkpoint_path= is timm's own local-weights loading path. It loads with\n        # strict=True internally and no way to relax that through create_model, so\n        # the model has to be built WITH its original 1000-class ImageNet classifier\n        # first (matching every key the checkpoint actually contains) — building it\n        # with num_classes=0 up front made the classifier.weight/bias keys in the\n        # checkpoint \"unexpected\" and the strict load failed. The classifier is\n        # stripped afterward via reset_classifier(0), which only swaps out the head\n        # module and leaves backbone.num_features (used below) untouched.\n        self.backbone = timm.create_model(\n            backbone_name, pretrained=False,\n            checkpoint_path=str(TIMM_CHECKPOINT_PATH),\n        )\n        self.backbone.reset_classifier(0)\n        log(f\"loaded TIMM {backbone_name} checkpoint from {TIMM_CHECKPOINT_PATH}\")\n\n        feat_dim = self.backbone.num_features\n        self.slice_pool = AttentionPool(feat_dim, hidden_dim=128)\n        self.target_head = TargetSlotHead(\n            feat_dim, n_views=len(SEQUENCE_PRIORITY), n_targets=n_targets,\n            hidden=256, p=0.3,\n        )\n\n    def forward(self, bag: torch.Tensor, mask: torch.Tensor, prior_bias: torch.Tensor) -> torch.Tensor:\n        b, v, s, c, h, w = bag.shape\n        flat = bag.view(b * v * s, c, h, w)\n        # channels_last lets cuDNN use tensor-core-friendly conv kernels on this\n        # NCHW workload — free speed on Ampere+ GPUs, no effect on the result.\n        # .contiguous() is required: .view() above can return a tensor that is\n        # already non-contiguous in a way that silently no-ops the format change.\n        flat = flat.contiguous(memory_format=torch.channels_last)\n        feats = self.backbone(flat)\n        feats = feats.view(b, v, s, -1)\n\n        feats = feats.view(b * v, s, -1)\n        slice_pooled, _ = self.slice_pool(feats)\n        slice_pooled = slice_pooled.view(b, v, -1)\n\n        logits = self.target_head(slice_pooled, mask, prior_bias)\n        return logits\n\n    def set_backbone_trainable(self, last_n_stages: int, grad_checkpoint: bool | None = None) -> None:\n        \"\"\"Freeze everything except the last `last_n_stages` MBConv stages.\n\n        `grad_checkpoint`, when given, is remembered on the model (`run_epoch` calls\n        this every training epoch to reapply BN train/eval state, without knowing\n        about checkpointing) — so passing it once at the epoch where the backbone\n        opens is enough; later calls reuse whatever was last set.\n        \"\"\"\n        if grad_checkpoint is not None:\n            self._grad_checkpoint_enabled = grad_checkpoint\n        grad_checkpoint = getattr(self, \"_grad_checkpoint_enabled\", False)\n\n        self._frozen_stage_count = last_n_stages\n        for p in self.backbone.parameters():\n            p.requires_grad = False\n        if hasattr(self.backbone, \"set_grad_checkpointing\"):\n            self.backbone.set_grad_checkpointing(False)\n        self.backbone.eval()  # Frozen ALL -> Block BN running-stats\n\n        if last_n_stages <= 0:\n            return\n        stages = list(self.backbone.blocks.children())\n        trainable_stages = stages[max(0, len(stages) - last_n_stages):]\n        for stage in trainable_stages:\n            for p in stage.parameters():\n                p.requires_grad = True\n            stage.train()  # only stage really learn need updated BN-stats\n        for attr in (\"conv_head\", \"bn2\"):\n            module = getattr(self.backbone, attr, None)\n            if module is not None:\n                for p in module.parameters():\n                    p.requires_grad = True\n                module.train()\n\n        if grad_checkpoint and hasattr(self.backbone, \"set_grad_checkpointing\"):\n            self.backbone.set_grad_checkpointing(True)\n\n\ndef unwrap_model(model: nn.Module) -> \"KneeMILModel\":\n    \"\"\"Return the underlying KneeMILModel whether or not `model` is wrapped in\n    nn.DataParallel. DataParallel only proxies forward()/train()/eval(); custom\n    methods like set_backbone_trainable and direct submodule access (.backbone,\n    .head, ...) need the unwrapped module.\n    \"\"\"\n    return model.module if isinstance(model, nn.DataParallel) else model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T18:12:07.56795Z","iopub.execute_input":"2026-08-07T18:12:07.568177Z","iopub.status.idle":"2026-08-07T18:12:07.585275Z","shell.execute_reply.started":"2026-08-07T18:12:07.568147Z","shell.execute_reply":"2026-08-07T18:12:07.584629Z"}},"outputs":[],"execution_count":null},{"id":"6afb4d67-983e-4a14-bf33-5731d5edb82f","cell_type":"markdown","source":"## 7. Loss, metrics, and the per-fold training loop\n\n`run_epoch` now takes a per-sample, per-target weight tensor `w` and uses mixed\nprecision (`torch.autocast` + `GradScaler`), which roughly halves activation memory\non top of the partial-unfreeze fix above.","metadata":{}},{"id":"bc544449-6781-4504-b30b-637151338aa7","cell_type":"code","source":"def compute_pos_weight(y: np.ndarray, w: np.ndarray) -> torch.Tensor:\n    pos = (w * y).sum(axis=0)\n    neg = (w * (1.0 - y)).sum(axis=0)\n    weight = np.clip(neg / np.maximum(pos, 1e-5), 0.5, 8.0)\n    return torch.tensor(weight, dtype=torch.float32)\n\n\ndef weighted_bce_loss(logits, y, w, pos_weight, focal_gamma: float = 0.0):\n    pos_term = F.softplus(-logits)   # -log(p),   the y==1 cross-entropy term\n    neg_term = F.softplus(logits)    # -log(1-p), the y==0 cross-entropy term\n\n    if focal_gamma > 0:\n        # p_t under the (soft, possibly non-0/1) target y: recovers the plain\n        # BCE p_t exactly when y in {0, 1}, and interpolates smoothly for the\n        # continuous weak-label scores in between. Detached: focal modulation\n        # should reweight the loss magnitude, not add its own gradient path.\n        raw_bce = (y * pos_term + (1.0 - y) * neg_term).detach()\n        pt = torch.exp(-raw_bce)\n        modulation = (1.0 - pt).clamp(min=0.0) ** focal_gamma\n    else:\n        modulation = 1.0\n\n    positive = pos_term * y * pos_weight.unsqueeze(0)\n    negative = neg_term * (1.0 - y)\n    element = (positive + negative) * modulation\n    return (element * w).sum() / w.sum().clamp_min(1.0)\n\n\ndef mean_auc(y_true: np.ndarray, y_pred: np.ndarray) -> tuple[float, dict]:\n    \"\"\"Macro-average AUC across the 12 targets; skips degenerate columns.\"\"\"\n    scores = {}\n    for i, name in enumerate(TARGETS):\n        col_true = (y_true[:, i] > 0.5).astype(int)\n        if len(np.unique(col_true)) < 2:\n            continue\n        try:\n            scores[name] = roc_auc_score(col_true, y_pred[:, i])\n        except ValueError:\n            continue\n    if not scores:\n        return float(\"nan\"), scores\n    return float(np.mean(list(scores.values()))), scores\n\n\ndef run_epoch(model, base_model, loader, optimizer, pos_weight, scaler, train: bool, ema=None):\n    model.train(mode=train)\n    if train:\n        base_model.set_backbone_trainable(base_model._frozen_stage_count)\n    total_loss, n_batches = 0.0, 0\n    all_true, all_pred, all_ids = [], [], []\n\n    for bag, mask, prior_bias, y, w, ids in loader:\n        bag = bag.to(DEVICE, non_blocking=True)\n        mask = mask.to(DEVICE, non_blocking=True)\n        prior_bias = prior_bias.to(DEVICE, non_blocking=True)\n        y = y.to(DEVICE, non_blocking=True)\n        w = w.to(DEVICE, non_blocking=True)\n        with torch.set_grad_enabled(train):\n            with torch.autocast(device_type=\"cuda\", enabled=DEVICE.type == \"cuda\"):\n                logits = model(bag, mask, prior_bias)\n                loss = weighted_bce_loss(logits, y, w, pos_weight.to(DEVICE), focal_gamma=FOCAL_GAMMA)\n            if train:\n                optimizer.zero_grad(set_to_none=True)\n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n                scaler.step(optimizer)\n                scaler.update()\n                if ema is not None:\n                    ema.update(model)\n\n        if not torch.isfinite(loss):\n            log(f\"  WARNING: non-finite loss for batch ids={ids}, skipping\")\n            continue\n        total_loss += loss.item()\n        n_batches += 1\n        all_true.append(y.detach().cpu().numpy())\n        all_pred.append(torch.sigmoid(logits.detach().float()).cpu().numpy())\n        all_ids.extend(ids)\n\n    y_true = np.concatenate(all_true)\n    y_pred = np.concatenate(all_pred)\n    auc, per_label = mean_auc(y_true, y_pred)\n    return total_loss / max(n_batches, 1), auc, per_label, y_true, y_pred, all_ids\n\ndef train_fold(fold: int, train_ids, val_ids, study_index, y_map, w_map, gold_ids):\n    log(f\"fold {fold}: train={len(train_ids)} val={len(val_ids)}\")\n    gold_ids_set = set(gold_ids)\n\n    train_ds = KneeMRIDataset(train_ids, study_index, y_map, w_map, train=True)\n    val_ds = KneeMRIDataset(val_ids, study_index, y_map, w_map, train=False)\n    train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True,\n                               num_workers=NUM_WORKERS, collate_fn=collate_bags, drop_last=True,\n                               persistent_workers=PERSISTENT_WORKERS, pin_memory=True)\n    val_loader = DataLoader(val_ds, batch_size=EVAL_BATCH_SIZE, shuffle=False,\n                             num_workers=NUM_WORKERS, collate_fn=collate_bags,\n                             persistent_workers=PERSISTENT_WORKERS, pin_memory=True)\n\n    y_train = np.stack([y_map[s] for s in train_ids])\n    w_train = np.stack([w_map[s] for s in train_ids])\n    pos_weight = compute_pos_weight(y_train, w_train)\n\n    base_model = KneeMILModel().to(DEVICE)\n    base_model.set_backbone_trainable(0)  # warm up the head first, backbone fully frozen\n\n    model = base_model\n\n    head_params = list(base_model.slice_pool.parameters()) + \\\n                  list(base_model.target_head.parameters())\n    optimizer = torch.optim.AdamW(\n        [{\"params\": head_params, \"lr\": LR_HEAD}], weight_decay=WEIGHT_DECAY\n    )\n    try:\n        scaler = torch.amp.GradScaler(\"cuda\", enabled=DEVICE.type == \"cuda\")\n    except Exception:\n        scaler = torch.cuda.amp.GradScaler(enabled=DEVICE.type == \"cuda\")\n\n    ema = ModelEmaV2(base_model, decay=EMA_DECAY) if USE_EMA else None\n\n    best_auc, best_state = -1.0, None\n    best_gold_true, best_gold_pred, best_gold_ids = None, None, None\n    best_val_true, best_val_pred, best_val_ids = None, None, None\n    scheduler = None\n\n    for epoch in range(1, EPOCHS + 1):\n        if epoch == 2:\n            base_model.set_backbone_trainable(UNFREEZE_LAST_STAGES, grad_checkpoint=USE_GRAD_CHECKPOINTING)\n            backbone_params = [p for p in base_model.backbone.parameters() if p.requires_grad]\n            optimizer.add_param_group({\"params\": backbone_params, \"lr\": LR_BACKBONE})\n            n_trainable = sum(p.numel() for p in backbone_params)\n            log(f\"  fold {fold}: unfroze last {UNFREEZE_LAST_STAGES} backbone stages \"\n                f\"({n_trainable / 1e6:.1f}M params, checkpointing={USE_GRAD_CHECKPOINTING})\")\n            \n            scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n                optimizer, T_max=max(1, EPOCHS - 1)\n            )\n\n        train_loss, train_auc, _, _, _, _ = run_epoch(\n            model, base_model, train_loader, optimizer, pos_weight, scaler, train=True, ema=ema\n        )\n        val_loss, val_auc, per_label, val_y_true, val_y_pred, val_ids = run_epoch(\n            model, base_model, val_loader, optimizer, pos_weight, scaler, train=False\n        )\n\n        gold_mask = np.array([sid in gold_ids_set for sid in val_ids])\n        if gold_mask.sum() >= 2:\n            gold_auc, _ = mean_auc(val_y_true[gold_mask], val_y_pred[gold_mask])\n        else:\n            gold_auc = float(\"nan\")\n\n        # EMA is evaluated on the same val split — this is an extra (cheap, since\n        # val is much smaller than train) inference-only pass, not extra training\n        # compute. Whichever of (raw, EMA) scores higher THIS epoch becomes the\n        # new best checkpoint; EMA usually starts winning once the backbone opens\n        # and per-step weight noise increases.\n        ema_val_auc = float(\"nan\")\n        ema_gold_auc = float(\"nan\")\n        ema_y_true = ema_y_pred = ema_ids = ema_gold_mask = None\n        # Skipped at epoch 1: the backbone is still fully frozen and EMA has barely\n        # diverged from the raw weights yet, so this extra val pass (same cost as\n        # the raw-model one above) buys almost no information there. Still updated\n        # every training step throughout (see run_epoch) — only the extra eval\n        # pass is skipped, not the EMA itself.\n        if ema is not None and epoch > 1:\n            _, ema_val_auc, _, ema_y_true, ema_y_pred, ema_ids = run_epoch(\n                ema.module, ema.module, val_loader, optimizer, pos_weight, scaler, train=False\n            )\n            ema_gold_mask = np.array([sid in gold_ids_set for sid in ema_ids])\n            if ema_gold_mask.sum() >= 2:\n                ema_gold_auc, _ = mean_auc(ema_y_true[ema_gold_mask], ema_y_pred[ema_gold_mask])\n\n        log(f\"  fold {fold} epoch {epoch}/{EPOCHS}: \"\n            f\"train_loss={train_loss:.4f} train_auc={train_auc:.4f} \"\n            f\"val_loss={val_loss:.4f} val_auc={val_auc:.4f} \"\n            f\"gold_auc={gold_auc:.4f} (n_gold_in_val={int(gold_mask.sum())}) \"\n            f\"ema_val_auc={ema_val_auc:.4f} ema_gold_auc={ema_gold_auc:.4f}\")\n\n        for auc_value, state_source, cand_true, cand_pred, cand_ids, cand_mask in (\n            (val_auc, base_model, val_y_true, val_y_pred, val_ids, gold_mask),\n            (ema_val_auc, ema.module if (ema is not None and epoch > 1) else None,\n             ema_y_true, ema_y_pred, ema_ids, ema_gold_mask),\n        ):\n            if state_source is not None and auc_value > best_auc:\n                best_auc = auc_value\n                best_state = {k: v.detach().cpu().clone() for k, v in state_source.state_dict().items()}\n                # Full val set (mostly weak-label targets, ~11-15/fold are gold) —\n                # this is what gets pooled into the OOF weak-label AUC below, the\n                # far-larger-n metric that should arbitrate real changes; gold_true\n                # stays a small, noisy, enrichment-biased secondary check only.\n                best_val_true, best_val_pred, best_val_ids = cand_true, cand_pred, cand_ids\n                if cand_mask is not None and cand_mask.sum() > 0:\n                    best_gold_true = cand_true[cand_mask]\n                    best_gold_pred = cand_pred[cand_mask]\n                    best_gold_ids = [cand_ids[i] for i in np.where(cand_mask)[0]]\n                else:\n                    best_gold_true, best_gold_pred, best_gold_ids = None, None, None\n\n        if scheduler is not None:\n            scheduler.step()\n\n        if time.time() - T0 > TIME_BUDGET_SECONDS:\n            log(\"  time budget reached, stopping fold early\")\n            break\n\n    # Exact-refine: a few more epochs on just the gold studies that landed in this\n    # fold's train split, at a much lower LR than main training used. Held-out val\n    # (including this fold's own gold-in-val subset) still arbitrates whether each\n    # refine epoch actually improves on the best main-training checkpoint.\n    gold_train_ids = [s for s in train_ids if s in gold_ids_set]\n    if EXACT_REFINE_EPOCHS > 0 and len(gold_train_ids) >= 2:\n        log(f\"  fold {fold}: exact-refine on {len(gold_train_ids)} gold studies \"\n            f\"for {EXACT_REFINE_EPOCHS} epoch(s)\")\n        refine_bs = min(BATCH_SIZE, len(gold_train_ids))\n        refine_ds = KneeMRIDataset(gold_train_ids, study_index, y_map, w_map, train=True)\n        refine_loader = DataLoader(refine_ds, batch_size=refine_bs, shuffle=True,\n                                    num_workers=NUM_WORKERS, collate_fn=collate_bags,\n                                    drop_last=False, persistent_workers=False, pin_memory=True)\n\n        refine_optimizer = torch.optim.AdamW(\n            [{\"params\": list(base_model.slice_pool.parameters()) +\n                        list(base_model.target_head.parameters()), \"lr\": LR_REFINE_HEAD},\n             {\"params\": [p for p in base_model.backbone.parameters() if p.requires_grad],\n              \"lr\": LR_REFINE_BACKBONE}],\n            weight_decay=WEIGHT_DECAY,\n        )\n        y_refine = np.stack([y_map[s] for s in gold_train_ids])\n        w_refine = np.stack([w_map[s] for s in gold_train_ids])\n        refine_pos_weight = compute_pos_weight(y_refine, w_refine)\n\n        for r_epoch in range(1, EXACT_REFINE_EPOCHS + 1):\n            r_loss, _, _, _, _, _ = run_epoch(\n                model, base_model, refine_loader, refine_optimizer, refine_pos_weight, scaler,\n                train=True, ema=ema,\n            )\n            r_val_loss, r_val_auc, _, r_val_true, r_val_pred, r_val_ids = run_epoch(\n                model, base_model, val_loader, refine_optimizer, refine_pos_weight, scaler, train=False\n            )\n            r_gold_mask = np.array([sid in gold_ids_set for sid in r_val_ids])\n            r_gold_auc = float(\"nan\")\n            if r_gold_mask.sum() >= 2:\n                r_gold_auc, _ = mean_auc(r_val_true[r_gold_mask], r_val_pred[r_gold_mask])\n            log(f\"  fold {fold} exact-refine {r_epoch}/{EXACT_REFINE_EPOCHS}: \"\n                f\"refine_loss={r_loss:.4f} val_auc={r_val_auc:.4f} \"\n                f\"gold_auc={r_gold_auc:.4f} (n_gold_in_val={int(r_gold_mask.sum())})\")\n\n            if r_val_auc > best_auc:\n                best_auc = r_val_auc\n                best_state = {k: v.detach().cpu().clone() for k, v in base_model.state_dict().items()}\n                best_val_true, best_val_pred, best_val_ids = r_val_true, r_val_pred, r_val_ids\n                if r_gold_mask.sum() > 0:\n                    best_gold_true = r_val_true[r_gold_mask]\n                    best_gold_pred = r_val_pred[r_gold_mask]\n                    best_gold_ids = [r_val_ids[i] for i in np.where(r_gold_mask)[0]]\n\n            if time.time() - T0 > TIME_BUDGET_SECONDS:\n                log(\"  time budget reached, stopping exact-refine early\")\n                break\n\n    if best_state is not None:\n        base_model.load_state_dict(best_state)\n    return (base_model, best_auc, best_gold_true, best_gold_pred, best_gold_ids,\n            best_val_true, best_val_pred, best_val_ids)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T18:12:07.586194Z","iopub.execute_input":"2026-08-07T18:12:07.586558Z","iopub.status.idle":"2026-08-07T18:12:07.608174Z","shell.execute_reply.started":"2026-08-07T18:12:07.586515Z","shell.execute_reply":"2026-08-07T18:12:07.607316Z"}},"outputs":[],"execution_count":null},{"id":"00bcd432-cdd3-4f0b-ad63-5c8b8c82e006","cell_type":"markdown","source":"## 8. Building weak + gold labels, and report-hash grouped cross-validation\n\nSome reports are duplicated verbatim across studies (a template read for an\nunremarkable knee); every study sharing that report gets the same extracted target\nvector. Splitting such a group across train/val would score the model on a target\nwhose source it already trained on, so folds are grouped by a hash of the\nnormalized report text rather than by `StudyInstanceUID` alone.","metadata":{}},{"id":"31163b3e-cf35-4038-8816-6370b7529c4c","cell_type":"code","source":"def compute_target_reliability(train_df: pd.DataFrame, gold_index: np.ndarray) -> np.ndarray:\n    \"\"\"How much to trust the weak-label extractor for each target individually,\n    measured as AUC of its raw score against the 58 gold labels for that target.\n\n    A single confidence formula applied uniformly to all 12 targets treats them as\n    equally trustworthy, which they are not: something like Fracture is usually\n    described unambiguously in a report, so the extractor tracks it well; something\n    like Synovitis or a graded OA severity is inherently fuzzier in the text itself,\n    so the extractor's score for it is closer to noise no matter how it's tuned.\n\n    Caveat: this uses the same 58 gold rows that also validate the model later\n    (scattered across folds by report_hash_folds). That is a mild, hard-to-avoid\n    leakage at this sample size — 58 rows split per-fold would be too few to\n    estimate 12 separate reliabilities at all — so treat `gold_auc` in the logs as\n    slightly optimistic, not as a completely clean holdout number.\n    \"\"\"\n    indexed = train_df.set_index(\"StudyInstanceUID\")\n    gold_rows = indexed.loc[gold_index]\n    weak_scores = np.zeros((len(gold_index), N_TARGETS), dtype=np.float32)\n    for i, (_study_id, row) in enumerate(gold_rows.iterrows()):\n        extracted = extract_report_labels(row.get(\"Report\", \"\") or \"\")\n        weak_scores[i] = [extracted[t] for t in TARGETS]\n    gold_y = gold_rows[TARGETS].to_numpy(dtype=np.float32)\n\n    reliability = np.full(N_TARGETS, 0.5, dtype=np.float32)  # default: no signal\n    for t in range(N_TARGETS):\n        col_true = (gold_y[:, t] > 0.5).astype(int)\n        if len(np.unique(col_true)) < 2:\n            continue  # can't measure AUC without both classes present in the gold set\n        try:\n            reliability[t] = roc_auc_score(col_true, weak_scores[:, t])\n        except ValueError:\n            continue\n    return reliability\n\n\ndef find_label_table() -> Path | None:\n    \"\"\"Look for an external, pre-computed label table (e.g. LLM-derived, from\n    reading each report with a language model instead of a regex lexicon)\n    mounted as a Kaggle Dataset anywhere under /kaggle/input.\n\n    Same column contract as extract_report_labels' output (each target plus a\n    f\"{{target}}__conf\" column), so it drops straight into the weighting formula\n    below in build_labels -- used for whichever studies it covers, the regex\n    lexicon fills in the rest. Entirely optional: if nothing matching is\n    mounted, this returns None and build_labels behaves exactly as before.\n    \"\"\"\n    base = Path(\"/kaggle/input\")\n    if not base.is_dir():\n        return None\n    needed_cols = set(TARGETS) | {t + \"__conf\" for t in TARGETS}\n    for csv_path in base.rglob(\"*.csv\"):\n        try:\n            cols = set(pd.read_csv(csv_path, nrows=0).columns)\n        except Exception:\n            continue\n        if \"StudyInstanceUID\" in cols and needed_cols.issubset(cols):\n            return csv_path\n    return None\n\n\ndef build_labels(train_df: pd.DataFrame) -> tuple[dict, dict, np.ndarray]:\n    \"\"\"Weak label (report-derived) + confidence, overridden by the 58 gold rows.\n\n    The per-sample confidence (from extract_report_labels, or from the external\n    table below where mounted) says how clearly THIS report stated a finding;\n    `trust` below says how much the extractor can be believed for THIS target\n    overall. The two are multiplied — a low-trust target stays down-weighted\n    even on its most confident-looking rows.\n    \"\"\"\n    indexed = train_df.set_index(\"StudyInstanceUID\")\n    gold = indexed[TARGETS]\n    gold = gold[gold.notna().all(axis=1)]\n\n    ext_table = None\n    ext_path = find_label_table()\n    if ext_path is not None:\n        ext_table = pd.read_csv(ext_path).set_index(\"StudyInstanceUID\")\n        n_covered = len(ext_table.index.intersection(indexed.index))\n        log(f\"external label table: {ext_path.name}, covers {n_covered}/{len(indexed)} \"\n            f\"studies (lexicon fills in the rest)\")\n\n    reliability = compute_target_reliability(train_df, gold.index)\n    # AUC 0.5 (no better than random) -> trust TARGET_TRUST_FLOOR; AUC 1.0\n    # (perfect on the gold rows) -> trust 1.\n    trust = np.clip((reliability - 0.5) * 2.0, 0.0, 1.0)\n    trust = TARGET_TRUST_FLOOR + (1.0 - TARGET_TRUST_FLOOR) * trust\n    log(\"per-target weak-label reliability (AUC vs gold): \" +\n        \", \".join(f\"{t}={r:.2f}\" for t, r in zip(TARGETS, reliability)))\n\n    y_map, w_map = {}, {}\n    for study_id, row in indexed.iterrows():\n        if ext_table is not None and study_id in ext_table.index:\n            ext_row = ext_table.loc[study_id]\n            y = ext_row[TARGETS].to_numpy(dtype=np.float32)\n            conf = ext_row[[t + \"__conf\" for t in TARGETS]].to_numpy(dtype=np.float32)\n        else:\n            extracted = extract_report_labels(row.get(\"Report\", \"\") or \"\")\n            y = np.array([extracted[t] for t in TARGETS], dtype=np.float32)\n            conf = np.array([extracted[t + \"__conf\"] for t in TARGETS], dtype=np.float32)\n        w = (0.15 + 0.85 * conf) * trust\n        if study_id in gold.index:\n            y = gold.loc[study_id].to_numpy(dtype=np.float32)\n            w = np.full(N_TARGETS, GOLD_WEIGHT_BOOST, dtype=np.float32)\n        y_map[study_id] = y\n        w_map[study_id] = w\n\n    log(f\"labels built for {len(y_map)} studies (exact={len(gold)}, \"\n        f\"report-derived={len(y_map) - len(gold)})\")\n    return y_map, w_map, gold.index.to_numpy()\n\n\ndef report_hash_folds(train_df: pd.DataFrame, n_folds: int = N_FOLDS) -> dict:\n    \"\"\"StudyInstanceUID -> fold index, grouped by a hash of the normalized report,\n    so studies sharing a duplicated report always land on the same side of a split.\"\"\"\n    indexed = train_df.set_index(\"StudyInstanceUID\")\n    fold_of = {}\n    for study_id, row in indexed.iterrows():\n        key = normalize(str(row.get(\"Report\", \"\") or study_id)).strip()\n        digest = hashlib.md5(key.encode(\"utf-8\")).hexdigest()\n        fold_of[study_id] = int(digest[:8], 16) % n_folds\n    return fold_of\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T18:12:07.60912Z","iopub.execute_input":"2026-08-07T18:12:07.609608Z","iopub.status.idle":"2026-08-07T18:12:07.624742Z","shell.execute_reply.started":"2026-08-07T18:12:07.609588Z","shell.execute_reply":"2026-08-07T18:12:07.624066Z"}},"outputs":[],"execution_count":null},{"id":"e7c75b4c-19ac-4a42-b1b9-992295044c1b","cell_type":"markdown","source":"## 9. Cross-validation loop","metadata":{}},{"id":"78f6d889-686b-4882-992c-bbcd77b29d0d","cell_type":"code","source":"def run_cross_validation(train_df: pd.DataFrame, study_index: dict) -> list[nn.Module]:\n    y_map, w_map, gold_ids = build_labels(train_df)\n    fold_of = report_hash_folds(train_df, N_FOLDS)\n\n    study_ids = train_df[\"StudyInstanceUID\"].to_numpy()\n    fold_id = np.array([fold_of[s] for s in study_ids])\n\n    fold_models, fold_aucs = [], []\n    oof_gold_true, oof_gold_pred, oof_gold_ids = [], [], []\n    oof_val_true, oof_val_pred, oof_val_ids = [], [], []\n    for fold in range(N_FOLDS):\n        val_mask = fold_id == fold\n        train_ids = study_ids[~val_mask]\n        val_ids = study_ids[val_mask]\n        if len(val_ids) == 0 or len(train_ids) < BATCH_SIZE:\n            log(f\"fold {fold + 1}: invalid split, skipped\")\n            continue\n\n        (model, auc, gold_true, gold_pred, gold_ids_here,\n         val_true_here, val_pred_here, val_ids_here) = train_fold(\n            fold + 1, train_ids, val_ids, study_index, y_map, w_map, gold_ids\n        )\n        fold_models.append(model)\n        fold_aucs.append(auc)\n        if gold_true is not None:\n            oof_gold_true.append(gold_true)\n            oof_gold_pred.append(gold_pred)\n            oof_gold_ids.extend(gold_ids_here)\n        if val_true_here is not None:\n            oof_val_true.append(val_true_here)\n            oof_val_pred.append(val_pred_here)\n            oof_val_ids.extend(val_ids_here)\n\n        try:\n            torch.save(model.state_dict(), OUTPUT_DIR / f\"model_fold{fold + 1}.pt\")\n        except (OSError, RuntimeError) as e:\n            # The trained model is already retained in `fold_models` for this\n            # session's own inference step, so a failed checkpoint write is not\n            # fatal to THIS run — only to resuming a future session from disk.\n            log(f\"WARNING: could not save checkpoint for fold {fold + 1} \"\n                f\"({e}); continuing without it (model stays in memory for \"\n                f\"this run's inference).\")\n\n        del model\n        gc_collect_and_empty_cache()\n\n        if time.time() - T0 > TIME_BUDGET_SECONDS:\n            log(\"time budget reached, stopping cross-validation early\")\n            break\n\n    log(f\"CV mean AUC = {np.nanmean(fold_aucs):.4f} across {len(fold_aucs)} folds\")\n\n    # Each of the 58 gold studies sits in exactly one fold's val split (folds are a\n    # non-overlapping partition), so pooling every fold's gold-in-val predictions\n    # recovers one prediction per gold study — up to n=58 instead of n~11-15 per\n    # fold. Standard error on an AUC estimate shrinks roughly with sqrt(n), so this\n    # is meaningfully more trustworthy than any single fold's gold_auc, though still\n    # not a large-sample number — treat it as the best available read on real\n    # (not report-derived) performance, not as a precise score.\n    # OOF weak-label AUC: every one of the ~4407 training studies sits in exactly\n    # one fold's val split, so pooling every fold's val predictions covers the\n    # whole corpus once each — n in the thousands rather than n<=58. Per\n    # https://www.kaggle.com/code/.../58-studies-cannot-see-a-0-01-gain (the\n    # simulation this project's gold_auc caveats have been citing all along):\n    # paired standard error on the 58-study gold AUC is ~0.0125, so a true 0.01\n    # gain only wins a head-to-head 79% of the time and 0.005 barely beats a coin\n    # flip. This metric is mostly weak (report-derived) labels rather than expert\n    # ground truth, so its absolute level is attenuated/noisy — but with ~76x the\n    # sample size, its STANDARD ERROR is roughly sqrt(76) =~ 8-9x smaller, which\n    # is what actually matters for deciding whether a change between two runs is\n    # real. Read it for direction/ranking between runs, and read OOF gold AUC\n    # for calibration against real labels — not the other way around.\n    if oof_val_true:\n        pooled_val_true = np.concatenate(oof_val_true)\n        pooled_val_pred = np.concatenate(oof_val_pred)\n        oof_weak_auc, oof_weak_per_label = mean_auc(pooled_val_true, pooled_val_pred)\n        log(f\"OOF weak-label AUC (pooled across folds, n={len(oof_val_ids)}) = {oof_weak_auc:.4f}\")\n        log(\"  per-target: \" + \", \".join(f\"{t}={v:.3f}\" for t, v in oof_weak_per_label.items()))\n    else:\n        log(\"OOF weak-label AUC: no val predictions were captured across folds\")\n\n    if oof_gold_true:\n        pooled_true = np.concatenate(oof_gold_true)\n        pooled_pred = np.concatenate(oof_gold_pred)\n        oof_auc, oof_per_label = mean_auc(pooled_true, pooled_pred)\n        log(f\"OOF gold AUC (pooled across folds, n={len(oof_gold_ids)}/58) = {oof_auc:.4f} \"\n            f\"(secondary check only — paired sigma on n<=58 is ~0.0125; do not\\n\"\n            f\"           read anything under ~0.02 vs a previous run as a real change)\")\n        log(\"  per-target: \" + \", \".join(f\"{t}={v:.3f}\" for t, v in oof_per_label.items()))\n    else:\n        log(\"OOF gold AUC: no gold studies were captured across folds\")\n\n    return fold_models\n\n\ndef gc_collect_and_empty_cache():\n    import gc\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T18:12:07.625757Z","iopub.execute_input":"2026-08-07T18:12:07.626045Z","iopub.status.idle":"2026-08-07T18:12:07.641724Z","shell.execute_reply.started":"2026-08-07T18:12:07.626024Z","shell.execute_reply":"2026-08-07T18:12:07.641095Z"}},"outputs":[],"execution_count":null},{"id":"2e8a260b-d4f3-4cc9-95f5-889368e743c1","cell_type":"markdown","source":"## 10. Inference and submission\n\nPredictions are rank-transformed before being written out: the competition score is\nthe unweighted mean of twelve per-label ROC AUCs, which only reads the *order* of\nscores within a label, not their scale. Averaging ranks across fold models is\ntherefore the correct way to combine them — averaging raw probabilities would let\nwhichever fold happens to be most confident dominate.","metadata":{}},{"id":"ad250cfd-ec88-4017-8fee-1f6ab43b764c","cell_type":"code","source":"def write_fallback_submission(test_ids: list[str]) -> None:\n    df = pd.DataFrame({\"StudyInstanceUID\": test_ids})\n    for t in TARGETS:\n        df[t] = 0.5\n    df.to_csv(OUTPUT_DIR / \"submission.csv\", index=False)\n    log(\"wrote fallback (benchmark) submission.csv\")\n\n\ndef _scale_bag(bag: torch.Tensor, scale: float) -> torch.Tensor:\n    \"\"\"Zoom the bag in/out by `scale` (center-crop for >1.0, pad for <1.0), then\n    resize back to the original H/W — a cheap tensor-level multi-scale TTA that\n    doesn't require re-decoding DICOM at a different crop.\n    \"\"\"\n    if scale == 1.0:\n        return bag\n    b, v, s, c, h, w = bag.shape\n    flat = bag.reshape(b * v * s, c, h, w)\n    new_h, new_w = max(1, int(round(h * scale))), max(1, int(round(w * scale)))\n    resized = F.interpolate(flat, size=(new_h, new_w), mode=\"bilinear\", align_corners=False)\n    if scale >= 1.0:\n        top, left = (new_h - h) // 2, (new_w - w) // 2\n        resized = resized[:, :, top:top + h, left:left + w]\n    else:\n        pad_h, pad_w = h - new_h, w - new_w\n        resized = F.pad(resized, (pad_w // 2, pad_w - pad_w // 2, pad_h // 2, pad_h - pad_h // 2))\n    return resized.reshape(b, v, s, c, h, w)\n\n\n@torch.no_grad()\ndef predict_test(models: list[nn.Module], test_ids: list[str], study_index: dict) -> pd.DataFrame:\n    ds = KneeMRIDataset(test_ids, study_index, y=None, w=None, train=False)\n    loader = DataLoader(ds, batch_size=EVAL_BATCH_SIZE, shuffle=False,\n                         num_workers=NUM_WORKERS, collate_fn=collate_bags,\n                         persistent_workers=PERSISTENT_WORKERS, pin_memory=True)\n\n    all_preds = np.zeros((len(test_ids), N_TARGETS), dtype=np.float32)\n    id_to_row = {sid: i for i, sid in enumerate(test_ids)}\n\n    for model in models:\n        model.eval()\n\n    for bag, mask, prior_bias, ids in loader:\n        bag = bag.to(DEVICE)\n        mask = mask.to(DEVICE)\n        prior_bias = prior_bias.to(DEVICE)\n        batch_pred = torch.zeros(bag.shape[0], N_TARGETS, device=DEVICE)\n        max_pool_idx = [TARGETS.index(t) for t in TTA_TARGET_POOL]\n\n        for model in models:\n            variant_probs = []\n            for scale in TTA_SCALES:\n                scaled = _scale_bag(bag, scale)\n                for variant in (scaled, torch.flip(scaled, dims=[-2])):  # vertical, matches augmentation\n                    with torch.autocast(device_type=\"cuda\", enabled=DEVICE.type == \"cuda\"):\n                        # NOT variant.contiguous(memory_format=torch.channels_last) here: variant\n                        # is (batch, view, slice, C, H, W) -- rank 6, not the rank-4 NCHW shape\n                        # channels_last requires. model.forward() already applies channels_last\n                        # correctly, internally, on the reshaped 4D tensor -- this call doesn't\n                        # need to do it too (that mismatch is exactly what crashed this run).\n                        logits = model(variant, mask, prior_bias)\n                    variant_probs.append(torch.sigmoid(logits.float()))\n\n            # Mean-pool across scale/flip variants for most targets; max-pool for\n            # TTA_TARGET_POOL targets (see cell 2). Combined per model first, then\n            # averaged across the fold ensemble below -- unchanged from before.\n            stacked = torch.stack(variant_probs, dim=0)  # (n_variants, batch, N_TARGETS)\n            combined = stacked.mean(dim=0)\n            combined[:, max_pool_idx] = stacked.max(dim=0).values[:, max_pool_idx]\n            batch_pred += combined\n\n        batch_pred /= len(models)\n        batch_pred = batch_pred.cpu().numpy()\n\n        for sid, row in zip(ids, batch_pred):\n            all_preds[id_to_row[sid]] = row\n\n    # Rank-transform per column: the metric reads only order, so this makes the\n    # submission robust to any one fold model being systematically over/under-confident.\n    ranked = pd.DataFrame(all_preds).rank(pct=True).to_numpy(dtype=np.float32)\n    sub = pd.DataFrame(ranked, columns=TARGETS)\n    sub.insert(0, \"StudyInstanceUID\", test_ids)\n    return sub\n\n\ndef main():\n    test_df = pd.read_csv(ROOT / \"test.csv\")\n    train_df = pd.read_csv(ROOT / \"train.csv\")\n    test_ids = test_df[\"StudyInstanceUID\"].tolist()\n\n    write_fallback_submission(test_ids)\n\n    train_series_lookup = load_series_lookup(ROOT / \"train_series.csv\")\n    test_series_lookup = load_series_lookup(ROOT / \"test_series.csv\")\n    log(f\"series lookup: train covers {len(train_series_lookup)} studies, \"\n        f\"test covers {len(test_series_lookup)} studies\")\n\n    train_ids_all = train_df[\"StudyInstanceUID\"].tolist()\n    log(\"indexing train series...\")\n    # cache_name changed (train -> train_gtplane): a pickle built before this\n    # patch was regex-only and would otherwise be loaded as-is, silently\n    # skipping the ground-truth lookup entirely on a resumed session.\n    train_index = build_study_index(TRAIN_DICOM_DIR, train_ids_all, cache_name=\"train_gtplane\",\n                                     series_lookup=train_series_lookup)\n    log(\"indexing test series...\")\n    test_index = build_study_index(TEST_DICOM_DIR, test_ids, cache_name=\"test_gtplane\",\n                                    series_lookup=test_series_lookup)\n\n    # Opportunistic: attempt this once, covering both train and test, before any\n    # training starts. If it doesn't fit in available RAM, build_ram_cache logs\n    # why and returns False -- everything below keeps working unchanged via the\n    # per-study disk cache, exactly as it did before this was added.\n    build_ram_cache(train_ids_all + test_ids, {**train_index, **test_index})\n\n    fold_models = run_cross_validation(train_df, train_index)\n\n    if not fold_models:\n        log(\"no trained models available, keeping fallback submission\")\n        return\n\n    submission = predict_test(fold_models, test_ids, test_index)\n    submission.to_csv(OUTPUT_DIR / \"submission.csv\", index=False)\n    log(\"wrote final submission.csv\")\n\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T18:12:07.6436Z","iopub.execute_input":"2026-08-07T18:12:07.64386Z","execution_failed":"2026-08-07T18:15:34.715Z"}},"outputs":[],"execution_count":null}]}