{"cells":[{"cell_type":"markdown","metadata":{},"source":"# RSNA Knee 0.920 - one-mount rank ensemble\n\nThis is the whole inference pipeline behind a **0.920 public leaderboard score** on\nRSNA Knee Abnormality Detection. That number is a submission we actually fired\n(ref 55616077, 2026-08-19), not a cross-validation estimate and not a projection.\nRun the notebook top to bottom and it writes `submission.csv`.\n\nEverything it needs is one public dataset mount, the competition data, and the\nDINOv2-small model. Nothing is trained here - every weight is loaded from a\npublic assets pack. It is a long run: stage 1 alone is guarded by an eight-hour\nwall clock, and the verified submission ran on a T4 GPU session.\n\n## How it works\n\nThree families of models see the same DICOM studies, and they are combined by\n**rank averaging** rather than by averaging probabilities. Each member's\nper-target predictions are turned into percentile ranks first, then mixed. Ranks\nmake differently-calibrated members comparable, which is what lets a\ntwenty-checkpoint transformer branch, a five-fold transformer branch and two banks\nof RadImageNet attention heads be blended with flat weights and still behave.\n\n1. **Stage 1 - DINOv2 branch.** DICOM series are sorted by plane, fat suppression,\n   slice order and laterality into six slots at 336 px. Twenty checkpoints run over\n   slice windows; each checkpoint's predictions become percentile ranks and are\n   averaged equally.\n2. **Stage 2 - DINOv3 folds.** Five fold models run on their own six-slot 336 px\n   view. Their fold ranks are averaged, then mixed with stage 1 to form the\n   transformer parent: `0.55 * DINOv2 + 0.45 * DINOv3`.\n3. **Stage 3 - RadImageNet branch.** One shared RadImageNet ResNet-50 encoder\n   extracts 2048-dimensional slice features. Five reference heads run on the\n   three-plane E10 layout and five E13 heads run on a four-slot fat-sensitive\n   layout; the two are mixed `0.50 / 0.50` and re-ranked. For ten of the twelve\n   targets the parent is then mixed `0.50 / 0.50` with that Rad rank - Baker's cyst\n   and Fracture skip that mix and keep the transformer parent. Finally the same five E13\n   heads run once more on the E11 slot layout, and the answer is\n   `0.85 * E10 rank + 0.15 * second-pass rank` across all twelve targets.\n\nThe one artifact written is `submission.csv`.\n\n## Credits\n\nThis is a documented reproduction of public community work. The modelling is not\nours.\n\n- **Mattia Angeli** - the 0.917 ensemble this lineage descends from:\n  [Bend the Knee to DINOv3 - Ensembled](https://www.kaggle.com/code/mattiaangeli/bend-the-knee-to-dinov3-ensembled).\n  Our reproduction records its source as `mattiaangeli/bend-the-knee-to-dinov3-the-original`,\n  version 78.\n- **tonylica** - the consolidated assets pack every weight below is loaded from,\n  [rsna-knee-bend-dinov3-0917-repro-assets](https://www.kaggle.com/datasets/tonylica/rsna-knee-bend-dinov3-0917-repro-assets),\n  plus four of the trained models in the ensemble.\n- **pilkwang** - twenty of the trained models here, and one of the label sets read\n  from the reports.\n- **stevenleehans** and **lixin73** - two further report-derived label sets, so one\n  reading could be checked against another.\n- **marwanmath** - the official RadImageNet ResNet-50 weights.\n- **prvsiyan** - the notebook the original was forked from (Apache 2.0).\n- **cf696666 (fishface)** - the earliest verified E10/E11 RadImageNet machinery,\n  including the four-slot E11 design.\n- **romantamrazov** - the fold-rank aggregation idea.\n- **ieshanmeghani** - isolating the public v15/E10 RadImageNet delta.\n- **sofiaanjenje** - running and publishing the five-fold E11 and E13 RadImageNet\n  heads, including the E13 fat-sensitive crop arm used here.\n- **saidmohamedomary** - publishing V48, which surfaced the value of adding the E13\n  arm to the deployed RadImageNet block.\n- **sakhawathossen** and **ranjithragavan07** - publishing the no-extra-pass\n  fold-balanced/smooth DINO aggregation that was tested against this.\n\nIf a contributor is missing from that list it is an oversight carried forward from\nthe source notebook, not a claim."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# ============================================================================\n# CONFIG - the whole knob set.\n#\n# Every default below is the value that produced the 0.920 public-LB submission\n# (ref 55616077, 2026-08-19). Change one at a time. Each comment says what we\n# actually measured about that knob, including where we measured nothing.\n# ============================================================================\n\nSMOKE = False\n# False is the scoring run. True resolves the mounts, reports what is present,\n# writes a schema-valid placeholder submission.csv at 0.5, and skips every\n# inference stage. It takes about a minute and scores nothing - it exists so a\n# fork can prove it is wired up before spending a full GPU session.\n\nTIME_BUDGET_HOURS = 8.0\n# Wall-clock guard for stage 1. When it runs out the stage surrenders members\n# rather than dying, and then refuses to hand on a partial ensemble.\n# Measured on the pilkwang baseline kernel: 3h -> 5h produced a byte-identical\n# submission and the same 0.891, so that kernel is not budget-throttled\n# (ref 55613482). We have never run that test on THIS pipeline, so treat 8.0 as\n# documented rather than measured.\n\nDINOV3_RANK_WEIGHT = 0.45\n# Transformer parent = (1 - w) * DINOv2 rank + w * DINOv3 fold rank, w = 0.45.\n# Shipped value, carried over unchanged. We did not ablate this one.\n\nRAD_ALPHA = 0.50\n# Stage 3, for ten of the twelve targets:\n#   (1 - RAD_ALPHA) * transformer-parent rank + RAD_ALPHA * RadImageNet rank.\n# Baker's cyst and Fracture are excluded and keep the transformer parent.\n\nRAD_E13_MEMBER_WEIGHT = 0.50\n# Inside the RadImageNet branch:\n#   (1 - w) * reference-head rank + w * E13-head rank.\n# Our nested CV preferred RAD_ALPHA 0.35 with this at 0.65. The leaderboard did\n# not: that pair scored 0.919 (ref 55620998). The flat 0.50 / 0.50 defaults are\n# the measured winner, so offline CV was not a reliable guide for these two.\n\nPASS2_ENABLED = True\nPASS2_WEIGHT = 0.15\n# The second pass re-runs the five E13 heads on the E11 slot layout and mixes\n#   (1 - PASS2_WEIGHT) * E10 rank + PASS2_WEIGHT * second-pass rank\n# across all twelve targets. Turning it off scored 0.919 (ref 55636360), so the\n# stage is worth about +0.001 in exchange for one extra encode pass over the\n# test set. Turn it off if you need the time back."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Mounts and a preflight report. The paths below are the layout the 0.920 run\n# used; each falls back to the older /kaggle/input/<slug> layout so a fork with a\n# different mount style resolves instead of dying on the first read. Nothing here\n# raises - the stages fail loudly on their own if something is genuinely absent.\nimport json\nfrom pathlib import Path\n\nimport pandas as pd\nimport torch\n\n\ndef _resolve(label, primary, alternates=()):\n    for cand in (primary, *alternates):\n        p = Path(cand)\n        if p.exists():\n            note = '' if str(p) == str(primary) else '   (fallback; verified run used ' + str(primary) + ')'\n            print(f'  {label:<7} {p}{note}')\n            return p\n    tried = ', '.join([str(primary), *[str(a) for a in alternates]])\n    print(f'  {label:<7} NOT FOUND - tried {tried}')\n    return Path(primary)\n\n\ndef _preflight():\n    print('mounts')\n    asset = _resolve(\n        'ASSET',\n        '/kaggle/input/datasets/tonylica/rsna-knee-bend-dinov3-0917-repro-assets',\n        ['/kaggle/input/rsna-knee-bend-dinov3-0917-repro-assets'],\n    )\n    root = _resolve(\n        'COMP',\n        '/kaggle/input/competitions/rsna-knee-abnormality-detection',\n        ['/kaggle/input/rsna-knee-abnormality-detection'],\n    )\n    dino = _resolve(\n        'DINOv2',\n        '/kaggle/input/models/metaresearch/dinov2/pytorch/small/1',\n        [\n            '/kaggle/input/models/metaresearch/dinov2/PyTorch/small/1',\n            '/kaggle/input/dinov2/pytorch/small/1',\n        ],\n    )\n    print()\n    print('inputs the stages will open')\n    needed = [\n        ('stage 1  DINOv2 package', asset / 'rsna-knee-weights' / 'manifest.json'),\n        ('stage 2  DINOv3 folds', asset / 'knee-mri-fold-weights'),\n        ('stage 3  RadImageNet encoder', asset / 'resnet-50-radimagenet-marwan' / 'ResNet50.pt'),\n        ('stage 3  reference heads', asset / 'rsna-knee-e9-radimagenet-heads-v15' / 'v52_radimagenet_heads.pt'),\n        ('stage 3  E13 heads', asset / 'kernel-sources' / 'rsna-knee-e13-train' / 'rsna_rad_e11' / 'v52_e11_heads.pt'),\n    ]\n    missing = 0\n    for label, path in needed:\n        present = path.exists()\n        flag = 'ok     ' if present else 'MISSING'\n        missing += 0 if present else 1\n        print(f'  {flag}  {label:<30} {path}')\n    try:\n        n_members = len(json.loads((asset / 'rsna-knee-weights' / 'manifest.json').read_text())['members'])\n    except Exception as exc:\n        n_members = 'unreadable (' + type(exc).__name__ + ')'\n    folds = asset / 'knee-mri-fold-weights'\n    n_folds = len(sorted(folds.glob('*_f*.pt'))) if folds.is_dir() else 0\n    print(f'  DINOv2 checkpoints in the manifest : {n_members}')\n    print(f'  DINOv3 fold checkpoints on disk    : {n_folds}')\n    print()\n    print(f'  torch {torch.__version__}, {torch.cuda.device_count()} CUDA device(s)')\n    for i in range(torch.cuda.device_count()):\n        print(f'    cuda:{i}  {torch.cuda.get_device_name(i)}')\n    try:\n        print(f'  test studies: {len(pd.read_csv(root / \"test.csv\")):,}')\n    except Exception as exc:\n        print(f'  test.csv unreadable: {type(exc).__name__}: {exc}')\n    return asset, root, dino, missing\n\n\nASSET, ROOT, DINO, MISSING_INPUTS = _preflight()\nCOMP = ROOT\n\nif SMOKE:\n    _sample = pd.read_csv(ROOT / 'sample_submission.csv')\n    _sample[[c for c in _sample.columns if c != 'StudyInstanceUID']] = 0.5\n    Path('/kaggle/working').mkdir(parents=True, exist_ok=True)\n    _sample.to_csv('/kaggle/working/submission.csv', index=False)\n    print()\n    print(f'SMOKE=True: wrote a {len(_sample)}-row placeholder submission.csv at 0.5; inference skipped.')\n    print(f'            {MISSING_INPUTS} required input(s) missing. Set SMOKE=False for the scoring run.')"},{"cell_type":"markdown","metadata":{},"source":"### Stage 1 - DINOv2 branch\n\nReads the test DICOMs, builds the six-slot 336 px cache, runs every DINOv2\ncheckpoint in the pack's manifest over slice windows, and writes the equal-weight\nrank mean to `submission.csv`. Twenty checkpoints in the run we verified - the\npreflight cell above prints the count it actually found. This is the long stage;\n`TIME_BUDGET_HOURS` guards it."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"from __future__ import annotations\n\nif SMOKE:\n    print('SMOKE=True: stage 1 (DINOv2 branch) skipped.')\nelse:\n    import os\n    import gc\n    import hashlib\n    import json\n    import re\n    import time\n    import traceback\n    import threading\n    from concurrent.futures import ThreadPoolExecutor\n    from pathlib import Path\n    import numpy as np\n    import pandas as pd\n    import pydicom\n    import torch\n    import torch.nn as nn\n    import torch.nn.functional as F\n    # ASSET, ROOT and DINO are resolved in the CONFIG/preflight cell above.\n    T0 = time.time()\n    DEVS = [torch.device(f'cuda:{i}') for i in range(torch.cuda.device_count())]\n    SEED = 2026\n    TARGETS = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture']\n    CROP_MM = 130.0\n    CACHE_IMG = 336\n    GROUP = 3\n    N_GROUP_MAX = 1\n    CACHE_FRACTION = 0.45\n    CACHE_BUDGET_MAX_GB = 24.0\n    CACHE_BUDGET_GB = 12.0\n    TEST_SHARE = 0.3\n    HDR_THREADS = 16\n    PIX_THREADS = 12\n    ORDER_THREADS = 32\n    ORDER_BUDGET_S = 5400\n    AUG_ROT_DEG = 8.0\n    AUG_SCALE = 0.08\n    AUG_SHIFT = 0.05\n    AUG_INTENSITY = 0.1\n    LAT_MIN_OFFSET_MM = 20.0\n    SLICE_BAND = (0.2, 0.8)\n    RULES_NATIVE = {'order': 'normal', 'lat': 'centre', 'slot_fallback': False, 'decode_fill': 'nearest'}\n    RULES_LEGACY = {'order': 'dominant_axis', 'lat': 'corner_x', 'slot_fallback': True, 'decode_fill': 'zero'}\n    RULES = dict(RULES_NATIVE)\n    LEGACY_LAT_OFFSET_MM = 5.0\n    EVAL_BATCH = 8\n    TIME_BUDGET = TIME_BUDGET_HOURS * 3600.0\n    SLOTS_RECOVERED = [('SAG_FLUID_FS', 'Sagittal', True, True), ('COR_FLUID_FS', 'Coronal', True, True), ('AX_FLUID_FS', 'Axial', True, True), ('SAG_FLUID_NOFS', 'Sagittal', True, False), ('COR_T1', 'Coronal', False, False), ('SAG_T1', 'Sagittal', False, False)]\n    SLOTS_PUBLIC = [('SAG_FLUID', 'Sagittal', None, True), ('COR_FLUID', 'Coronal', None, True), ('AX_FLUID', 'Axial', None, True), ('SAG_STRUCT', 'Sagittal', None, False), ('COR_STRUCT', 'Coronal', None, False), ('AX_STRUCT', 'Axial', None, False)]\n    SLOT_SCHEME = os.environ.get('SLOT_SCHEME', 'recovered')\n    SLOTS = SLOTS_PUBLIC if SLOT_SCHEME == 'public' else SLOTS_RECOVERED\n    N_SLOT = len(SLOTS)\n    POOL_PARTS = {'cls_mean': 2, 'cls_mean_focal': 3}\n    SLOT_PRIOR_TABLE = {'ACL': (0, 3, 5), 'MCL': (1, 4), 'Medial Meniscus': (0, 1, 3, 4), 'Lateral Meniscus': (0, 1, 3, 4), 'Medial OA': (1, 4, 5), 'Lateral OA': (1, 4, 5), 'PF OA': (0, 2, 5), 'Effusion': (0, 2), 'Synovitis': (0, 2), \"Baker's\": (0,), 'Contusion': (0, 1, 2), 'Fracture': (0, 1, 2, 4, 5)}\n    SLOT_PRIOR_STRENGTH = 0.55\n    FATSAT_OPTS = {'FS', 'FATSAT', 'FAT_SAT', 'FSAT'}\n    _SEP = re.compile('[_\\\\-.]')\n    _FATSAT_RX = re.compile('\\\\bfs\\\\b|fatsat|fat sat|\\\\bstir\\\\b|\\\\bspair\\\\b|\\\\bspir\\\\b|\\\\bwe\\\\b|water excit|\\\\btirm\\\\b|\\\\bsting\\\\b|\\\\bfatsup\\\\b')\n    _T1_RX = re.compile('\\\\bt1\\\\b|\\\\bt1w\\\\b')\n    _T2_RX = re.compile('\\\\bt2\\\\b|\\\\bt2w\\\\b')\n    _PD_RX = re.compile('\\\\bpd\\\\b|\\\\bpdw\\\\b|proton|\\\\bdp\\\\b|dens')\n\n    def log(msg):\n        print(f'[{time.time() - T0:7.1f}s] {msg}', flush=True)\n    IMG = CACHE_IMG\n\n    def available_gb():\n        try:\n            with open('/proc/meminfo') as fh:\n                info = {k.strip(): v for k, v in (l.split(':', 1) for l in fh if ':' in l)}\n            return int(info['MemAvailable'].split()[0]) / 1024 ** 2\n        except Exception:\n            return CACHE_BUDGET_GB / CACHE_FRACTION\n\n    def plan_cache(n_study, n_test=0):\n        avail = available_gb()\n        budget = min(avail * CACHE_FRACTION, CACHE_BUDGET_MAX_GB)\n        n_total = n_study + max(n_test, int(TEST_SHARE * n_study))\n        per_slice = n_total * N_SLOT * IMG * IMG\n        afford = int(budget * 1024 ** 3 // max(per_slice, 1))\n        groups = max(1, min(N_GROUP_MAX, afford // GROUP))\n        log(f'memory: {avail:.1f} GB available, {budget:.1f} GB to the cache; sizing for {n_study} train + {n_total - n_study} test studies -> {groups} group(s) of {GROUP} = {groups * GROUP} slices per slot' + (f' (wanted {N_GROUP_MAX})' if groups < N_GROUP_MAX else ''))\n        return groups\n    N_GROUP = plan_cache(len(pd.read_csv(ROOT / 'train.csv')), len(pd.read_csv(ROOT / 'test.csv')))\n    CACHE_SLICES = GROUP * N_GROUP\n    HDR_TAGS = ['SeriesDescription', 'SequenceName', 'ScanOptions', 'ScanningSequence', 'RepetitionTime', 'EchoTime', 'Laterality', 'PixelSpacing', 'Rows', 'Columns', 'RescaleSlope', 'RescaleIntercept', 'ImagePositionPatient', 'ImageOrientationPatient']\n\n    def _hdr_vec(s, n):\n        if not isinstance(s, str):\n            return None\n        try:\n            v = [float(x) for x in s.split('|')]\n        except ValueError:\n            return None\n        return np.array(v) if len(v) >= n else None\n\n    def side_from_geometry(h):\n        cx = {}\n        for r in h.itertuples(index=False):\n            ipp = _hdr_vec(getattr(r, 'ImagePositionPatient', None), 3)\n            iop = _hdr_vec(getattr(r, 'ImageOrientationPatient', None), 6)\n            ps = _hdr_vec(getattr(r, 'PixelSpacing', None), 2)\n            rows, cols = (getattr(r, 'Rows', None), getattr(r, 'Columns', None))\n            if ipp is None or iop is None or ps is None or (not rows) or (not cols):\n                continue\n            try:\n                c = ipp[:3] + iop[:3] * ps[1] * float(cols) / 2 + iop[3:6] * ps[0] * float(rows) / 2\n            except (TypeError, ValueError):\n                continue\n            cx.setdefault(r.StudyInstanceUID, []).append(float(c[0]))\n        out = {}\n        for st, xs in cx.items():\n            m = float(np.median(xs))\n            out[st] = None if abs(m) < LAT_MIN_OFFSET_MM else 'R' if m < 0 else 'L'\n        return out\n\n    def side_from_corner_x(h):\n        out = {}\n        for st, g in h.groupby('StudyInstanceUID'):\n            xs = []\n            for r in g.itertuples(index=False):\n                ipp = _hdr_vec(getattr(r, 'ImagePositionPatient', None), 3)\n                if ipp is not None and np.isfinite(ipp).all():\n                    xs.append(float(ipp[0]))\n            if not xs:\n                out[st] = None\n                continue\n            x = float(np.median(xs))\n            out[st] = None if abs(x) < LEGACY_LAT_OFFSET_MM else 'R' if x < 0 else 'L'\n        return out\n\n    def lat_of(h, tag=''):\n        geo = side_from_corner_x(h) if RULES['lat'] == 'corner_x' else side_from_geometry(h)\n        d, n_tag, n_geo, n_none, n_disagree = ({}, 0, 0, 0, 0)\n        for st, g in h.groupby('StudyInstanceUID'):\n            v = [str(x).strip().upper() for x in g['Laterality'].dropna()]\n            if RULES['lat'] == 'corner_x' and 'ImageLaterality' in g.columns:\n                v += [str(x).strip().upper() for x in g['ImageLaterality'].dropna()]\n            v = [x[0] for x in v if x and x[0] in ('L', 'R')]\n            side = v[0] if v else None\n            if side is not None:\n                n_tag += 1\n                if geo.get(st) is not None and geo[st] != side:\n                    n_disagree += 1\n            else:\n                side = geo.get(st)\n                n_geo += side is not None\n                n_none += side is None\n            d[st] = side\n        log(f'{tag}laterality: {n_tag} from the tag, {n_geo} from geometry, {n_none} unresolved; tag and geometry disagree on {n_disagree} ({n_disagree / max(n_tag, 1):.1%} of the tagged)')\n        return d\n\n    def probe(item):\n        split, study, series, path = item\n        row = {'split': split, 'StudyInstanceUID': study, 'SeriesInstanceUID': series, 'dir': path}\n        try:\n            files = sorted((e.name for e in os.scandir(path) if e.name.endswith('.dcm')))\n            row['files'] = files\n            row['n_slices'] = len(files)\n            if not files:\n                return row\n            ds = pydicom.dcmread(os.path.join(path, files[len(files) // 2]), stop_before_pixels=True, force=True)\n            for t in HDR_TAGS:\n                v = getattr(ds, t, None)\n                if v is None:\n                    row[t] = None\n                elif isinstance(v, (list, tuple)) or type(v).__name__ == 'MultiValue':\n                    row[t] = '|'.join((str(x) for x in v))\n                else:\n                    row[t] = str(v)\n        except Exception as exc:\n            row['err'] = str(exc)[:120]\n        return row\n\n    def walk(split):\n        base = ROOT / split\n        items = []\n        if not base.is_dir():\n            return pd.DataFrame(columns=['split', 'StudyInstanceUID', 'SeriesInstanceUID', 'dir', 'files', 'n_slices'] + HDR_TAGS)\n        for study in os.scandir(base):\n            if study.is_dir():\n                for series in os.scandir(study.path):\n                    if series.is_dir():\n                        items.append((split, study.name, series.name, series.path))\n        with ThreadPoolExecutor(max_workers=HDR_THREADS) as pool:\n            rows = list(pool.map(probe, items))\n        return pd.DataFrame(rows)\n\n    def annotate(df):\n        desc = df['SeriesDescription'].fillna('') + ' ' + df['SequenceName'].fillna('')\n        desc = desc.str.lower().str.replace(_SEP, ' ', regex=True)\n        opts = df['ScanOptions'].fillna('').str.upper().str.split('|')\n        opts_fs = opts.apply(lambda ts: any((t.strip() in FATSAT_OPTS for t in ts)))\n        df['fatsat'] = desc.str.contains(_FATSAT_RX) | opts_fs\n        tr = pd.to_numeric(df['RepetitionTime'], errors='coerce')\n        te = pd.to_numeric(df['EchoTime'], errors='coerce')\n        gre = df['ScanningSequence'].fillna('').str.upper().str.contains('GR')\n        t1, t2, pdw = (desc.str.contains(_T1_RX), desc.str.contains(_T2_RX), desc.str.contains(_PD_RX))\n        df['weight'] = np.where(t1 & ~t2 & ~pdw, 'T1', np.where(t2 & ~pdw, 'T2', np.where(pdw, 'PD', np.where(gre, 'GRE', np.where(tr < 800, 'T1', np.where(te > 60, 'T2', np.where(tr >= 800, 'PD', 'UNK')))))))\n        df['fluid'] = np.isin(df['weight'], ['PD', 'T2'])\n        df['px'] = pd.to_numeric(df['PixelSpacing'].fillna('').str.split('|').str[0].replace('', np.nan), errors='coerce')\n        return df\n\n    def pick_slots(series_df, plane_map):\n        series_df = series_df.copy()\n        series_df['plane'] = series_df['SeriesInstanceUID'].map(plane_map)\n        out = {}\n        for study, g in series_df.groupby('StudyInstanceUID'):\n            chosen = {}\n            for name, plane, fluid, fs in SLOTS:\n                sel = (g['plane'] == plane) & (g['fatsat'] == fs)\n                if fluid is not None:\n                    sel &= g['fluid'] == fluid\n                cand = g[sel]\n                if len(cand) == 0 and RULES['slot_fallback'] and (fluid is False):\n                    cand = g[(g['plane'] == plane) & ~g['fatsat']]\n                if len(cand):\n                    chosen[name] = cand.sort_values('n_slices', ascending=False).iloc[0]\n            out[study] = chosen\n        return out\n    ORDER_TAGS = [(32, 50), (32, 55), (32, 19)]\n    DECODE_FAILED = []\n\n    def _natural_key(name):\n        return tuple((int(x) if x.isdigit() else x.lower() for x in re.split('(\\\\d+)', str(name))))\n\n    def _order_dominant_axis(rec):\n        files, d = (rec['files'], rec['dir'])\n        rows = []\n        for pos, f in enumerate(files):\n            ipp = inst = None\n            try:\n                ds = pydicom.dcmread(os.path.join(d, f), force=True, stop_before_pixels=True, specific_tags=['ImagePositionPatient', 'InstanceNumber'])\n                raw = getattr(ds, 'ImagePositionPatient', None)\n                if raw is not None and len(raw) >= 3:\n                    c = np.asarray(raw[:3], dtype=np.float64)\n                    if np.isfinite(c).all():\n                        ipp = c\n                n = getattr(ds, 'InstanceNumber', None)\n                if n is not None:\n                    inst = float(n)\n            except Exception:\n                pass\n            rows.append((f, ipp, inst, pos))\n        placed = [r for r in rows if r[1] is not None]\n        need = max(2, int(0.8 * len(rows)))\n        if len(placed) >= need:\n            xyz = np.stack([r[1] for r in placed])\n            axis = int(np.argmax(np.ptp(xyz, axis=0)))\n            spare = float(np.nanmedian(xyz[:, axis]))\n            rows.sort(key=lambda r: (float(r[1][axis]) if r[1] is not None else spare, r[2] if r[2] is not None else float('inf'), r[3]))\n        elif sum((r[2] is not None for r in rows)) >= need:\n            rows.sort(key=lambda r: (r[2] if r[2] is not None else float('inf'), r[3]))\n        else:\n            rows.sort(key=lambda r: _natural_key(r[0]))\n        return ([r[0] for r in rows], True)\n\n    def order_slices(rec):\n        if RULES['order'] == 'dominant_axis':\n            return _order_dominant_axis(rec)\n        files, d = (rec['files'], rec['dir'])\n        keyed = []\n        for f in files:\n            k = None\n            try:\n                ds = pydicom.dcmread(os.path.join(d, f), force=True, stop_before_pixels=True, specific_tags=ORDER_TAGS)\n                iop = np.asarray(ds.ImageOrientationPatient, dtype=float)\n                ipp = np.asarray(ds.ImagePositionPatient, dtype=float)\n                k = float(np.dot(ipp, np.cross(iop[:3], iop[3:])))\n            except Exception:\n                try:\n                    k = float(ds.InstanceNumber)\n                except Exception:\n                    k = None\n            keyed.append((k, f))\n        if any((k is None for k, _ in keyed)):\n            return (files, False)\n        return ([f for _, f in sorted(keyed, key=lambda t: t[0])], True)\n\n    def read_slot(rec, n_slice=None, out_size=None):\n        n_slice = GROUP if n_slice is None else n_slice\n        out_size = IMG if out_size is None else out_size\n        files, d, px = (rec.get('ordered') or rec['files'], rec['dir'], rec['px'])\n        n = len(files)\n        if n == 0:\n            return None\n        lo, hi = (int(SLICE_BAND[0] * (n - 1)), int(SLICE_BAND[1] * (n - 1)))\n        idx = np.unique(np.linspace(lo, hi, n_slice).astype(int)) if hi > lo else np.array([n // 2])\n        while len(idx) < n_slice:\n            idx = np.append(idx, idx[-1])\n        planes = []\n        for i in idx[:n_slice]:\n            try:\n                ds = pydicom.dcmread(os.path.join(d, files[int(i)]), force=True)\n                a = ds.pixel_array.astype(np.float32)\n                sl = float(getattr(ds, 'RescaleSlope', 1) or 1)\n                ic = float(getattr(ds, 'RescaleIntercept', 0) or 0)\n                a = a * sl + ic\n            except Exception:\n                a = None\n            planes.append(a)\n        got = [k for k, p in enumerate(planes) if p is not None]\n        if RULES['decode_fill'] == 'zero':\n            if not got:\n                DECODE_FAILED.append(rec.get('SeriesInstanceUID', d))\n            planes = [np.zeros((out_size, out_size), np.float32) if p is None else p for p in planes]\n            got = list(range(len(planes)))\n        if not got:\n            DECODE_FAILED.append(rec.get('SeriesInstanceUID', d))\n            return None\n        if len(got) < len(planes):\n            DECODE_FAILED.append(rec.get('SeriesInstanceUID', d))\n            for k, p in enumerate(planes):\n                if p is None:\n                    planes[k] = planes[min(got, key=lambda j: abs(j - k))]\n        shp = planes[0].shape\n        planes = [p if p.shape == shp else np.zeros(shp, np.float32) for p in planes]\n        vol = np.stack(planes)\n        if px and np.isfinite(px) and (px > 0):\n            want = int(round(CROP_MM / px))\n            h, w = shp\n            if 16 < want < min(h, w):\n                cy, cx = (h // 2, w // 2)\n                half = want // 2\n                vol = vol[:, max(0, cy - half):cy + half, max(0, cx - half):cx + half]\n        lo_v, hi_v = np.percentile(vol, [1, 99])\n        vol = np.clip((vol - lo_v) / max(hi_v - lo_v, 1e-06), 0, 1)\n        t = torch.from_numpy(np.ascontiguousarray(vol)).unsqueeze(0)\n        t = F.interpolate(t, size=(out_size, out_size), mode='bilinear', align_corners=False)\n        return (t.squeeze(0) * 255).round().clamp(0, 255).to(torch.uint8)\n\n    def normalise_laterality(img, plane, lat):\n        if lat != 'R':\n            return img\n        if plane in ('Coronal', 'Axial'):\n            return torch.flip(img, dims=[-1])\n        return torch.flip(img, dims=[0])\n    ORDER_CACHE = os.environ.get('RSNA_ORDER_CACHE') or None\n\n    def build_cache(slot_map, plane_map, lat_map, tag):\n        studies = sorted(slot_map)\n        sidx = {s: i for i, s in enumerate(studies)}\n        cache = np.zeros((len(studies), N_SLOT, CACHE_SLICES, IMG, IMG), np.uint8)\n        mask = np.zeros((len(studies), N_SLOT), np.float32)\n        log(f'{tag}: cache {cache.shape} = {cache.nbytes / 1024 ** 3:.1f} GB')\n        jobs = [(st, k, plane, slot_map[st][name]) for st in studies for k, (name, plane, _, _) in enumerate(SLOTS) if name in slot_map[st]]\n        n_job = len(jobs)\n        t_ord = time.time()\n        n_slice_total = sum((len(j[3]['files']) for j in jobs))\n        log(f'{tag}: ordering {len(jobs)} slot-series ({n_slice_total} slice headers)')\n        ok = done = 0\n        CHUNK_O = 1024\n        seen = {}\n        if ORDER_CACHE and Path(ORDER_CACHE).is_file():\n            try:\n                import json as _json\n                seen = _json.loads(Path(ORDER_CACHE).read_text())\n            except (OSError, ValueError):\n                seen = {}\n            hit = 0\n            for _, _, _, rec in jobs:\n                e = seen.get(rec['SeriesInstanceUID'])\n                if e and len(e['files']) == len(rec['files']):\n                    rec['ordered'] = e['files']\n                    ok += int(e['good'])\n                    hit += 1\n            jobs = [j for j in jobs if 'ordered' not in j[3]]\n            log(f'{tag}: {hit} slot-series ordered from {ORDER_CACHE}, {len(jobs)} to read')\n        with ThreadPoolExecutor(max_workers=ORDER_THREADS) as pool:\n            for c0 in range(0, len(jobs), CHUNK_O):\n                block = jobs[c0:c0 + CHUNK_O]\n                for (_, _, _, rec), (files, good) in zip(block, pool.map(lambda j: order_slices(j[3]), block)):\n                    rec['ordered'] = files\n                    ok += int(good)\n                    done += 1\n                    if ORDER_CACHE:\n                        seen[rec['SeriesInstanceUID']] = {'files': files, 'good': bool(good)}\n                budget = min(ORDER_BUDGET_S, max(60.0, (TIME_BUDGET - (time.time() - T0)) * 0.35))\n                if time.time() - t_ord > budget:\n                    log(f'{tag}: ordering budget spent at {done}/{len(jobs)}; the rest keep file order')\n                    break\n        if ORDER_CACHE and done:\n            import json as _json\n            _t = Path(ORDER_CACHE).with_suffix('.tmp')\n            _t.write_text(_json.dumps(seen))\n            _t.replace(Path(ORDER_CACHE))\n        log(f'{tag}: ordered {ok}/{n_job} by geometry ({n_job - ok} kept arbitrary) in {time.time() - t_ord:.0f}s')\n        jobs = [(st, k, plane, slot_map[st][name]) for st in studies for k, (name, plane, _, _) in enumerate(SLOTS) if name in slot_map[st]]\n        log(f'{tag}: decoding {len(jobs)} slot-series')\n        n_failed_before = len(DECODE_FAILED)\n        CHUNK = 512\n        done = 0\n        with ThreadPoolExecutor(max_workers=PIX_THREADS) as pool:\n            for c0 in range(0, len(jobs), CHUNK):\n                block = jobs[c0:c0 + CHUNK]\n                for (st, k, plane, _), img in zip(block, pool.map(lambda j: read_slot(j[3], CACHE_SLICES, IMG), block)):\n                    done += 1\n                    if img is None:\n                        continue\n                    cache[sidx[st], k] = normalise_laterality(img, plane, lat_map.get(st)).numpy()\n                    mask[sidx[st], k] = 1.0\n                if done % 4096 < CHUNK:\n                    log(f'  {tag} {done}/{len(jobs)}')\n                if time.time() - T0 > TIME_BUDGET:\n                    log(f'  {tag}: time budget reached during decode')\n                    break\n        n_failed = len(DECODE_FAILED) - n_failed_before\n        log(f'{tag}: {int(mask.sum())}/{len(jobs)} slots filled' + (f'; {n_failed} series had a slice that would not decode' if n_failed else ''))\n        gc.collect()\n        return (studies, cache, mask)\n\n    class SlotHead(nn.Module):\n\n        def __init__(self, dim, n_slot, n_out, hidden=256, p=0.2, prior=False):\n            super().__init__()\n            self.proj = nn.Sequential(nn.LayerNorm(dim), nn.Linear(dim, hidden), nn.GELU())\n            self.slot_emb = nn.Parameter(torch.randn(n_slot, hidden) * 0.02)\n            self.query = nn.Parameter(torch.randn(n_out, hidden) * 0.02)\n            self.drop = nn.Dropout(p)\n            self.out = nn.Linear(hidden, n_out)\n            self.hidden = hidden\n            p_ = torch.zeros(n_out, n_slot)\n            if prior and n_slot == len(SLOTS) and (n_out == len(TARGETS)):\n                for t, slots in SLOT_PRIOR_TABLE.items():\n                    if t in TARGETS:\n                        p_[TARGETS.index(t), list(slots)] = SLOT_PRIOR_STRENGTH\n            self.prior = prior\n            if prior:\n                self.register_buffer('slot_prior', p_)\n\n        def forward(self, x, mask):\n            h = self.proj(x) + self.slot_emb\n            att = torch.einsum('bsh,oh->bos', h, self.query) / self.hidden ** 0.5\n            if self.prior:\n                att = att + self.slot_prior.unsqueeze(0)\n            att = att.masked_fill(mask.unsqueeze(1) < 0.5, -10000.0).softmax(-1)\n            ctx = self.drop(torch.einsum('bos,bsh->boh', att, h))\n            return (ctx * self.out.weight.unsqueeze(0)).sum(-1) + self.out.bias\n\n    class Model(nn.Module):\n\n        def __init__(self, backbone, dim, pool='cls_mean', prior=False):\n            super().__init__()\n            self.backbone = backbone\n            self.pool = pool\n            self.head = SlotHead(dim * POOL_PARTS[pool], N_SLOT, len(TARGETS), prior=prior)\n            self.register_buffer('mean', torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1))\n            self.register_buffer('std', torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1))\n\n        def forward(self, imgs, mask, img_size=None):\n            B, S = imgs.shape[:2]\n            x = imgs.reshape(B * S, *imgs.shape[2:]).float().div_(255.0)\n            if img_size is not None and img_size != x.shape[-1]:\n                x = F.interpolate(x, size=(img_size, img_size), mode='bilinear', align_corners=False)\n            x = (x - self.mean) / self.std\n            out = self.backbone(pixel_values=x).last_hidden_state\n            patch = out[:, 1:]\n            parts = [out[:, 0], patch.mean(1)]\n            if self.pool == 'cls_mean_focal':\n                k = max(1, patch.shape[1] // 8)\n                parts.append(patch.topk(k, dim=1).values.mean(1))\n            feat = torch.cat(parts, dim=1).reshape(B, S, -1)\n            return self.head(feat, mask)\n\n    def build_model(unfreeze_last, source=None, variant='small', pool='cls_mean', prior=False):\n        from transformers import AutoModel\n        p = source if source is not None else find_dinov2(variant)\n        if p is None:\n            raise FileNotFoundError('DINOv2 weights not attached')\n        bb = AutoModel.from_pretrained(str(p))\n        n_layer = len(bb.encoder.layer)\n        for prm in bb.parameters():\n            prm.requires_grad = False\n        for blk in bb.encoder.layer[max(0, n_layer - unfreeze_last):]:\n            for prm in blk.parameters():\n                prm.requires_grad = True\n        for prm in bb.layernorm.parameters():\n            prm.requires_grad = True\n        dim = bb.config.hidden_size\n        trainable = sum((p.numel() for p in bb.parameters() if p.requires_grad))\n        log(f'backbone: {n_layer} blocks, last {unfreeze_last} trainable ({trainable / 1000000.0:.1f}M params), feature dim {dim * POOL_PARTS[pool]}')\n        return Model(bb, dim, pool=pool, prior=prior)\n    FINGERPRINT_TOL = 0.002\n\n    def fingerprint(model, dev, img_size, n_slot=None, group=None, seed=None):\n        n_slot = N_SLOT if n_slot is None else n_slot\n        group = GROUP if group is None else group\n        seed = SEED if seed is None else seed\n        g = torch.Generator().manual_seed(seed)\n        imgs = torch.randint(0, 256, (2, n_slot, group, img_size, img_size), generator=g, dtype=torch.uint8).to(dev)\n        mask = torch.ones(2, n_slot, device=dev)\n        mask[1, -1] = 0.0\n        was_training = model.training\n        model.eval()\n        with torch.no_grad():\n            out = model(imgs, mask, img_size).float().cpu().numpy()\n        if was_training:\n            model.train()\n        return out\n\n    def check_fingerprint(model, dev, img_size, expected, tol=FINGERPRINT_TOL, tag=''):\n        got = fingerprint(model, dev, img_size)\n        exp = np.asarray(expected, np.float32)\n        if got.shape != exp.shape:\n            raise WeightsError(f'{tag}fingerprint shape {got.shape} != stored {exp.shape}: the architecture is not the one these weights were fitted to')\n        d = float(np.abs(got - exp).max())\n        if d > tol:\n            raise WeightsError(f'{tag}fingerprint differs by {d:.4g} (tolerance {tol:g}). The weights load but do not compute what they computed when fitted - preprocessing, resolution or architecture has moved between the two runs.')\n        log(f'{tag}fingerprint matches within {d:.2g}')\n        return d\n\n    class WeightsError(RuntimeError):\n        pass\n    TTA_OVERLAP = True\n    TTA_POOL = 'prob'\n    PUBLIC_FRONTIER_TARGET_POOL = {'Fracture': 'max', 'Contusion': 'max', 'Medial Meniscus': 'max', 'Lateral Meniscus': 'max', 'ACL': 'top2', 'MCL': 'top2', \"Baker's\": 'max'}\n    TTA_TARGET_POOL = {**PUBLIC_FRONTIER_TARGET_POOL, 'Synovitis': 'original_mean'}\n    LEGACY_FOLD_SOFTPOOL_BETA = {'ACL': 6.0, 'MCL': 6.0, 'Medial Meniscus': 8.0, 'Lateral Meniscus': 8.0, \"Baker's\": 8.0, 'Contusion': 8.0, 'Fracture': 10.0}\n    LEGACY_FOLD_SOFTPOOL_ALPHA = {'ACL': 0.2, 'MCL': 0.2, 'Medial Meniscus': 0.25, 'Lateral Meniscus': 0.25, \"Baker's\": 0.2, 'Contusion': 0.2, 'Fracture': 0.15}\n\n    def window_starts(n_slice, group, overlap=None):\n        overlap = TTA_OVERLAP if overlap is None else overlap\n        if overlap and n_slice >= group:\n            return list(range(n_slice - group + 1))\n        return [g * group for g in range(max(n_slice // group, 1))]\n\n    def apply_target_window_pool(values, probs, logits, original_probs, mapping, target_idx):\n        for target, mode in mapping.items():\n            j = target_idx[target]\n            if mode == 'max':\n                values[:, j] = probs[:, :, j].max(0).values\n            elif mode == 'mean':\n                values[:, j] = probs[:, :, j].mean(0)\n            elif mode == 'logit_mean':\n                values[:, j] = torch.sigmoid(logits[:, :, j].mean(0))\n            elif mode == 'original_mean':\n                values[:, j] = original_probs[:, :, j].mean(0)\n            elif mode in ('top2', 'top3'):\n                k = min(int(mode[3:]), probs.shape[0])\n                values[:, j] = probs[:, :, j].topk(k, dim=0).values.mean(0)\n            else:\n                raise ValueError(f'unknown TTA pooling mode for {target}: {mode}')\n        return values\n\n    def legacy_fold_soft_window_pool(original_probs, target_idx):\n        values = original_probs.mean(0).clone()\n        for target, beta in LEGACY_FOLD_SOFTPOOL_BETA.items():\n            j = target_idx[target]\n            x = original_probs[:, :, j]\n            weight = torch.softmax(float(beta) * x, dim=0)\n            values[:, j] = (weight * x).sum(0)\n        return values\n\n    @torch.no_grad()\n    def predict_member(model, cache, mask, idx, dev, img_size, group=None, pool=None, starts=None, jitter=False, jitter_seed=SEED, return_public_frontier=False):\n        group = GROUP if group is None else group\n        pool = TTA_POOL if pool is None else pool\n        starts = window_starts(cache.shape[2], group) if starts is None else list(starts)\n        if not starts:\n            raise ValueError('predict_member was given no windows to average over')\n        target_idx = {t: j for j, t in enumerate(TARGETS)}\n        unknown = (set(TTA_TARGET_POOL) | set(PUBLIC_FRONTIER_TARGET_POOL)) - set(target_idx)\n        if unknown:\n            raise ValueError(f'unknown target(s) in TTA_TARGET_POOL: {unknown}')\n        jitter_gen = torch.Generator(device=dev)\n        jitter_gen.manual_seed(int(jitter_seed) % (2 ** 63 - 1))\n        model.eval()\n        out, public_frontier_out, public_soft_out = ([], [], [])\n        for b in range(0, len(idx), EVAL_BATCH):\n            sel = idx[b:b + EVAL_BATCH]\n            m = torch.from_numpy(mask[sel]).to(dev)\n            win_probs, win_logits, win_original_probs = ([], [], [])\n            for st in starts:\n                rows = torch.from_numpy(np.ascontiguousarray(cache[sel, :, st:st + group])).to(dev)\n                views = [rows] + ([augment(rows, generator=jitter_gen)] if jitter else [])\n                view_probs, view_logits = ([], [])\n                for view in views:\n                    with torch.autocast('cuda', enabled=dev.type == 'cuda'):\n                        z = model(view, m, img_size).float()\n                    view_logits.append(z)\n                    view_probs.append(torch.sigmoid(z))\n                win_logits.append(torch.stack(view_logits).mean(0))\n                win_probs.append(torch.stack(view_probs).mean(0))\n                win_original_probs.append(view_probs[0])\n            probs = torch.stack(win_probs)\n            logits = torch.stack(win_logits)\n            original_probs = torch.stack(win_original_probs)\n            v = torch.sigmoid(logits.mean(0)) if pool == 'logit' else probs.mean(0)\n            v = apply_target_window_pool(v, probs, logits, original_probs, TTA_TARGET_POOL, target_idx)\n            out.append(v.cpu().numpy())\n            if return_public_frontier:\n                public_v = apply_target_window_pool(original_probs.mean(0), original_probs, logits, original_probs, PUBLIC_FRONTIER_TARGET_POOL, target_idx)\n                public_frontier_out.append(public_v.cpu().numpy())\n                public_soft = legacy_fold_soft_window_pool(original_probs, target_idx)\n                public_soft_out.append(public_soft.cpu().numpy())\n        primary = np.concatenate(out) if out else np.zeros((0, len(TARGETS)), np.float32)\n        if not return_public_frontier:\n            return primary\n        public_frontier = np.concatenate(public_frontier_out) if public_frontier_out else np.zeros((0, len(TARGETS)), np.float32)\n        public_soft = np.concatenate(public_soft_out) if public_soft_out else np.zeros((0, len(TARGETS)), np.float32)\n        return (primary, public_frontier, public_soft)\n    BUILD_LOCK = threading.Lock()\n    STATE_LOCK = threading.Lock()\n\n    def _run_member(path, m, dev, Cte, Mte, idx, starts, jitter):\n        t0 = time.time()\n        with BUILD_LOCK:\n            if 'state' in m:\n                state, fp = (m['state'], None)\n            else:\n                ck = torch.load(Path(path) / m['file'], map_location='cpu', weights_only=False)\n                state, fp = (ck['model'], ck.get('fingerprint'))\n            model = build_model(int(m['config']['unfreeze_last']), variant=m['config']['variant'], pool=m['config'].get('pool', 'cls_mean'), prior=bool(m['config'].get('prior', False))).to(dev)\n            model.load_state_dict(state)\n            if fp is not None:\n                check_fingerprint(model, dev, IMG, fp, tag=f\"{m['id']}: \")\n            else:\n                log(f\"  {m['id']}: no stored fingerprint (legacy bundle) -- accepted at reduced weight\")\n        t_ready = time.time()\n        jitter_seed = SEED + int(hashlib.sha256(str(m['id']).encode()).hexdigest()[:8], 16)\n        public_member = 'state' not in m\n        predicted = predict_member(model, Cte, Mte, idx, dev, IMG, starts=starts, jitter=jitter, jitter_seed=jitter_seed, return_public_frontier=public_member)\n        if public_member:\n            p, public_p, public_soft = predicted\n        else:\n            p, public_p, public_soft = (predicted, None, None)\n        t_done = time.time()\n        del model, state\n        gc.collect()\n        if dev.type == 'cuda':\n            with torch.cuda.device(dev):\n                torch.cuda.empty_cache()\n        passes = len(starts) * (2 if jitter else 1)\n        return (p, public_p, public_soft, (t_ready - t0, (t_done - t_ready) / max(passes, 1)))\n\n    def _combine(per_member):\n        all_ids = sorted({s for m in per_member for s in m['ids']})\n        pos = {s: i for i, s in enumerate(all_ids)}\n        acc = np.zeros((len(all_ids), len(TARGETS)), np.float64)\n        tot = np.zeros(len(TARGETS), np.float64)\n        for m in per_member:\n            target_weight = m.get('target_weight')\n            w = np.asarray(target_weight if target_weight is not None else [float(m.get('weight', 1.0))] * len(TARGETS), dtype=np.float64)\n            if w.shape != (len(TARGETS),) or np.any(w < 0):\n                raise ValueError(f\"invalid target weights for {m.get('id')}: {w}\")\n            r = pd.DataFrame(m['pred']).rank(pct=True).to_numpy()\n            acc[[pos[s] for s in m['ids']]] += r * w[None, :]\n            tot += w\n        if np.any(tot <= 0):\n            raise ValueError(f'at least one target has no ensemble vote: {tot}')\n        return (all_ids, acc / tot[None, :])\n\n    def combine_public_members_by_fold(per_member, pred_key='pred'):\n        all_ids = sorted({study for member in per_member for study in member['ids']})\n        position = {study: i for i, study in enumerate(all_ids)}\n        groups = {}\n        for i, member in enumerate(per_member):\n            fold = member.get('fold')\n            key = f'fold_{fold}' if fold is not None else f'member_{i}'\n            groups.setdefault(key, []).append(member)\n        fold_ranks, diagnostics = ([], [])\n        for key, members_in_fold in sorted(groups.items()):\n            matrices = []\n            for member in members_in_fold:\n                values = np.full((len(all_ids), len(TARGETS)), np.nan, np.float64)\n                values[[position[study] for study in member['ids']]] = np.asarray(member[pred_key], np.float64)\n                if np.isnan(values).any():\n                    raise WeightsError(f\"{member.get('id')}: incomplete {pred_key} coverage\")\n                matrices.append(values)\n            raw_fold_mean = np.mean(matrices, axis=0)\n            fold_ranks.append(pd.DataFrame(raw_fold_mean).rank(method='average', pct=True).to_numpy(np.float64))\n            diagnostics.append({'ensemble_group': key, 'members': len(members_in_fold)})\n        if len(fold_ranks) != 5:\n            raise WeightsError(f'legacy branch requires five folds, found {len(fold_ranks)}')\n        return (all_ids, np.mean(fold_ranks, axis=0), pd.DataFrame(diagnostics))\n\n    def blend_legacy_frontier_and_soft(frontier_rank, soft_rank):\n        output = np.asarray(frontier_rank, np.float64).copy()\n        for j, target in enumerate(TARGETS):\n            alpha = float(LEGACY_FOLD_SOFTPOOL_ALPHA.get(target, 0.0))\n            if alpha:\n                output[:, j] = (1.0 - alpha) * frontier_rank[:, j] + alpha * soft_rank[:, j]\n        return output\n\n    def infer_from_package(path, dev=None):\n        man = json.loads((Path(path) / 'manifest.json').read_text())\n        members = man['members']\n        log(f'weights package: {len(members)} member(s) from {path}; {len(DEVS)} device(s)')\n        test_df = pd.read_csv(ROOT / 'test.csv')\n        test_series = pd.read_csv(ROOT / 'test_series.csv')\n        plane_map = dict(zip(test_series['SeriesInstanceUID'], test_series['Anatomical_Plane']))\n        hte = annotate(walk('test_series'))\n        log(f'test header pass: {len(hte)} series')\n        groups = {}\n        for m in members:\n            groups.setdefault(m['pixel_group'], []).append(m)\n        groups.update(legacy_group_members())\n        per_member, public_frontier_members = ([], [])\n        est = {'fixed': None, 'win': None}\n\n        def bank(m, ids, pred, starts, jitter, public_pred=None, public_soft=None):\n            if float(np.std(pred)) < 1e-09:\n                log(f\"  {m['id']}: degenerate predictions; not banked\")\n                return\n            with STATE_LOCK:\n                per_member.append({'id': m['id'], 'fold': m.get('fold'), 'ids': ids, 'pred': pred, 'weight': m.get('weight', 1.0), 'target_weight': m.get('target_weight'), 'holdout': m.get('holdout')})\n                if public_pred is not None and len(starts) == len(starts_full):\n                    if float(np.std(public_pred)) < 1e-09:\n                        raise WeightsError(f\"{m['id']}: degenerate public-frontier prediction\")\n                    public_frontier_members.append({'id': m['id'], 'fold': m.get('fold'), 'ids': ids, 'pred': public_pred, 'soft_pred': public_soft})\n                elif public_pred is not None:\n                    log(f\"  {m['id']}: public-frontier vote omitted because only {len(starts)} / {len(starts_full)} windows completed\")\n                all_ids, acc = _combine(per_member)\n                write_submission(acc, all_ids, test_df, 'submission.csv')\n                log(f\"  banked {m['id']} fold {m.get('fold', '?')} ({len(starts)} window(s){(', jitter' if jitter else '')}); submission.csv = weighted rank mean of {len(per_member)} member(s)\")\n        for gi, (key, gm) in enumerate(groups.items(), 1):\n            cfg = json.loads(key)\n            adopt_config_globals(cfg)\n            log(f\"decode group {gi}/{len(groups)}: {cfg['img']}px x {cfg['slices']} slices, crop {cfg['crop_mm']} mm -> {len(gm)} member(s)\")\n            st_te, Cte, Mte = build_cache(pick_slots(hte, plane_map), plane_map, lat_of(hte, 'test '), f'test g{gi}')\n            idx = np.arange(len(st_te))\n            starts_full = window_starts(Cte.shape[2], GROUP)\n            pending = sorted(gm, key=lambda m: -(m.get('holdout') or 0))\n            left_after = sum((len(g) for j, (_, g) in enumerate(groups.items(), 1) if j > gi))\n\n            def pop_next():\n                with STATE_LOCK:\n                    if not pending:\n                        return (None, None, False)\n                    left = TIME_BUDGET - (time.time() - T0)\n                    remaining = len(pending) + left_after\n                    slots_left = -(-remaining // len(DEVS))\n                    starts, jit = (starts_full, False)\n                    if est['fixed'] is not None and est['win'] is not None:\n                        afford = max(left * 0.9, 0.0)\n                        room = afford / max(slots_left, 1)\n                        if est['fixed'] + est['win'] > room:\n                            log(f'  {left / 60:.0f} min left: surrendering {len(pending)} member(s); not one more fits')\n                            pending.clear()\n                            return (None, None, False)\n                        jit = est['fixed'] + 2 * len(starts_full) * est['win'] <= room * 0.6\n                        per_win = est['win'] * (2 if jit else 1)\n                        n_win = int((room - est['fixed']) / per_win) if per_win > 0 else len(starts_full)\n                        n_win = max(1, min(len(starts_full), n_win))\n                        if n_win < len(starts_full):\n                            mid = (len(starts_full) - n_win) // 2\n                            starts = starts_full[mid:mid + n_win]\n                    return (pending.pop(0), starts, jit)\n\n            def worker(dev):\n                others = [d for d in DEVS if d is not dev]\n                while True:\n                    m, starts, jit = pop_next()\n                    if m is None:\n                        return\n                    for attempt, d in enumerate([dev] + others[:1]):\n                        try:\n                            p, public_p, public_soft, (fs, ws) = _run_member(path, m, d, Cte, Mte, idx, starts, jit)\n                            with STATE_LOCK:\n                                est['fixed'], est['win'] = (fs, ws)\n                            bank(m, st_te, p, starts, jit, public_p, public_soft)\n                            break\n                        except Exception as exc:\n                            log(f\"  MEMBER {m['id']} failed on {d} ({type(exc).__name__}: {exc}); \" + ('retrying on peer device' if attempt == 0 and others else 'dropped -- costs one vote, not the run'))\n                            if d.type == 'cuda':\n                                with torch.cuda.device(d):\n                                    torch.cuda.empty_cache()\n            threads = [threading.Thread(target=worker, args=(d,)) for d in DEVS]\n            for t in threads:\n                t.start()\n            for t in threads:\n                t.join()\n            del Cte, Mte\n            gc.collect()\n        if not per_member:\n            raise WeightsError('no member produced predictions; submission stays at 0.5')\n        all_ids, acc = _combine(per_member)\n        sub = write_submission(acc, all_ids, test_df, 'submission.csv')\n        log(f'final submission.csv = weighted rank mean of {len(per_member)} member(s); {sub.shape}; nulls {int(sub[TARGETS].isna().sum().sum())}')\n        if len(public_frontier_members) == len(members):\n            frontier_ids, frontier_acc = _combine(public_frontier_members)\n            frontier_sub = write_submission(frontier_acc, frontier_ids, test_df, 'submission_public_0899.csv')\n            log(f'submission_public_0899.csv = exact no-jitter public-frontier rank mean of {len(public_frontier_members)} member(s); {frontier_sub.shape}; nulls {int(frontier_sub[TARGETS].isna().sum().sum())}')\n            fold_ids, fold_frontier, fold_diagnostics = combine_public_members_by_fold(public_frontier_members, 'pred')\n            soft_ids, fold_soft, _ = combine_public_members_by_fold(public_frontier_members, 'soft_pred')\n            if fold_ids != soft_ids:\n                raise WeightsError('legacy hard/soft study order mismatch')\n            legacy_prediction = blend_legacy_frontier_and_soft(fold_frontier, fold_soft)\n            legacy_sub = write_submission(legacy_prediction, fold_ids, test_df, 'submission_legacy_fold_blend.csv')\n            fold_diagnostics.to_csv('legacy_fold_diagnostics.csv', index=False)\n            log(f'legacy DINO aggregation written from five folds; {legacy_sub.shape}')\n        else:\n            log(f'public-frontier fallback not emitted: {len(public_frontier_members)} / {len(members)} required public members completed')\n        return sub\n\n    def adopt_config_globals(cfg):\n        global IMG, CACHE_IMG, GROUP, CACHE_SLICES, N_GROUP, CROP_MM, SLICE_BAND, RULES\n        CACHE_IMG = IMG = int(cfg['img'])\n        GROUP = int(cfg['group'])\n        CACHE_SLICES = int(cfg['slices'])\n        N_GROUP = max(CACHE_SLICES // GROUP, 1)\n        CROP_MM = float(cfg['crop_mm'])\n        SLICE_BAND = tuple((float(x) for x in cfg['band']))\n        rules = cfg.get('rules') or RULES_NATIVE\n        unknown = {k: v for k, v in rules.items() if k not in RULES_NATIVE or v not in (RULES_NATIVE[k], RULES_LEGACY[k])}\n        if unknown:\n            raise WeightsError(f'the members record pixel rules this pipeline cannot reproduce: {unknown}')\n        RULES = {**RULES_NATIVE, **rules}\n        if [s[0] for s in SLOTS] != list(cfg['slots']):\n            raise WeightsError(f\"the members were fitted on slots {cfg['slots']} and this pipeline defines {[s[0] for s in SLOTS]}; a weight would be read against the wrong slot\")\n\n    def augment(imgs, generator=None):\n        lead = imgs.shape[:-3]\n        x = imgs.reshape(-1, *imgs.shape[-3:]).float()\n        n, dev = (x.shape[0], x.device)\n        rot = (torch.rand(n, device=dev, generator=generator) - 0.5) * 2 * (AUG_ROT_DEG * np.pi / 180)\n        sc = 1.0 + torch.rand(n, device=dev, generator=generator) * AUG_SCALE\n        tx = (torch.rand(n, device=dev, generator=generator) - 0.5) * 2 * AUG_SHIFT\n        ty = (torch.rand(n, device=dev, generator=generator) - 0.5) * 2 * AUG_SHIFT\n        cos, sin = (torch.cos(rot) / sc, torch.sin(rot) / sc)\n        theta = torch.zeros(n, 2, 3, device=dev, dtype=torch.float32)\n        theta[:, 0, 0], theta[:, 0, 1], theta[:, 0, 2] = (cos, -sin, tx)\n        theta[:, 1, 0], theta[:, 1, 1], theta[:, 1, 2] = (sin, cos, ty)\n        grid = F.affine_grid(theta, x.shape, align_corners=False)\n        x = F.grid_sample(x, grid, mode='bilinear', padding_mode='border', align_corners=False)\n        scale = 1.0 + (torch.rand(n, 1, 1, 1, device=dev, generator=generator) - 0.5) * 2 * AUG_INTENSITY\n        x = (x * scale).clamp(0, 255)\n        return x.reshape(*lead, *x.shape[-3:]).to(imgs.dtype)\n\n    def write_submission(pred, studies, test_df, path):\n        sub = pd.DataFrame(pd.DataFrame(pred).rank(pct=True).values, columns=TARGETS)\n        sub.insert(0, 'StudyInstanceUID', studies)\n        sub = test_df[['StudyInstanceUID']].merge(sub, on='StudyInstanceUID', how='left')\n        sub[TARGETS] = sub[TARGETS].fillna(0.5)\n        sub.to_csv(path, index=False)\n        return sub\n\n    def find_dinov2(variant='small'):\n        if not (DINO / 'config.json').is_file():\n            raise FileNotFoundError(DINO)\n        return DINO\n\n    def legacy_group_members():\n        return {}\n\n    def run_dinov2():\n        path = ASSET / 'rsna-knee-weights'\n        infer_from_package(path, DEVS[0])\n        public = Path('/kaggle/working/submission_public_0899.csv')\n        if not public.is_file():\n            raise RuntimeError('public DINOv2 frontier was not produced')\n        public.replace('/kaggle/working/submission.csv')\n        for name in ('submission_legacy_fold_blend.csv', 'legacy_fold_diagnostics.csv'):\n            candidate = Path('/kaggle/working') / name\n            if candidate.is_file():\n                candidate.unlink()\n    run_dinov2()"},{"cell_type":"markdown","metadata":{},"source":"### Stage 2 - DINOv3 folds\n\nRuns the five fold checkpoints on their own slot view and blends their fold-rank\nmean into the file stage 1 wrote, at `DINOV3_RANK_WEIGHT`. The cell snapshots and\nrestores the globals it borrows, so stage 3 still sees stage 1's helpers."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"if SMOKE:\n    print('SMOKE=True: stage 2 (DINOv3 folds) skipped.')\nelse:\n    _A5_SAVED = dict(globals())\n    import gc, os, time, warnings\n    from concurrent.futures import ProcessPoolExecutor, as_completed\n    from pathlib import Path\n    import cv2\n    import numpy as np\n    import pandas as pd\n    import pydicom\n    import timm\n    import torch\n    import torch.nn as nn\n    import torch.nn.functional as F\n    warnings.filterwarnings('ignore')\n    cv2.setNumThreads(1)\n    CROP_MM = 130.0\n    SIZE = 336\n    SLICE_BAND = (0.12, 0.88)\n    N_SLICE = 16\n    INTENSITY = 'slice'\n    SLOTS = [('Sagittal', 1), ('Sagittal', 0), ('Coronal', 1), ('Coronal', 0), ('Axial', 1), ('Axial', 0)]\n    N_SLOT = len(SLOTS)\n    LABELS = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture']\n    # COMP is resolved in the CONFIG/preflight cell above.\n    CKPT = ASSET / 'knee-mri-fold-weights'\n    DEV = 'cuda' if torch.cuda.is_available() else 'cpu'\n    assert DEV == 'cuda', \"GPU requested but not visible; refusing silent CPU fallback\"\n    print(f'competition : {COMP}')\n    print(f'checkpoints : {CKPT}')\n    print(f'device      : {DEV}')\n    for i in range(torch.cuda.device_count() if DEV == 'cuda' else 0):\n        cc = torch.cuda.get_device_capability(i)\n        print(f'  gpu{i}       : {torch.cuda.get_device_name(i)} sm_{cc[0]}{cc[1]}, {torch.cuda.get_device_properties(i).total_memory / 2 ** 30:.0f} GiB, native bf16={cc >= (8, 0)}')\n    SERIES_ROOT = COMP / 'test_series'\n    if not SERIES_ROOT.exists():\n        SERIES_ROOT = COMP / 'train_series'\n    print('series root:', SERIES_ROOT)\n\n    def ordered_files(sdir, cap=64):\n        keyed = []\n        for f in sdir.glob('*.dcm'):\n            try:\n                ds = pydicom.dcmread(str(f), stop_before_pixels=True)\n                keyed.append((int(ds.InstanceNumber), str(f)))\n            except Exception:\n                continue\n            if len(keyed) >= cap * 4:\n                break\n        return [f for _, f in sorted(keyed)]\n\n    def series_side(path):\n        try:\n            return float(pydicom.dcmread(path, stop_before_pixels=True).ImagePositionPatient[0])\n        except Exception:\n            return 0.0\n\n    def read_crop(path):\n        try:\n            ds = pydicom.dcmread(path)\n            arr = ds.pixel_array.astype(np.float32)\n        except Exception:\n            return None\n        try:\n            ps = float(ds.PixelSpacing[0])\n        except Exception:\n            ps = CROP_MM / max(arr.shape)\n        half = int(round(CROP_MM / ps / 2))\n        cy, cx = (arr.shape[0] // 2, arr.shape[1] // 2)\n        y0, y1 = (max(0, cy - half), min(arr.shape[0], cy + half))\n        x0, x1 = (max(0, cx - half), min(arr.shape[1], cx + half))\n        crop = arr[y0:y1, x0:x1]\n        return None if crop.size == 0 else crop\n\n    def window(crop, lo, hi, flip):\n        c = np.clip((crop - lo) / max(hi - lo, 1e-06), 0, 1)\n        img = cv2.resize(c, (SIZE, SIZE), interpolation=cv2.INTER_AREA)\n        return img[:, ::-1].copy() if flip else img\n\n    def render(path, flip):\n        crop = read_crop(path)\n        if crop is None:\n            return None\n        lo, hi = np.percentile(crop[::4, ::4], [1, 99])\n        return window(crop, lo, hi, flip)\n\n    def build_study(args):\n        idx, study, recs = args\n        out = np.zeros((N_SLOT, N_SLICE, SIZE, SIZE), np.uint8)\n        mask = np.zeros(N_SLOT, np.uint8)\n        rows = pd.DataFrame(recs)\n        if len(rows):\n            for s_i, (plane, fs) in enumerate(SLOTS):\n                sub = rows[(rows.Anatomical_Plane == plane) & (rows.Fat_Suppression == fs)]\n                if sub.empty:\n                    continue\n                files = ordered_files(SERIES_ROOT / study / sub.iloc[0].SeriesInstanceUID)\n                if not files:\n                    continue\n                flip = plane != 'Sagittal' and series_side(files[0]) < 0\n                lo, hi = SLICE_BAND\n                i0 = int(round(lo * (len(files) - 1)))\n                i1 = int(round(hi * (len(files) - 1)))\n                avail = list(range(i0, i1 + 1))\n                if len(avail) >= N_SLICE:\n                    picks = [avail[int(round(t))] for t in np.linspace(0, len(avail) - 1, N_SLICE)]\n                    off = 0\n                else:\n                    picks, off = (avail, (N_SLICE - len(avail)) // 2)\n                if INTENSITY == 'series':\n                    crops = [read_crop(files[p]) for p in picks]\n                    got = [x for x in crops if x is not None]\n                    if got:\n                        samp = np.concatenate([x[::4, ::4].ravel() for x in got])\n                        lo_, hi_ = np.percentile(samp, [1, 99])\n                        for c, x in enumerate(crops):\n                            if x is None:\n                                x = read_crop(files[min(len(files) - 1, picks[c] + 1)])\n                            if x is not None:\n                                out[s_i, off + c] = (window(x, lo_, hi_, flip) * 255).astype(np.uint8)\n                else:\n                    for c, p in enumerate(picks):\n                        img = render(files[p], flip)\n                        if img is None:\n                            img = render(files[min(len(files) - 1, p + 1)], flip)\n                        if img is not None:\n                            out[s_i, off + c] = (img * 255).astype(np.uint8)\n                mask[s_i] = len(picks)\n        return (idx, out, mask)\n    sub_df = pd.read_csv(COMP / 'sample_submission.csv')\n    ser_csv = pd.read_csv(COMP / 'test_series.csv')\n    if not (COMP / 'test_series').exists():\n        ser_csv = pd.read_csv(COMP / 'train_series.csv')\n    ser_csv = ser_csv.loc[:, ~ser_csv.columns.duplicated()]\n    studies = sub_df.StudyInstanceUID.tolist()\n    by = {s: g.to_dict('records') for s, g in ser_csv[ser_csv.StudyInstanceUID.isin(set(studies))].groupby('StudyInstanceUID')}\n    print(f'{len(studies):,} test studies, {len(by):,} with series metadata')\n    N_SLOT_TYPES, MASK_IDX = (6, 0)\n\n    def segment_softmax(scores, sidx, B):\n        T, K = scores.shape\n        idx = sidx.unsqueeze(1).expand(-1, K)\n        m = torch.full((B, K), float('-inf'), device=scores.device, dtype=scores.dtype)\n        m = m.scatter_reduce(0, idx, scores, reduce='amax', include_self=True)\n        e = (scores - m[sidx]).exp()\n        s = torch.zeros(B, K, device=scores.device, dtype=scores.dtype).index_add_(0, sidx, e)\n        return e / s[sidx].clamp(min=1e-06)\n\n    class MeanMaxPool(nn.Module):\n\n        def forward(self, f, sidx, B, slot=None, return_attn=False):\n            D = f.shape[1]\n            cnt = torch.zeros(B, device=f.device, dtype=f.dtype).index_add_(0, sidx, torch.ones(f.shape[0], device=f.device, dtype=f.dtype))\n            mean = torch.zeros(B, D, device=f.device, dtype=f.dtype).index_add_(0, sidx, f)\n            mean = mean / cnt.clamp(min=1).unsqueeze(1)\n            mx = torch.full((B, D), -10000.0, device=f.device, dtype=f.dtype)\n            mx = mx.scatter_reduce(0, sidx.unsqueeze(1).expand(-1, D), f, reduce='amax', include_self=True)\n            return (torch.cat([mean, mx], 1), None)\n\n    class LabelAttentionPool(nn.Module):\n\n        def __init__(self, d, n_labels=12, n_heads=4, slot_bias=True):\n            super().__init__()\n            self.d, self.k, self.h = (d, n_labels, n_heads)\n            self.q = nn.Parameter(torch.randn(n_labels, d) * 0.02)\n            self.key, self.val = (nn.Linear(d, d), nn.Linear(d, d))\n            self.slot_bias = nn.Parameter(torch.zeros(n_labels, N_SLOT_TYPES + 1)) if slot_bias else None\n\n        def forward(self, f, sidx, B, slot=None, return_attn=False):\n            scores = self.key(f) @ self.q.t() / self.d ** 0.5\n            if self.slot_bias is not None and slot is not None:\n                scores = scores + self.slot_bias.t()[slot]\n            a = segment_softmax(scores, sidx, B)\n            out = torch.zeros(B, self.k, self.d, device=f.device, dtype=f.dtype)\n            out = out.index_add_(0, sidx, a.unsqueeze(-1) * self.val(f).unsqueeze(1))\n            return (out, a)\n\n    class TokenXAttnPool(nn.Module):\n\n        def __init__(self, d, n_labels=12, n_heads=6, dropout=0.2):\n            super().__init__()\n            self.d, self.k = (d, n_labels)\n            self.q = nn.Parameter(torch.randn(n_labels, d) * 0.02)\n            self.slot_emb = nn.Embedding(N_SLOT_TYPES + 1, d, padding_idx=0)\n            self.kv_norm = nn.LayerNorm(d)\n            self.attn = nn.MultiheadAttention(d, n_heads, dropout=dropout, batch_first=True)\n\n        def forward(self, tok, sidx, B, slot=None, return_attn=False):\n            T, N, D = tok.shape\n            cnt = torch.bincount(sidx, minlength=B)\n            S = int(cnt.max().item())\n            starts = torch.cumsum(cnt, 0) - cnt\n            pos = torch.arange(T, device=tok.device) - starts[sidx]\n            kv = tok + self.slot_emb(slot).unsqueeze(1)\n            pad = tok.new_zeros(B, S, N, D)\n            pad[sidx, pos] = kv\n            keep = torch.zeros(B, S, dtype=torch.bool, device=tok.device)\n            keep[sidx, pos] = True\n            kpm = ~keep.repeat_interleave(N, dim=1)\n            pad = self.kv_norm(pad.reshape(B, S * N, D))\n            q = self.q.unsqueeze(0).expand(B, -1, -1)\n            att, w = self.attn(q, pad, pad, key_padding_mask=kpm, need_weights=return_attn, average_attn_weights=True)\n            cls = tok[:, 0]\n            mean = torch.zeros(B, D, device=tok.device, dtype=tok.dtype).index_add_(0, sidx, cls) / cnt.clamp(min=1).unsqueeze(1)\n            mx = torch.full((B, D), -10000.0, device=tok.device, dtype=tok.dtype)\n            mx = mx.scatter_reduce(0, sidx.unsqueeze(1).expand(-1, D), cls, reduce='amax', include_self=True)\n            base = torch.cat([mean, mx], 1).unsqueeze(1).expand(-1, self.k, -1)\n            return (torch.cat([att, base], -1), w)\n\n    class ViTSlotToken(nn.Module):\n\n        def __init__(self, vit, n_cat, dim=None):\n            super().__init__()\n            self.vit = vit\n            d = dim or vit.embed_dim\n            self.tok = nn.Embedding(n_cat + 1, d, padding_idx=MASK_IDX)\n            self.num_features = vit.num_features\n            self._orig_prefix = getattr(vit, 'num_prefix_tokens', 1)\n            vit.num_prefix_tokens = self._orig_prefix + 1\n            for blk in vit.blocks:\n                a = getattr(blk, 'attn', None)\n                if a is not None and hasattr(a, 'num_prefix_tokens'):\n                    a.num_prefix_tokens = a.num_prefix_tokens + 1\n\n        @staticmethod\n        def _maybe(mod, x):\n            return x if mod is None else mod(x)\n\n        def forward_features(self, x, cat):\n            v = self.vit\n            x = v.patch_embed(x)\n            pos = v._pos_embed(x)\n            rope = None\n            if isinstance(pos, tuple):\n                x, rope = pos\n            else:\n                x = pos\n            x = self._maybe(getattr(v, 'patch_drop', None), x)\n            x = self._maybe(getattr(v, 'norm_pre', None), x)\n            npt = self._orig_prefix\n            tok = self.tok(cat).unsqueeze(1)\n            x = torch.cat([x[:, :npt], tok, x[:, npt:]], dim=1)\n            if rope is not None:\n                if getattr(v, 'rope_mixed', False):\n                    for i, blk in enumerate(v.blocks):\n                        x = blk(x, rope=rope[i])\n                else:\n                    for blk in v.blocks:\n                        x = blk(x, rope=rope)\n            else:\n                x = v.blocks(x)\n            return v.norm(x)\n\n        def forward_head(self, x, pre_logits=True):\n            return self.vit.forward_head(x, pre_logits=pre_logits)\n    IMAGENET_MEAN = (0.485, 0.456, 0.406)\n    IMAGENET_STD = (0.229, 0.224, 0.225)\n\n    class _GatedDepthBlock(nn.Module):\n\n        def __init__(self, n_slice, dropout=0.0, ls_init=0.1):\n            super().__init__()\n            self.norm = nn.GroupNorm(1, n_slice)\n            self.v = nn.Conv2d(n_slice, n_slice, 1)\n            self.g = nn.Conv2d(n_slice, n_slice, 1)\n            self.out = nn.Conv2d(n_slice, n_slice, 1)\n            self.gamma = nn.Parameter(torch.full((n_slice, 1, 1), ls_init))\n            self.drop = nn.Dropout2d(dropout) if dropout else nn.Identity()\n\n        def forward(self, x):\n            z = self.norm(x)\n            return x + self.gamma * self.drop(self.out(self.v(z) * F.silu(self.g(z))))\n\n    class DepthCompress(nn.Module):\n\n        def __init__(self, n_slice=16, out_ch=3, depth=1, dropout=0.0, ls_init=0.1, imagenet=True, proj_noise=0.25):\n            super().__init__()\n            self.imagenet = imagenet\n            self.blocks = nn.ModuleList([_GatedDepthBlock(n_slice, dropout, ls_init) for _ in range(depth)])\n            self.proj = nn.Conv2d(n_slice, out_ch, 1, bias=True)\n            if imagenet:\n                self.register_buffer('mu', torch.tensor(IMAGENET_MEAN).view(1, -1, 1, 1))\n                self.register_buffer('sd', torch.tensor(IMAGENET_STD).view(1, -1, 1, 1))\n\n        def forward(self, x):\n            keep = (x.amax(dim=1, keepdim=True) > 0).to(x.dtype)\n            z = x\n            for b in self.blocks:\n                z = b(z)\n            z = self.proj(z)\n            if self.imagenet:\n                z = (z - self.mu.to(z.dtype)) / self.sd.to(z.dtype)\n            return z * keep\n    N_PLANE, N_CONTRAST = (3, 2)\n    _PLANE_OF = lambda s: torch.clamp(s - 1, 0, 5) // 2\n    _CONTRAST_OF = lambda s: torch.clamp(s - 1, 0, 5) % 2\n\n    class SlotDepthMixer(nn.Module):\n\n        def __init__(self, n_slice=16, ksize=5, alpha_max=0.25):\n            super().__init__()\n            self.n_slice, self.ksize, self.r = (n_slice, ksize, ksize // 2)\n            self.alpha_max = alpha_max\n            b = torch.tensor([1.0, 4.0, 6.0, 4.0, 1.0])\n            self.register_buffer('base', b.log()[self.r:])\n            n_u = self.r + 1\n            self.shared = nn.Parameter(torch.zeros(n_u))\n            self.plane_k = nn.Parameter(torch.zeros(N_PLANE, n_u))\n            self.contrast_k = nn.Parameter(torch.zeros(N_CONTRAST, n_u))\n            self.g0 = nn.Parameter(torch.zeros(()))\n            self.gate_p = nn.Parameter(torch.zeros(N_PLANE))\n            self.gate_c = nn.Parameter(torch.zeros(N_CONTRAST))\n            idx = torch.arange(n_slice)\n            self.register_buffer('off', idx[None, :] - idx[:, None])\n\n        def kernel(self, slot):\n            p, c = (_PLANE_OF(slot), _CONTRAST_OF(slot))\n            half = self.base + self.shared + self.plane_k[p] + self.contrast_k[c]\n            full = torch.cat([half.flip(-1)[..., :self.r], half], dim=-1)\n            return F.softmax(full, dim=-1)\n\n        def alpha(self, slot):\n            p, c = (_PLANE_OF(slot), _CONTRAST_OF(slot))\n            return self.alpha_max * torch.tanh(self.g0 + self.gate_p[p] + self.gate_c[c])\n\n        def forward(self, x, slot, vmask):\n            T, S, H, W = x.shape\n            if vmask is None:\n                raise ValueError('stem=mixer requires the padding mask')\n            k = self.kernel(slot)\n            v = vmask.to(k.dtype)\n            d = self.off + self.r\n            inb = (d >= 0) & (d < self.ksize)\n            kk = k[:, d.clamp(0, self.ksize - 1)] * inb\n            M = kk * v[:, None, :]\n            den = M.sum(-1, keepdim=True)\n            eye = torch.eye(S, device=x.device, dtype=M.dtype).expand(T, S, S)\n            ok = (den > 1e-06) & v[:, :, None].bool()\n            M = torch.where(ok, M / den.clamp(min=1e-06), eye)\n            a = self.alpha(slot)[:, None, None]\n            Aop = ((1.0 - a) * eye + a * M).to(x.dtype)\n            if x.is_contiguous(memory_format=torch.channels_last) and (not x.is_contiguous()):\n                y = torch.bmm(x.permute(0, 2, 3, 1).reshape(T, H * W, S), Aop.transpose(1, 2))\n                return y.reshape(T, H, W, S).permute(0, 3, 1, 2)\n            return torch.bmm(Aop, x.reshape(T, S, H * W)).reshape(T, S, H, W)\n\n    def _seg_mean_max(v, sidx, B):\n        D = v.shape[1]\n        cnt = torch.zeros(B, device=v.device, dtype=v.dtype).index_add_(0, sidx, torch.ones(v.shape[0], device=v.device, dtype=v.dtype))\n        mean = torch.zeros(B, D, device=v.device, dtype=v.dtype).index_add_(0, sidx, v)\n        mean = mean / cnt.clamp(min=1).unsqueeze(1)\n        mx = torch.full((B, D), -10000.0, device=v.device, dtype=v.dtype)\n        mx = mx.scatter_reduce(0, sidx.unsqueeze(1).expand(-1, D), v, reduce='amax', include_self=True)\n        return torch.cat([mean, mx], 1)\n\n    def _pad_kv(x, sidx, B, norm):\n        T, P, D = x.shape\n        cnt = torch.bincount(sidx, minlength=B)\n        S = int(cnt.max().item())\n        starts = torch.cumsum(cnt, 0) - cnt\n        pos = torch.arange(T, device=x.device) - starts[sidx]\n        pad = x.new_zeros(B, S, P, D)\n        pad[sidx, pos] = x\n        keep = torch.zeros(B, S, dtype=torch.bool, device=x.device)\n        keep[sidx, pos] = True\n        return (norm(pad.reshape(B, S * P, D)), ~keep.repeat_interleave(P, dim=1))\n\n    class _GatedDelta(nn.Module):\n\n        def __init__(self, d, n_labels, n_heads, dropout):\n            super().__init__()\n            self.q = nn.Parameter(torch.randn(n_labels, d) * 0.02)\n            self.kv_norm = nn.LayerNorm(d)\n            self.attn = nn.MultiheadAttention(d, n_heads, dropout=dropout, batch_first=True)\n            self.d_norm = nn.LayerNorm(d)\n            self.dw = nn.Parameter(torch.randn(n_labels, d) * (1.0 / d ** 0.5))\n            self.db = nn.Parameter(torch.zeros(n_labels))\n            self.gate = nn.Parameter(torch.zeros(n_labels))\n\n        def delta(self, pat, sidx, B, return_attn):\n            kv, kpm = _pad_kv(pat, sidx, B, self.kv_norm)\n            q = self.q.unsqueeze(0).expand(B, -1, -1)\n            att, w = self.attn(q, kv, kv, key_padding_mask=kpm, need_weights=return_attn, average_attn_weights=True)\n            return ((self.d_norm(att) * self.dw).sum(-1) + self.db, w)\n\n    class TokenResidualPool(_GatedDelta):\n\n        def __init__(self, d, n_labels=12, n_heads=6, pe=64, dropout=0.2):\n            super().__init__(d, n_labels, n_heads, dropout)\n            self.base = nn.Sequential(nn.LayerNorm(2 * d + pe), nn.Dropout(dropout), nn.Linear(2 * d + pe, n_labels))\n\n        def forward(self, tok, slot, sidx, B, pres, return_attn=False):\n            base = self.base(torch.cat([_seg_mean_max(tok[:, 1:].mean(1), sidx, B), pres], 1))\n            d_, w = self.delta(tok[:, 1:], sidx, B, return_attn)\n            return (base + self.gate * d_, w)\n\n    class CodexResidualPool(_GatedDelta):\n\n        def __init__(self, d, n_labels=12, n_heads=6, pe=64, dropout=0.2):\n            super().__init__(d, n_labels, n_heads, dropout)\n            self.base = nn.Sequential(nn.LayerNorm(2 * d + pe), nn.Dropout(dropout), nn.Linear(2 * d + pe, n_labels))\n\n        def forward(self, tok, slot, sidx, B, pres, return_attn=False):\n            base = self.base(torch.cat([_seg_mean_max(tok[:, 0], sidx, B), pres], 1))\n            d_, w = self.delta(tok[:, 1:], sidx, B, return_attn)\n            return (base + self.gate * d_, w)\n\n    class ClsAddPool(nn.Module):\n\n        def __init__(self, d, n_labels=12, pe=64, dropout=0.2):\n            super().__init__()\n            self.net = nn.Sequential(nn.LayerNorm(4 * d + pe), nn.Dropout(dropout), nn.Linear(4 * d + pe, n_labels))\n\n        def forward(self, tok, slot, sidx, B, pres, return_attn=False):\n            return (self.net(torch.cat([_seg_mean_max(tok[:, 1:].mean(1), sidx, B), _seg_mean_max(tok[:, 0], sidx, B), pres], 1)), None)\n\n    class Readout(nn.Module):\n\n        def __init__(self, pool, d, n_labels=12, pe=64):\n            super().__init__()\n            self.pool_kind, self.k = (pool, n_labels)\n            self.pres_emb = nn.Embedding(N_SLOT_TYPES + 1, pe, padding_idx=0)\n            if pool in ('xres', 'clsadd', 'xcodex'):\n                self.pool = {'xres': TokenResidualPool, 'clsadd': ClsAddPool, 'xcodex': CodexResidualPool}[pool](d, n_labels, pe=pe)\n            elif pool in ('attn', 'xattn'):\n                if pool == 'xattn':\n                    self.pool = TokenXAttnPool(d, n_labels)\n                    wd = 3 * d + pe\n                else:\n                    self.pool = LabelAttentionPool(d, n_labels)\n                    wd = d + pe\n                self.norm = nn.LayerNorm(wd)\n                self.w = nn.Parameter(torch.randn(n_labels, wd) * (1.0 / wd ** 0.5))\n                self.b = nn.Parameter(torch.zeros(n_labels))\n            else:\n                self.pool = MeanMaxPool()\n                self.net = nn.Sequential(nn.LayerNorm(2 * d + pe), nn.Dropout(0.2), nn.Linear(2 * d + pe, n_labels))\n            self.drop = nn.Dropout(0.2)\n\n        def forward(self, f, slot, sidx, B, return_attn=False):\n            pe = self.pres_emb(slot)\n            pres = torch.zeros(B, pe.shape[1], device=f.device, dtype=f.dtype).index_add_(0, sidx, pe)\n            if self.pool_kind in ('xres', 'clsadd', 'xcodex'):\n                return self.pool(f, slot, sidx, B, pres)[0]\n            pooled, attn = self.pool(f, sidx, B, slot=slot, return_attn=return_attn)\n            if self.pool_kind in ('attn', 'xattn'):\n                x = torch.cat([pooled, pres.unsqueeze(1).expand(-1, self.k, -1)], -1)\n                x = self.drop(self.norm(x))\n                return (x * self.w).sum(-1) + self.b\n            return self.net(torch.cat([pooled, pres], 1))\n\n    class Net(nn.Module):\n\n        def __init__(self, enc, cond, n_meta=0, pool='mean_max', stem='native', n_slice=16):\n            super().__init__()\n            self.enc, self.cond = (enc, cond)\n            self.compress = DepthCompress(n_slice, 3) if stem == 'compress' else None\n            self.mixer = SlotDepthMixer(n_slice) if stem == 'mixer' else None\n            self.tokens = pool in ('xattn', 'xres', 'clsadd', 'xcodex')\n            D = enc.num_features\n            self.meta_mlp = nn.Sequential(nn.LayerNorm(n_meta), nn.Linear(n_meta, 128), nn.GELU(), nn.Linear(128, D)) if n_meta > 0 else None\n            self.readout = Readout(pool, D)\n            if cond == 'post':\n                self.slot_emb = nn.Embedding(N_SLOT_TYPES + 1, D, padding_idx=MASK_IDX)\n\n        def forward(self, im, slot, smeta, sidx, B, vm=None):\n            if self.mixer is not None:\n                im = self.mixer(im, slot, vm)\n            if self.compress is not None:\n                im = self.compress(im)\n            f = self.enc.forward_features(im, slot) if self.cond == 'token' else self.enc.forward_features(im)\n            if self.tokens:\n                inner = getattr(self.enc, 'vit', self.enc)\n                orig = getattr(self.enc, '_orig_prefix', getattr(inner, 'num_prefix_tokens', 1))\n                f = torch.cat([f[:, :1], f[:, orig:]], 1)\n            else:\n                f = self.enc.forward_head(f, pre_logits=True)\n                if f.dim() > 2:\n                    f = f.flatten(1)\n            ex = (lambda v: v.unsqueeze(1)) if self.tokens else lambda v: v\n            if self.cond == 'post':\n                f = f + ex(self.slot_emb(slot))\n            if self.meta_mlp is not None and smeta.shape[1] > 0:\n                mt = self.meta_mlp(smeta)\n                f = torch.cat([f, mt.unsqueeze(1)], 1) if self.tokens else f + mt\n            return self.readout(f, slot, sidx, B)\n    models = []\n    for ckpt_path in sorted(CKPT.glob('*_f*.pt')):\n        z = torch.load(ckpt_path, map_location='cpu', weights_only=False)\n        cfg = z['cfg']\n        _stem = cfg.get('stem', 'native')\n        _in = 3 if _stem == 'compress' else cfg.get('n_slice', 16)\n        enc = timm.create_model(cfg['backbone'], pretrained=False, num_classes=0, in_chans=_in, **{'img_size': cfg['img']} if 'vit_' in cfg['backbone'] else {})\n        if cfg['cond'] == 'token':\n            enc = ViTSlotToken(enc, N_SLOT_TYPES)\n        m = Net(enc, cfg['cond'], cfg.get('n_meta', 0), cfg['pool'], stem=_stem, n_slice=cfg.get('n_slice', 16))\n        missing, unexpected = m.load_state_dict(z['state_dict'], strict=False)\n        assert not missing, f'missing {missing[:5]}'\n        assert not unexpected, f'unexpected {unexpected[:5]}'\n        models.append(m.eval())\n        print(f\"loaded {ckpt_path.name}  fold {z['fold']}  {cfg['backbone']} pool={cfg['pool']} meta={cfg['meta']}\")\n    CFG = cfg\n    assert CFG.get('n_meta', 0) == 0, f\"checkpoint expects {CFG['n_meta']} metadata features -- build slot_meta for the TEST studies and pass it to predict() before submitting\"\n    print(f\"\\n{len(models)} fold models ready | input norm: {CFG.get('norm', 'none')}\")\n    AMP_PREF = 'bf16'\n\n    def amp_for(dev):\n        if not str(dev).startswith('cuda'):\n            return (torch.float32, False)\n        cc = torch.cuda.get_device_capability(dev)\n        if AMP_PREF == 'bf16':\n            return (torch.bfloat16, True)\n        if AMP_PREF == 'fp16':\n            return (torch.float16, True)\n        if AMP_PREF == 'fp32':\n            return (torch.float32, False)\n        return (torch.bfloat16 if cc >= (8, 0) else torch.float16, True)\n    AMP_DT, AMP_ON = amp_for(DEV)\n    WORKERS = max(1, min(4, os.cpu_count() or 4))\n    CHUNK = 48\n    MICRO = 8\n    models = [m.to(DEV).eval() for m in models]\n    print(f\"device {DEV} | amp {str(AMP_DT).split('.')[-1]} (on={AMP_ON}) | workers {WORKERS} | chunk {CHUNK} | micro {MICRO}\")\n\n    def _norm_(im):\n        k = CFG.get('norm', 'none')\n        if k == 'zscore':\n            m = (im > 0).float()\n            n = m.sum(dim=(1, 2, 3), keepdim=True).clamp(min=1.0)\n            mu = (im * m).sum(dim=(1, 2, 3), keepdim=True) / n\n            var = (((im - mu) * m) ** 2).sum(dim=(1, 2, 3), keepdim=True) / n\n            return (im - mu) / (var.sqrt() + 1e-06) * m\n        if k == 'imagenet':\n            m = (im > 0).float()\n            return (im - 0.485) / 0.229 * m\n        return im\n\n    @torch.no_grad()\n    def _micro(images, masks):\n        dev = DEV\n        ims, slots, sidx, vms = ([], [], [], [])\n        for b in range(len(masks)):\n            present = np.nonzero(masks[b] > 0)[0]\n            if len(present) == 0:\n                continue\n            blk = images[b][present]\n            ims.append(torch.from_numpy(blk))\n            vms.append(torch.from_numpy(blk.reshape(blk.shape[0], blk.shape[1], -1).max(2) > 0))\n            slots.append(torch.from_numpy(present + 1).long())\n            sidx.append(torch.full((len(present),), b, dtype=torch.long))\n        out = np.full((len(models), len(masks), len(LABELS)), np.nan, np.float32)\n        if not ims:\n            return out\n        im = _norm_(torch.cat(ims).to(dev, non_blocking=True).float().div_(255.0))\n        sl = torch.cat(slots).to(dev)\n        si = torch.cat(sidx).to(dev)\n        vm = torch.cat(vms).to(dev)\n        sm = torch.zeros(len(sl), CFG.get('n_meta', 0), device=dev)\n        per = torch.zeros(len(models), len(masks), len(LABELS), device=dev, dtype=torch.float32)\n        with torch.autocast('cuda' if str(dev).startswith('cuda') else 'cpu', dtype=AMP_DT, enabled=AMP_ON):\n            for fold_index, model in enumerate(models):\n                per[fold_index] = torch.sigmoid(model(im, sl, sm, si, len(masks), vm=vm).float())\n        got = per.cpu().numpy()\n        keep = np.array([(masks[b] > 0).any() for b in range(len(masks))])\n        out[:, keep] = got[:, keep]\n        return out\n\n    def predict(images, masks):\n        out = np.full((len(models), len(masks), len(LABELS)), np.nan, np.float32)\n        for a in range(0, len(masks), MICRO):\n            b = min(a + MICRO, len(masks))\n            out[:, a:b] = _micro(images[a:b], masks[a:b])\n        return out\n    preds = np.full((len(models), len(studies), len(LABELS)), np.nan, np.float32)\n    t0, done = (time.time(), 0)\n    with ProcessPoolExecutor(max_workers=WORKERS) as ex:\n        for c0 in range(0, len(studies), CHUNK):\n            block = studies[c0:c0 + CHUNK]\n            imgs = np.zeros((len(block), N_SLOT, N_SLICE, SIZE, SIZE), np.uint8)\n            msks = np.zeros((len(block), N_SLOT), np.uint8)\n            futs = [ex.submit(build_study, (i, s, by.get(s, []))) for i, s in enumerate(block)]\n            for f in as_completed(futs):\n                try:\n                    i, a, k = f.result()\n                    imgs[i], msks[i] = (a, k)\n                except Exception as e:\n                    print(f'  study failed: {type(e).__name__}: {e}')\n            preds[:, c0:c0 + len(block)] = predict(imgs, msks)\n            done += len(block)\n            el = time.time() - t0\n            print(f'  {done:,}/{len(studies):,}  {el / 60:.1f}m  eta {el / done * (len(studies) - done) / 60:.1f}m', flush=True)\n            del imgs, msks\n            gc.collect()\n    print(f'\\ninference done in {(time.time() - t0) / 60:.1f} min')\n    A5_W = DINOV3_RANK_WEIGHT\n    A5_LABELS = list(LABELS)\n    _a5_ok = np.isfinite(preds).all(axis=(0, 2))\n    _a5_rank_mean = np.zeros((len(studies), len(LABELS)), np.float64)\n    for fold_index in range(preds.shape[0]):\n        fold = preds[fold_index][_a5_ok]\n        ordinal = fold.argsort(0).argsort(0).astype(np.float64)\n        _a5_rank_mean[_a5_ok] += ordinal / max(len(fold) - 1, 1)\n    _a5_rank_mean /= preds.shape[0]\n    _a5_rank_mean[~_a5_ok] = np.nan\n    A5_PREDS = dict(zip(sub_df['StudyInstanceUID'].astype(str), _a5_rank_mean.astype(np.float32)))\n    for _a5k, _a5v in _A5_SAVED.items():\n        globals()[_a5k] = _a5v\n    del _A5_SAVED, _a5k, _a5v\n    _a5_sub = pd.read_csv('/kaggle/working/submission.csv', dtype={'StudyInstanceUID': str})\n    assert _a5_sub.columns.tolist()[1:] == A5_LABELS, 'submission schema drift'\n    if A5_W > 0:\n        _a5_ours = np.stack([A5_PREDS[_u] for _u in _a5_sub['StudyInstanceUID'].astype(str)])\n        _a5_base_rank = _a5_sub[A5_LABELS].rank(method='average', pct=True)\n        _a5_ours_rank = pd.DataFrame(_a5_ours, columns=A5_LABELS, index=_a5_sub.index).rank(method='average', pct=True)\n        _a5_sub[A5_LABELS] = (1.0 - A5_W) * _a5_base_rank + A5_W * _a5_ours_rank\n        assert np.isfinite(_a5_sub[A5_LABELS].to_numpy()).all()\n        _a5_sub.to_csv('/kaggle/working/submission.csv', index=False)"},{"cell_type":"markdown","metadata":{},"source":"### Stage 3 - RadImageNet branch\n\nOne shared RadImageNet ResNet-50 encoder, two banks of five attention heads, and\nthe optional second pass. Rewrites `submission.csv` in place with the final\ntwelve-target prediction."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"from __future__ import annotations\n\nif SMOKE:\n    print('SMOKE=True: stage 3 (RadImageNet branch) skipped.')\nelse:\n    import contextlib as _rad_contextlib\n    import gc as _rad_gc\n    import hashlib as _rad_hashlib\n    import json as _rad_json\n    import os as _rad_os\n    import re as _rad_re\n    import time as _rad_time\n    from concurrent.futures import ThreadPoolExecutor as _RadThreadPool\n    from pathlib import Path as _RadPath\n    import numpy as _rad_np\n    import pandas as _rad_pd\n    import pydicom as _rad_pydicom\n    import torch as _rad_torch\n    import torch.nn as _rad_nn\n    import torch.nn.functional as _rad_F\n    from torchvision.models import resnet50 as _rad_resnet50\n    _RAD_LABELS = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture']\n    _RAD_ALPHA = RAD_ALPHA\n    _RAD_EXCLUDE = (\"Baker's\", 'Fracture')\n    _RAD_REFERENCE_HEADS_SHA256 = '0f465649799ecfbccaac1767844639e7ced44e1bc9babde6e4bac7c5d9b89eaa'\n    _RAD_ENCODER_SHA256 = '08629f7e7bd3e29b8ee9522ca3f65ce4d010a7ddf74f0ea3c7e3f3d0bbab0734'\n    _RAD_E13_HEADS_SHA256 = 'ad9f19af73bfdf4e49263c0e45060dc3cb239e1195039b26dc8c0a3a6bcd1a8a'\n    _RAD_E13_MEMBER_WEIGHT = RAD_E13_MEMBER_WEIGHT\n    _RAD_V48_SECOND_ALPHA = PASS2_WEIGHT\n    _RAD_TOKEN_DIM, _RAD_HEAD_DIM = (2048, 512)\n    _RAD_E11_SLOTS = [('SAG_NOFS', 'Sagittal', None, False), ('COR_NOFS', 'Coronal', None, False), ('AX_NOFS', 'Axial', None, False), ('SAG_FS', 'Sagittal', None, True)]\n    _RAD_E11_CROP_MM = 130.0\n    _RAD_E13_SLOTS = [('SAG_FS', 'Sagittal', None, True), ('COR_FS', 'Coronal', None, True), ('AX_FS', 'Axial', None, True), ('SAG_NOFS', 'Sagittal', None, False)]\n    _RAD_E13_CROP_MM = 130.0\n    _RAD_E13_CACHE_SLICES = 8\n    _RAD_E13_IMG = 224\n    SLOTS = [('SAG_FS', 'Sagittal', None, True), ('COR_FS', 'Coronal', None, True), ('AX_FS', 'Axial', None, True)]\n    N_SLOT = len(SLOTS)\n    CACHE_SLICES = 8\n\n    def _rad_sha256(path, chunk=8 << 20):\n        digest = _rad_hashlib.sha256()\n        with open(path, 'rb') as handle:\n            for block in iter(lambda: handle.read(chunk), b''):\n                digest.update(block)\n        return digest.hexdigest()\n\n    def _rad_find_file(name, expected_sha=None, explicit_env=None):\n        files = {_RAD_ENCODER_SHA256: ASSET / 'resnet-50-radimagenet-marwan/ResNet50.pt', _RAD_REFERENCE_HEADS_SHA256: ASSET / 'rsna-knee-e9-radimagenet-heads-v15/v52_radimagenet_heads.pt', _RAD_E13_HEADS_SHA256: ASSET / 'kernel-sources/rsna-knee-e13-train/rsna_rad_e11/v52_e11_heads.pt'}\n        path = files.get(expected_sha)\n        if path is None or not path.is_file():\n            raise FileNotFoundError(name)\n        if _rad_sha256(path) != expected_sha:\n            raise RuntimeError(f'hash mismatch for {path}')\n        return path\n\n    class _RadEncoder(_rad_nn.Module):\n\n        def __init__(self):\n            super().__init__()\n            self.backbone = _rad_nn.Sequential(*list(_rad_resnet50(weights=None).children())[:-2])\n\n        def forward(self, image):\n            return self.backbone(image).mean(dim=(2, 3))\n\n    class _RadHead(_rad_nn.Module):\n\n        def __init__(self):\n            super().__init__()\n            self.project = _rad_nn.Sequential(_rad_nn.LayerNorm(_RAD_TOKEN_DIM), _rad_nn.Linear(_RAD_TOKEN_DIM, _RAD_HEAD_DIM), _rad_nn.GELU())\n            self.plane = _rad_nn.Parameter(_rad_torch.randn(N_SLOT, _RAD_HEAD_DIM) * 0.01)\n            self.position = _rad_nn.Parameter(_rad_torch.randn(CACHE_SLICES, _RAD_HEAD_DIM) * 0.01)\n            self.query = _rad_nn.Parameter(_rad_torch.randn(len(_RAD_LABELS), _RAD_HEAD_DIM) * 0.02)\n            self.attn = _rad_nn.MultiheadAttention(_RAD_HEAD_DIM, 8, dropout=0.1, batch_first=True)\n            self.fuse = _rad_nn.Sequential(_rad_nn.LayerNorm(_RAD_HEAD_DIM * 4), _rad_nn.Linear(_RAD_HEAD_DIM * 4, _RAD_HEAD_DIM), _rad_nn.GELU(), _rad_nn.Dropout(0.15))\n            self.weight = _rad_nn.Parameter(_rad_torch.randn(len(_RAD_LABELS), _RAD_HEAD_DIM) * 0.02)\n            self.bias = _rad_nn.Parameter(_rad_torch.zeros(len(_RAD_LABELS)))\n\n        def forward(self, feature, mask):\n            token = self.project(feature.float())\n            token = token.view(len(token), N_SLOT, CACHE_SLICES, _RAD_HEAD_DIM)\n            token = token + self.plane[None, :, None] + self.position[None, None]\n            token = token.flatten(1, 2)\n            key_padding = mask <= 0\n            all_empty = key_padding.all(1)\n            if all_empty.any():\n                key_padding = key_padding.clone()\n                key_padding[all_empty, 0] = False\n            query = self.query.unsqueeze(0).expand(len(token), -1, -1)\n            attended = query + self.attn(query, token, token, key_padding_mask=key_padding, need_weights=False)[0]\n            denominator = mask.sum(1, keepdim=True).clamp_min(1).unsqueeze(-1)\n            mean = (token * mask.unsqueeze(-1)).sum(1, keepdim=True) / denominator\n            mean = mean.expand(-1, len(_RAD_LABELS), -1)\n            fused = self.fuse(_rad_torch.cat([attended, mean, _rad_torch.abs(attended - mean), attended * mean], dim=-1))\n            return (fused * self.weight.unsqueeze(0)).sum(-1) + self.bias\n\n    def _rad_load_public_heads(device, expected_sha):\n        heads_path = _rad_find_file('v52_radimagenet_heads.pt', expected_sha)\n        payload = _rad_torch.load(heads_path, map_location='cpu', weights_only=True)\n        expected = {'version': 'v52-radimagenet-resnet50-official-1', 'targets': _RAD_LABELS, 'encoder_sha256': _RAD_ENCODER_SHA256, 'encoder_source_commit': '0ce16f7375db4236e646829d1eca61cdb4282133', 'img': 224, 'slices_per_plane': 8, 'feature': 'global_average_pool'}\n        for key, value in expected.items():\n            if payload.get(key) != value:\n                raise RuntimeError(f'public-v15 head contract drift for {key}')\n        folds = payload.get('folds')\n        if not isinstance(folds, list) or len(folds) != 5:\n            raise RuntimeError('public-v15 bundle requires exactly five heads')\n        if sorted((int(record.get('fold', -1)) for record in folds)) != list(range(5)):\n            raise RuntimeError('public-v15 fold identity drift')\n        heads = []\n        for record in folds:\n            head = _RadHead().to(device).eval()\n            head.load_state_dict(record['state_dict'], strict=True)\n            heads.append(head)\n        return (heads, str(heads_path))\n\n    def _rad_load_e13_heads(device):\n        heads_path = _rad_find_file('v52_e11_heads.pt', _RAD_E13_HEADS_SHA256)\n        payload = _rad_torch.load(heads_path, map_location='cpu', weights_only=False)\n        expected = {'version': 'e11-radimagenet-resnet50-diverse-1', 'targets': _RAD_LABELS, 'encoder_sha256': _RAD_ENCODER_SHA256, 'slots': [list(slot) for slot in _RAD_E13_SLOTS], 'crop_mm': _RAD_E13_CROP_MM, 'img': _RAD_E13_IMG, 'slices_per_plane': _RAD_E13_CACHE_SLICES, 'feature': 'global_average_pool'}\n        for key, value in expected.items():\n            if payload.get(key) != value:\n                raise RuntimeError(f'E13 head contract drift for {key}')\n        folds = payload.get('folds')\n        if not isinstance(folds, list) or len(folds) != 5:\n            raise RuntimeError('E13 bundle requires exactly five heads')\n        if sorted((int(record.get('fold', -1)) for record in folds)) != list(range(5)):\n            raise RuntimeError('E13 fold identity drift')\n        heads = []\n        for record in folds:\n            head = _RadHead().to(device).eval()\n            head.load_state_dict(record['state_dict'], strict=True)\n            heads.append(head)\n        return (heads, str(heads_path))\n\n    @_rad_torch.inference_mode()\n    def _rad_encode(encoder, pixels, slot_mask, device):\n        n, slots, slices, height, width = pixels.shape\n        features = _rad_np.zeros((n, slots * slices, _RAD_TOKEN_DIM), _rad_np.float16)\n        token_mask = _rad_np.repeat(slot_mask[:, :, None], slices, axis=2).reshape(n, -1)\n        valid = _rad_np.flatnonzero(token_mask.reshape(-1) > 0)\n        flat = pixels.reshape(-1, height, width)\n        batch = 192 if device.type == 'cuda' and _rad_torch.cuda.device_count() > 1 else 96 if device.type == 'cuda' else 8\n        for start in range(0, len(valid), batch):\n            indices = valid[start:start + batch]\n            image = _rad_torch.from_numpy(flat[indices]).to(device).float().div_(127.5).sub_(1.0)\n            image = image.unsqueeze(1).expand(-1, 3, -1, -1).contiguous()\n            amp = _rad_torch.autocast('cuda') if device.type == 'cuda' else _rad_contextlib.nullcontext()\n            with amp:\n                feature = encoder(image)\n            values = feature.float().cpu().numpy()\n            if not _rad_np.isfinite(values).all():\n                raise RuntimeError('V36 non-finite RadImageNet feature')\n            features.reshape(-1, _RAD_TOKEN_DIM)[indices] = values.astype(_rad_np.float16)\n        return (features, token_mask.astype(_rad_np.float32))\n\n    @_rad_torch.inference_mode()\n    def _rad_predict_head(head, features, masks, device, batch=64):\n        predictions = []\n        for start in range(0, len(features), batch):\n            image = _rad_torch.from_numpy(features[start:start + batch]).to(device)\n            mask = _rad_torch.from_numpy(masks[start:start + batch]).to(device)\n            amp = _rad_torch.autocast('cuda') if device.type == 'cuda' else _rad_contextlib.nullcontext()\n            with amp:\n                predictions.append(_rad_torch.sigmoid(head(image, mask)).float().cpu())\n        return _rad_torch.cat(predictions).numpy()\n\n    def _rad_rank_columns(values):\n        return _rad_pd.DataFrame(_rad_np.asarray(values, dtype=_rad_np.float64)).rank(method='average', pct=True).to_numpy(_rad_np.float64)\n\n    def _rad_validate(frame, expected_ids):\n        if frame.columns.tolist() != ['StudyInstanceUID', *_RAD_LABELS]:\n            raise RuntimeError('V36 submission schema drift')\n        ids = frame['StudyInstanceUID'].astype(str).tolist()\n        if ids != list(map(str, expected_ids)) or len(ids) != len(set(ids)):\n            raise RuntimeError('V36 submission study identity/order drift')\n        values = frame[_RAD_LABELS].to_numpy(_rad_np.float64)\n        if not _rad_np.isfinite(values).all() or values.min() < 0 or values.max() > 1:\n            raise RuntimeError('V36 invalid submission values')\n\n    def _rad_main():\n        work = _RadPath('/kaggle/working')\n        primary = work / 'submission.csv'\n        test = _rad_pd.read_csv(ROOT / 'test.csv', dtype={'StudyInstanceUID': str})\n        expected_ids = test.StudyInstanceUID.astype(str).tolist()\n        baseline = _rad_pd.read_csv(primary, dtype={'StudyInstanceUID': str})\n        _rad_validate(baseline, expected_ids)\n        device = _rad_torch.device('cuda:0')\n        test_series = _rad_pd.read_csv(ROOT / 'test_series.csv', dtype={'StudyInstanceUID': str, 'SeriesInstanceUID': str})\n        plane = dict(zip(test_series.SeriesInstanceUID, test_series.Anatomical_Plane))\n\n        def cache(slots, crop, tag, threshold):\n            globals().update(SLOTS=list(slots), N_SLOT=len(slots), CACHE_SLICES=8, IMG=224, CACHE_IMG=224, CROP_MM=float(crop), RULES=dict(RULES_LEGACY))\n            headers = annotate(walk('test_series'))\n            studies, pixels, masks = build_cache(pick_slots(headers, plane), plane, lat_of(headers, tag + ' '), tag)\n            positions = {str(uid): index for index, uid in enumerate(studies)}\n            missing = [uid for uid in expected_ids if uid not in positions]\n            if missing:\n                raise RuntimeError(f'{len(missing)} studies absent from {tag}')\n            order = _rad_np.asarray([positions[uid] for uid in expected_ids], dtype=_rad_np.int64)\n            pixels, masks = (pixels[order], masks[order])\n            tokens = int(_rad_np.repeat(masks[:, :, None], CACHE_SLICES, axis=2).sum())\n            if tokens < int(threshold * len(test) * N_SLOT * CACHE_SLICES):\n                raise RuntimeError(f'insufficient slices for {tag}: {tokens}')\n            return (pixels, masks)\n        public_slots = [('SAG_FS', 'Sagittal', None, True), ('COR_FS', 'Coronal', None, True), ('AX_FS', 'Axial', None, True)]\n        pixels, masks = cache(public_slots, 10000.0, 'test-e10', 0.85)\n        encoder_path = _rad_find_file('ResNet50.pt', _RAD_ENCODER_SHA256)\n        encoder = _RadEncoder()\n        encoder.load_state_dict(_rad_torch.load(encoder_path, map_location='cpu', weights_only=True), strict=True)\n        encoder.eval().to(device)\n        for parameter in encoder.parameters():\n            parameter.requires_grad_(False)\n        if _rad_torch.cuda.device_count() > 1:\n            encoder = _rad_nn.DataParallel(encoder, device_ids=list(range(_rad_torch.cuda.device_count())))\n        reference_heads, _ = _rad_load_public_heads(device, _RAD_REFERENCE_HEADS_SHA256)\n        features, token_mask = _rad_encode(encoder, pixels, masks, device)\n        reference_predictions = [_rad_predict_head(head, features, token_mask, device) for head in reference_heads]\n        reference_probability = _rad_np.mean(_rad_np.stack(reference_predictions), axis=0)\n        reference_rank = _rad_rank_columns(reference_probability)\n        del reference_predictions, reference_heads\n        del reference_probability, features, token_mask, pixels, masks\n        _rad_gc.collect()\n        _rad_torch.cuda.empty_cache()\n        globals().update(SLOTS=list(_RAD_E13_SLOTS), N_SLOT=len(_RAD_E13_SLOTS), CACHE_SLICES=_RAD_E13_CACHE_SLICES, IMG=_RAD_E13_IMG, CACHE_IMG=_RAD_E13_IMG, CROP_MM=_RAD_E13_CROP_MM, RULES=dict(RULES_LEGACY))\n        e13_heads, _ = _rad_load_e13_heads(device)\n        pixels, masks = cache(_RAD_E13_SLOTS, _RAD_E13_CROP_MM, 'test-e13', 0.85)\n        features, token_mask = _rad_encode(encoder, pixels, masks, device)\n        e13_predictions = [_rad_predict_head(head, features, token_mask, device) for head in e13_heads]\n        e13_probability = _rad_np.mean(_rad_np.stack(e13_predictions), axis=0)\n        e13_rank = _rad_rank_columns(e13_probability)\n        reference_rank = _rad_rank_columns((1.0 - _RAD_E13_MEMBER_WEIGHT) * reference_rank + _RAD_E13_MEMBER_WEIGHT * e13_rank)\n        del e13_predictions, e13_probability, e13_rank\n        del features, token_mask, pixels, masks\n        _rad_gc.collect()\n        _rad_torch.cuda.empty_cache()\n        baseline_rank = _rad_rank_columns(baseline[_RAD_LABELS].to_numpy())\n        e10 = baseline.copy()\n        for index, target in enumerate(_RAD_LABELS):\n            if target not in _RAD_EXCLUDE:\n                e10[target] = (1.0 - _RAD_ALPHA) * baseline_rank[:, index] + _RAD_ALPHA * reference_rank[:, index]\n        _rad_validate(e10, expected_ids)\n        final = e10.copy()\n        if PASS2_ENABLED:\n            pixels, masks = cache(_RAD_E11_SLOTS, _RAD_E11_CROP_MM, 'test-v48-pass2', 0.55)\n            features, token_mask = _rad_encode(encoder, pixels, masks, device)\n            pass2_predictions = [_rad_predict_head(head, features, token_mask, device) for head in e13_heads]\n            pass2_probability = _rad_np.mean(_rad_np.stack(pass2_predictions), axis=0)\n            pass2_rank = _rad_rank_columns(pass2_probability)\n            final[_RAD_LABELS] = (1.0 - _RAD_V48_SECOND_ALPHA) * _rad_rank_columns(e10[_RAD_LABELS].to_numpy()) + _RAD_V48_SECOND_ALPHA * pass2_rank\n        else:\n            print('PASS2_ENABLED=False: second-pass E13 rescoring skipped (measured -0.001 on the public LB).')\n        _rad_validate(final, expected_ids)\n        final.to_csv(primary, index=False)\n    _rad_main()"},{"cell_type":"markdown","metadata":{},"source":"## What we measured (real leaderboard receipts)\n\nEvery row below is a submission that was actually scored, not a CV estimate. They\nare here so a reader can skip the dead ends we already paid for.\n\n| Change from the defaults | Public LB | Submission ref |\n|---|---|---|\n| none - this notebook as shipped | **0.920** | 55616077 (2026-08-19) |\n| `PASS2_ENABLED = False` | 0.919 | 55636360 (2026-08-20) |\n| `RAD_ALPHA` 0.50 -> 0.35 and `RAD_E13_MEMBER_WEIGHT` 0.50 -> 0.65 | 0.919 | 55620998 |\n| `TIME_BUDGET` 3h -> 5h, on the pilkwang baseline kernel (not this pipeline) | 0.891, byte-identical file | 55613482 |\n\nWhat they mean:\n\n- **The second pass is worth about +0.001.** Real, small, and it costs a full extra\n  encode pass over the test set. If you are short on time that is the first thing\n  to drop, and you should expect to land near 0.919.\n- **The retuned blend weights lost.** That 0.35 / 0.65 pair is what our nested CV\n  preferred; the leaderboard put it 0.001 below the flat 0.50 / 0.50 defaults. For\n  these two knobs, offline CV was not a reliable guide to the public LB - which is\n  worth knowing before you spend a day tuning them.\n- **The time budget did nothing where we tested it.** On the pilkwang baseline\n  kernel, 3h and 5h produced byte-for-byte identical submissions and the same\n  0.891, so that kernel is not budget-throttled. We have not run the same test on\n  this pipeline, so `TIME_BUDGET_HOURS` here is documented rather than measured.\n\nTwo of those three are losses and the third is a null. One of the losses is a\nchange our own offline validation recommended."},{"cell_type":"markdown","metadata":{},"source":"## Provenance, stated plainly\n\nThis notebook documents and reproduces public community artifacts. Other people\ntrained these models and published them; what we did was mount them, wire the\nstages together, verify the result on the leaderboard, and write down what each\nknob is worth. That is reproduction and measurement, not novel modelling, and it\nseems worth saying before anyone forks it.\n\nOne number we cannot account for: the upstream notebook this reproduces displays\n**0.922** for its own submission, and this pipeline measures **0.920** from the\nsame recipe. We do not know what explains the 0.002. The cause is not published\nanywhere we could find, and guessing at it here would not help anyone. If you find\nit, it is worth a comment."}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12"}},"nbformat":4,"nbformat_minor":4}