{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# RSNA Knee DINO-RadImageNet + MaxSpan + ConvNeXt Rank Ensemble\n\n## Model inventory\n\nThis inference ensemble uses 35 checkpoint members and one shared encoder:\n\n- 20 DINOv2-small checkpoints from five folds;\n- 5 DINOv3-small fold checkpoints;\n- 5 reference RadImageNet attention heads;\n- 5 E13 RadImageNet attention heads;\n- 1 shared RadImageNet ResNet-50 encoder for the ten RadImageNet heads.\n\nThe five E13 heads are reused on a second image-slot layout. Reuse adds a second prediction view, not five additional unique checkpoints.\n\n### Why only DINOv2 is an external model\n\nEach DINOv3 fold file (`m_f0.pt` through `m_f4.pt`) is a complete fine-tuned network, not a head or delta checkpoint. The code calls `timm.create_model(..., pretrained=False)` only to construct the ViT-S/16 architecture, then loads all 162 DINOv3 backbone tensors plus the slot-conditioning and readout tensors from that fold file. The original DINOv3 weights are therefore already embedded in the consolidated checkpoint dataset and no separate DINOv3 model mount is required.\n\nThe DINOv2 branch is packaged differently: it constructs its encoder with `AutoModel.from_pretrained` from the attached `metaresearch/dinov2` model before applying each competition checkpoint. That is why DINOv2 remains the notebook's single external model source.\n\n## Pipeline\n\n1. DICOM series are classified by plane, fat suppression, ordering, and laterality. The DINOv2 branch builds six slots at 336 px and evaluates the 20 checkpoints over slice windows.\n2. Each DINOv2 checkpoint produces the pinned public-frontier prediction. Checkpoint predictions are converted to percentile ranks and equally averaged.\n3. The five DINOv3 fold models run on their six-slot 336 px representation. Their fold predictions are converted to ranks and averaged.\n4. The transformer parent is `0.55 * DINOv2 rank ensemble + 0.45 * DINOv3 rank ensemble`.\n5. The shared RadImageNet encoder extracts 2048-dimensional slice features. Five reference heads run on the three-plane E10 layout, while five E13 heads run on a four-slot fat-sensitive layout.\n6. Reference-head ranks and E13 ranks are mixed `0.50 / 0.50`, then ranked again to form the Rad branch.\n7. For ten targets, E10 mixes `0.50 * transformer parent rank + 0.50 * Rad rank`. Baker's cyst and Fracture preserve the transformer parent at this stage.\n8. The same five E13 heads run again on the E11 slot layout. The final prediction is `0.85 * E10 branch rank + 0.15 * second-pass E13 rank` for all twelve targets.\n\nSau nhánh MaxSpan, notebook chạy thêm 5 fold ConvNeXt Spatial MIL từ dataset\n`checkpoint-conv`. Nhánh này được kiểm tra tương quan trước khi nhận trọng số nhỏ\n5%; tỉ lệ cuối là `0.475 DINO/Rad + 0.475 MaxSpan + 0.050 ConvNeXt`.\nNotebook vẫn lưu riêng baseline 50/50 để có đối chứng không bị ghi đè.\n","metadata":{"papermill":{"duration":0.005285,"end_time":"2026-08-26T08:49:13.777968+00:00","exception":false,"start_time":"2026-08-26T08:49:13.772683+00:00","status":"completed"}}},{"cell_type":"code","source":"from __future__ import annotations\nimport os\nimport gc\nimport hashlib\nimport json\nimport re\nimport time\nimport traceback\nimport threading\nfrom concurrent.futures import ThreadPoolExecutor\nfrom pathlib import Path\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nASSET = Path('/kaggle/input/datasets/tonylica/rsna-knee-bend-dinov3-0917-repro-assets')\nROOT = Path('/kaggle/input/competitions/rsna-knee-abnormality-detection')\nDINO = Path('/kaggle/input/models/metaresearch/dinov2/pytorch/small/1')\nT0 = time.time()\nDEVS = [torch.device(f'cuda:{i}') for i in range(torch.cuda.device_count())]\nSEED = 2026\nTARGETS = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture']\nCROP_MM = 130.0\nCACHE_IMG = 336\nGROUP = 3\nN_GROUP_MAX = 1\nCACHE_FRACTION = 0.45\nCACHE_BUDGET_MAX_GB = 24.0\nCACHE_BUDGET_GB = 12.0\nTEST_SHARE = 0.3\nHDR_THREADS = 16\nPIX_THREADS = 12\nORDER_THREADS = 32\nORDER_BUDGET_S = 5400\nAUG_ROT_DEG = 8.0\nAUG_SCALE = 0.08\nAUG_SHIFT = 0.05\nAUG_INTENSITY = 0.1\nLAT_MIN_OFFSET_MM = 20.0\nSLICE_BAND = (0.2, 0.8)\nRULES_NATIVE = {'order': 'normal', 'lat': 'centre', 'slot_fallback': False, 'decode_fill': 'nearest'}\nRULES_LEGACY = {'order': 'dominant_axis', 'lat': 'corner_x', 'slot_fallback': True, 'decode_fill': 'zero'}\nRULES = dict(RULES_NATIVE)\nLEGACY_LAT_OFFSET_MM = 5.0\nEVAL_BATCH = 8\nTIME_BUDGET = 8.0 * 3600\nSLOTS_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)]\nSLOTS_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)]\nSLOT_SCHEME = os.environ.get('SLOT_SCHEME', 'recovered')\nSLOTS = SLOTS_PUBLIC if SLOT_SCHEME == 'public' else SLOTS_RECOVERED\nN_SLOT = len(SLOTS)\nPOOL_PARTS = {'cls_mean': 2, 'cls_mean_focal': 3}\nSLOT_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)}\nSLOT_PRIOR_STRENGTH = 0.55\nFATSAT_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\ndef log(msg):\n    print(f'[{time.time() - T0:7.1f}s] {msg}', flush=True)\nIMG = CACHE_IMG\n\ndef 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\ndef 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\nN_GROUP = plan_cache(len(pd.read_csv(ROOT / 'train.csv')), len(pd.read_csv(ROOT / 'test.csv')))\nCACHE_SLICES = GROUP * N_GROUP\nHDR_TAGS = ['SeriesDescription', 'SequenceName', 'ScanOptions', 'ScanningSequence', 'RepetitionTime', 'EchoTime', 'Laterality', 'PixelSpacing', 'Rows', 'Columns', 'RescaleSlope', 'RescaleIntercept', 'ImagePositionPatient', 'ImageOrientationPatient']\n\ndef _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\ndef 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\ndef 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\ndef 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\ndef 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\ndef 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\ndef 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\ndef 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\nORDER_TAGS = [(32, 50), (32, 55), (32, 19)]\nDECODE_FAILED = []\n\ndef _natural_key(name):\n    return tuple((int(x) if x.isdigit() else x.lower() for x in re.split('(\\\\d+)', str(name))))\n\ndef _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\ndef 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\ndef 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\ndef 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])\nORDER_CACHE = os.environ.get('RSNA_ORDER_CACHE') or None\n\ndef 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\nclass 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\nclass 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\ndef 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)\nFINGERPRINT_TOL = 0.002\n\ndef 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\ndef 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\nclass WeightsError(RuntimeError):\n    pass\nTTA_OVERLAP = True\nTTA_POOL = 'prob'\nPUBLIC_FRONTIER_TARGET_POOL = {'Fracture': 'max', 'Contusion': 'max', 'Medial Meniscus': 'max', 'Lateral Meniscus': 'max', 'ACL': 'top2', 'MCL': 'top2', \"Baker's\": 'max'}\nTTA_TARGET_POOL = {**PUBLIC_FRONTIER_TARGET_POOL, 'Synovitis': 'original_mean'}\nLEGACY_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}\nLEGACY_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\ndef 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\ndef 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\ndef 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()\ndef 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)\nBUILD_LOCK = threading.Lock()\nSTATE_LOCK = threading.Lock()\n\ndef _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\ndef _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\ndef 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\ndef 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\ndef 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\ndef 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\ndef 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\ndef 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\ndef find_dinov2(variant='small'):\n    if not (DINO / 'config.json').is_file():\n        raise FileNotFoundError(DINO)\n    return DINO\n\ndef legacy_group_members():\n    return {}\n\ndef 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()\nrun_dinov2()","metadata":{"lines_to_next_cell":2,"papermill":{"duration":72.184316,"end_time":"2026-08-26T08:50:25.96625+00:00","exception":false,"start_time":"2026-08-26T08:49:13.781934+00:00","status":"completed"},"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T13:45:03.804955Z","iopub.execute_input":"2026-08-30T13:45:03.805351Z","iopub.status.idle":"2026-08-30T13:46:15.099588Z","shell.execute_reply.started":"2026-08-30T13:45:03.805317Z","shell.execute_reply":"2026-08-30T13:46:15.098877Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"_A5_SAVED = dict(globals())\nimport gc, os, time, warnings\nfrom concurrent.futures import ProcessPoolExecutor, as_completed\nfrom pathlib import Path\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nwarnings.filterwarnings('ignore')\ncv2.setNumThreads(1)\nCROP_MM = 130.0\nSIZE = 336\nSLICE_BAND = (0.12, 0.88)\nN_SLICE = 16\nINTENSITY = 'slice'\nSLOTS = [('Sagittal', 1), ('Sagittal', 0), ('Coronal', 1), ('Coronal', 0), ('Axial', 1), ('Axial', 0)]\nN_SLOT = len(SLOTS)\nLABELS = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture']\nCOMP = Path('/kaggle/input/competitions/rsna-knee-abnormality-detection')\nCKPT = ASSET / 'knee-mri-fold-weights'\nDEV = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(f'competition : {COMP}')\nprint(f'checkpoints : {CKPT}')\nprint(f'device      : {DEV}')\nfor 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)}')\nSERIES_ROOT = COMP / 'test_series'\nif not SERIES_ROOT.exists():\n    SERIES_ROOT = COMP / 'train_series'\nprint('series root:', SERIES_ROOT)\n\ndef 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\ndef 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\ndef 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\ndef 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\ndef 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\ndef 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)\nsub_df = pd.read_csv(COMP / 'sample_submission.csv')\nser_csv = pd.read_csv(COMP / 'test_series.csv')\nif not (COMP / 'test_series').exists():\n    ser_csv = pd.read_csv(COMP / 'train_series.csv')\nser_csv = ser_csv.loc[:, ~ser_csv.columns.duplicated()]\nstudies = sub_df.StudyInstanceUID.tolist()\nby = {s: g.to_dict('records') for s, g in ser_csv[ser_csv.StudyInstanceUID.isin(set(studies))].groupby('StudyInstanceUID')}\nprint(f'{len(studies):,} test studies, {len(by):,} with series metadata')\nN_SLOT_TYPES, MASK_IDX = (6, 0)\n\ndef 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\nclass 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\nclass 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\nclass 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\nclass 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)\nIMAGENET_MEAN = (0.485, 0.456, 0.406)\nIMAGENET_STD = (0.229, 0.224, 0.225)\n\nclass _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\nclass 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\nN_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\nclass 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\ndef _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\ndef _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\nclass _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\nclass 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\nclass 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\nclass 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\nclass 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\nclass 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)\nmodels = []\nfor 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']}\")\nCFG = cfg\nassert 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\"\nprint(f\"\\n{len(models)} fold models ready | input norm: {CFG.get('norm', 'none')}\")\nAMP_PREF = 'bf16'\n\ndef 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)\nAMP_DT, AMP_ON = amp_for(DEV)\nWORKERS = max(1, min(4, os.cpu_count() or 4))\nCHUNK = 48\nMICRO = 8\nmodels = [m.to(DEV).eval() for m in models]\nprint(f\"device {DEV} | amp {str(AMP_DT).split('.')[-1]} (on={AMP_ON}) | workers {WORKERS} | chunk {CHUNK} | micro {MICRO}\")\n\ndef _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()\ndef _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\ndef 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\npreds = np.full((len(models), len(studies), len(LABELS)), np.nan, np.float32)\nt0, done = (time.time(), 0)\nwith 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()\nprint(f'\\ninference done in {(time.time() - t0) / 60:.1f} min')\nA5_W = 0.45\nA5_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)\nfor 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\nA5_PREDS = dict(zip(sub_df['StudyInstanceUID'].astype(str), _a5_rank_mean.astype(np.float32)))\nfor _a5k, _a5v in _A5_SAVED.items():\n    globals()[_a5k] = _a5v\ndel _A5_SAVED, _a5k, _a5v\n_a5_sub = pd.read_csv('/kaggle/working/submission.csv', dtype={'StudyInstanceUID': str})\nassert _a5_sub.columns.tolist()[1:] == A5_LABELS, 'submission schema drift'\nif 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)","metadata":{"lines_to_next_cell":2,"papermill":{"duration":13.781777,"end_time":"2026-08-26T08:50:39.756592+00:00","exception":false,"start_time":"2026-08-26T08:50:25.974815+00:00","status":"completed"},"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T13:46:15.101141Z","iopub.execute_input":"2026-08-30T13:46:15.101696Z","iopub.status.idle":"2026-08-30T13:46:30.069201Z","shell.execute_reply.started":"2026-08-30T13:46:15.101665Z","shell.execute_reply":"2026-08-30T13:46:30.068345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from __future__ import annotations\nimport contextlib as _rad_contextlib\nimport gc as _rad_gc\nimport hashlib as _rad_hashlib\nimport json as _rad_json\nimport os as _rad_os\nimport re as _rad_re\nimport time as _rad_time\nfrom concurrent.futures import ThreadPoolExecutor as _RadThreadPool\nfrom pathlib import Path as _RadPath\nimport numpy as _rad_np\nimport pandas as _rad_pd\nimport pydicom as _rad_pydicom\nimport torch as _rad_torch\nimport torch.nn as _rad_nn\nimport torch.nn.functional as _rad_F\nfrom 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 = 0.5\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 = 0.5\n_RAD_V48_SECOND_ALPHA = 0.15\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\nSLOTS = [('SAG_FS', 'Sagittal', None, True), ('COR_FS', 'Coronal', None, True), ('AX_FS', 'Axial', None, True)]\nN_SLOT = len(SLOTS)\nCACHE_SLICES = 8\n\ndef _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\ndef _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\nclass _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\nclass _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\ndef _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\ndef _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()\ndef _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()\ndef _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\ndef _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\ndef _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\ndef _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    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 = e10.copy()\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    _rad_validate(final, expected_ids)\n    final.to_csv(primary, index=False)\n_rad_main()","metadata":{"lines_to_next_cell":2,"papermill":{"duration":9.556835,"end_time":"2026-08-26T08:50:49.321857+00:00","exception":false,"start_time":"2026-08-26T08:50:39.765022+00:00","status":"completed"},"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T13:46:30.07036Z","iopub.execute_input":"2026-08-30T13:46:30.070703Z","iopub.status.idle":"2026-08-30T13:46:40.537706Z","shell.execute_reply.started":"2026-08-30T13:46:30.070675Z","shell.execute_reply":"2026-08-30T13:46:40.537101Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!/usr/bin/env python3\n\"\"\"Knee MRI - Wide-Span Dense Corpus -- RSNA Knee Abnormality Detection.\n\nThe same CoAtNet 384 model and the same training recipe as the 0.926 submission, on a\ncorpus rebuilt once more. That submission sampled each series across 6 to 94 percent of\nthe stack, which was worth 0.006 over the previous model. This one keeps that span and\ntakes 64 slices per study instead of 44, so small structures are less likely to fall\nbetween the sampled slices. The proportions across the five series slots are unchanged.\n\nScoring still uses 42 slice windows per study, exactly as the 0.926 run did, so the only\ndifference between the two is how densely the corpus samples each series. Taking more\nslices per study costs nothing at inference, since the model scores a fixed number of\nwindows either way.\n\nOn the 58 radiologist-labelled held-out studies this reaches 0.9167 macro-AUC against\n0.9054 for the model behind the current submission, with lateral meniscus, the weakest\nfinding, moving from 0.799 to 0.881. These weights are the single best epoch; averaging\nthe best three epochs gave no gain on this run, unlike the previous one.\n\nInference runs in fp16 on a Turing or newer GPU, with an automatic fp32 retry so no\nstudy is ever dropped. Internet is off and timm is forced offline.\n\"\"\"\nimport os, sys, glob, time, json, gc\nos.environ.setdefault(\"HF_HUB_OFFLINE\", \"1\")\nos.environ.setdefault(\"TRANSFORMERS_OFFLINE\", \"1\")\nos.environ.setdefault(\"HF_HUB_DISABLE_TELEMETRY\", \"1\")\nimport numpy as np\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport timm\n# T4 (Turing) cuDNN v9 has fp16/fp32 conv engines but NOT bf16 for these shapes\n# (\"GET was unable to find an engine...\"); benchmark lets it pick a valid algo for\n# the fixed (1,24,3,res,res) input.\ntorch.backends.cudnn.benchmark = True\ntorch.backends.cuda.matmul.allow_tf32 = True\n\n# ---- fixed config (must match training exactly) -----------------------------\nIMG = 336\nCROP_MM = 140.0\n# 64 slices per study instead of 44, same proportions. Must match the corpus the weights\n# were trained on (knee_corpus_v4.py).\nSLOTS = [(\"Sagittal\", 1, 18), (\"Sagittal\", 0, 14), (\"Coronal\", 1, 12),\n         (\"Coronal\", 0, 8), (\"Axial\", -1, 12)]\nMAXS = sum(s[2] for s in SLOTS)                     # 64\nK_EVAL = 42   # every window position the volume holds, not an evenly spaced subset\nNORM = \"imagenet\"\nLAB = [\"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\", \"Medial OA\", \"Lateral OA\",\n       \"PF OA\", \"Effusion\", \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\"]\n_MEAN = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)\n_STD = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)\n\n# Three arms: (weights filename, fallback arch, fallback res). ck carries arch+res too.\n# Selected 2026-08-19 by greedy forward selection AND exhaustive subset search over a 7-arm\n# panel on the 45-study gold set (phase2/blend_panel.py); both agree on this exact set.\n# Singles: coatnet384 0.9025 | swinbase384 0.8825 | effv2l480 0.8716.\n# Blend {coatnet+swin+effv2l} = 0.9068 (2-arm {coatnet+swin} = 0.9059, coatnet alone 0.9025).\n# Dropped as redundant: cnn336 (0.8833, the former champion), cnbase384 (0.8754),\n# cnlarge384 (0.8752), maxvit384 (0.8438).\n#\n# SINGLE ARM: coatnet_rmlp_2_rw_384 retrained on the EXPANDED 4,349-study corpus.\n#\n# Why one arm and not the 3-arm blend: on the live leaderboard CoAtNet alone scored 0.914 while\n# every blend scored 0.914-0.915, so ensembling is worth ~+0.001 there -- the ~+0.010 it showed\n# on the old 45-study gold set was gold-set noise. One arm is also 1/3 the kernel runtime.\n#\n# Corpus expansion: the corpus previously held 3,200 of the 4,349 labelled studies and only 45\n# of the 58 gold studies. Rebuilt to 4,407 studies (+37.8% training data, 58-study gate).\n#\n# Measured on the 58-study gate (the incumbent re-scored on the SAME gate for a fair compare):\n#   incumbent CoAtNet (3,155-study corpus) 0.8923\n#   this model       (4,349-study corpus) 0.9054   (+0.0131, better in 92.7% of 2000 bootstraps)\n# Biggest gains land on the findings that were capping us: Lateral Meniscus +0.071,\n# Fracture +0.057, Lateral OA +0.048, Medial Meniscus +0.035, ACL +0.028.\nARMS = [\n    {\"file\": \"raptor_ft_coatnet_v4_full.pt\", \"arch\": \"coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k\", \"res\": 384, \"w\": 1.0},\n]\n\n\n# ============================================================================\n# Model -- verbatim from finetune_raptor.py\n# ============================================================================\ndef build_backbone(arch, pretrained=False):\n    # maxvit/maxxvit/coatnet are conv-attention hybrids: NO CLS token, NO interpolatable\n    # pos-embed -> avg pool. The \"vit\" substring in \"coatnet\"/\"maxvit\" must NOT route them\n    # down the ViT path (mirrors finetune_raptor.py exactly).\n    hybrid = arch.startswith((\"maxvit\", \"maxxvit\", \"coatnet\", \"coat_\", \"convnext\"))\n    is_vit = (not hybrid) and any(k in arch for k in (\"vit\", \"deit\", \"dinov2\", \"eva\", \"beit\"))\n    kw = dict(pretrained=pretrained, num_classes=0, in_chans=3)\n    if is_vit:\n        kw.update(global_pool=\"token\", dynamic_img_size=True)\n    else:\n        kw.update(global_pool=\"avg\")\n    return timm.create_model(arch, **kw)\n\n\nclass RaptorClassifier(nn.Module):\n    def __init__(self, backbone, F_dim=768, n=12, drop=0.2):\n        super().__init__()\n        self.backbone = backbone\n        self.norm = nn.LayerNorm(F_dim)\n        self.att = nn.Sequential(nn.Linear(F_dim, 256), nn.Tanh(), nn.Dropout(drop),\n                                 nn.Linear(256, n))\n        self.clsW = nn.Parameter(torch.zeros(n, F_dim))\n        self.clsb = nn.Parameter(torch.zeros(n))\n        nn.init.trunc_normal_(self.clsW, std=0.02)\n        self.n = n\n\n    def encode(self, x):\n        B, K = x.shape[:2]\n        f = self.backbone(x.flatten(0, 1))\n        return f.view(B, K, -1)\n\n    def head(self, feats):\n        h = self.norm(feats)\n        a = self.att(h)\n        a = torch.softmax(a, dim=1)\n        pooled = torch.einsum(\"bkn,bkf->bnf\", a, h)\n        logits = (pooled * self.clsW).sum(-1) + self.clsb\n        return logits\n\n    def forward(self, x):\n        return self.head(self.encode(x))\n\n\ndef load_model(pt_path, arch_default, res_default, device, ngpu=1):\n    ck = torch.load(pt_path, map_location=\"cpu\", weights_only=False)\n    arch = ck.get(\"arch\", arch_default)\n    ck_res = int(ck.get(\"res\", res_default))\n    bb = build_backbone(arch, pretrained=False)\n    model = RaptorClassifier(bb, F_dim=bb.num_features)\n    model.load_state_dict(ck[\"model\"], strict=True)\n    model.eval().to(device)\n    # NOTE: DataParallel removed on purpose. On the full hidden test it drove a system-RAM OOM\n    # (per-forward module replication over many studies); a single T4 handles K_EVAL=24 windows\n    # fine. Arms are also run SEQUENTIALLY (see main) so peak RAM == one model, not two.\n    del ck\n    gc.collect()\n    return model, ck_res\n\n\n# ============================================================================\n# Eval windowing -- verbatim from finetune_raptor.py StudyWindows (train=False)\n# ============================================================================\ndef _eval_centers(mask, D, k):\n    valid = np.where(mask > 0)[0]\n    if len(valid) < 3:\n        valid = np.arange(min(3, D))\n    lo, hi = int(valid.min()), int(valid.max())\n    cs = [c for c in range(lo + 1, hi) if c - 1 >= lo and c + 1 <= hi]\n    if not cs:\n        cs = [max(1, min((lo + hi) // 2, D - 2))]\n    idx = np.linspace(0, len(cs) - 1, k).round().astype(int)\n    return [cs[i] for i in idx]\n\n\ndef eval_windows(vol, mask, k, res, norm=NORM):\n    D = vol.shape[0]\n    cs = _eval_centers(mask, D, k)\n    wins = np.empty((len(cs), 3, res, res), np.float32)\n    for j, c in enumerate(cs):\n        c = max(1, min(c, D - 2))\n        tri = np.stack([vol[c - 1], vol[c], vol[c + 1]], 0).astype(np.float32) / 255.0\n        t = torch.from_numpy(tri)\n        if t.shape[-1] != res:\n            t = F.interpolate(t[None], size=(res, res), mode=\"bilinear\",\n                              align_corners=False)[0]\n        wins[j] = t.numpy()\n    x = torch.from_numpy(wins)\n    if norm == \"imagenet\":\n        x = (x - _MEAN) / _STD\n    return x\n\n\n@torch.no_grad()\ndef infer_probs(model, xwins, device):\n    x = xwins.unsqueeze(0).to(device)\n    use_cuda = device != \"cpu\" and str(device).startswith(\"cuda\")\n    if use_cuda:\n        # fp16 conv on T4 is fully cuDNN-supported (bf16 is NOT -> \"no engine\").\n        try:\n            with torch.autocast(\"cuda\", dtype=torch.float16):\n                o = torch.sigmoid(model(x).float())\n            return o[0].cpu().numpy()\n        except RuntimeError:\n            # fp32 always has a Turing conv engine; slower but never drops a study.\n            torch.cuda.empty_cache()\n            o = torch.sigmoid(model(x).float())\n            return o[0].cpu().numpy()\n    o = torch.sigmoid(model(x).float())\n    return o[0].cpu().numpy()\n\n\ndef rankpct(x):                                   # per-column percentile rank in [0,1]\n    order = x.argsort(0).argsort(0).astype(np.float64)\n    return order / max(1, (x.shape[0] - 1))\n\n\n# ============================================================================\n# Preprocessing -- verbatim from kprep2/dino_preprocess.py, retargeted to TEST\n# ============================================================================\ndef _make_reader():\n    import pydicom, cv2\n    from pydicom.pixel_data_handlers.util import apply_modality_lut\n\n    def order_and_meta(sdir):\n        fs = glob.glob(sdir + \"/*.dcm\"); recs = []; ps_list = []\n        for f in fs:\n            try:\n                h = pydicom.dcmread(f, stop_before_pixels=True)\n                iop = getattr(h, 'ImageOrientationPatient', None)\n                ipp = getattr(h, 'ImagePositionPatient', None)\n                if iop is not None and ipp is not None and len(iop) == 6:\n                    r = np.array(iop[:3], float); c = np.array(iop[3:], float)\n                    n = np.cross(r, c); pos = float(np.dot(np.array(ipp, float), n))\n                else:\n                    pos = float(getattr(h, 'InstanceNumber', 0) or 0)\n                ps = getattr(h, 'PixelSpacing', None); ps = float(ps[0]) if ps is not None else 0.5\n                ps_list.append(ps); recs.append((pos, f, ps))\n            except Exception:\n                recs.append((0.0, f, 0.5))\n        recs.sort(key=lambda x: x[0])\n        med_ps = float(np.median(ps_list)) if ps_list else 0.5\n        return [(f, ps) for _, f, ps in recs], med_ps\n\n    def read_px(f):\n        d = pydicom.dcmread(f)\n        a = apply_modality_lut(d.pixel_array, d).astype(np.float32)\n        if str(getattr(d, 'PhotometricInterpretation', '')) == 'MONOCHROME1':\n            a = a.max() - a\n        return a\n\n    def mm_crop_resize(a, ps):\n        h, w = a.shape; cpx = int(round(CROP_MM / max(ps, 1e-3)))\n        cpx = min(cpx, min(h, w)); y0 = (h - cpx) // 2; x0 = (w - cpx) // 2\n        a = a[y0:y0 + cpx, x0:x0 + cpx]\n        return cv2.resize(a, (IMG, IMG), interpolation=cv2.INTER_AREA)\n\n    return order_and_meta, read_px, mm_crop_resize\n\n\ndef _pick_series_for_slot(rows, plane, fluid, used):\n    cands = [r for r in rows if r['Anatomical_Plane'] == plane and r['SeriesInstanceUID'] not in used]\n    if fluid in (0, 1):\n        pref = [r for r in cands if int(r.get('Fluid_Sensitive', 0) or 0) == fluid]\n        if pref:\n            return pref[0]\n    return cands[0] if cands else None\n\n\ndef build_study(sid, ser_records, tsdir, reader):\n    order_and_meta, read_px, mm_crop_resize = reader\n    rows = ser_records.get(sid, [])\n    vol = np.zeros((MAXS, IMG, IMG), np.uint8); idx = 0; used = set()\n    for plane, fluid, k in SLOTS:\n        r = _pick_series_for_slot(rows, plane, fluid, used)\n        if r is None:\n            idx += k; continue\n        used.add(r['SeriesInstanceUID'])\n        files, med_ps = order_and_meta(f\"{tsdir}/{sid}/{r['SeriesInstanceUID']}\")\n        if not files:\n            idx += k; continue\n        # wide span: the collateral ligaments and lateral meniscus live in the\n        # peripheral slices the old 0.15-0.85 crop threw away. Must match the corpus\n        # the weights were trained on (knee_corpus_v2.py, SPAN_LO/SPAN_HI).\n        n = len(files); lo, hi = int(n * 0.06), int(n * 0.94) - 1; hi = max(hi, lo)\n        picks = np.linspace(lo, hi, k).round().astype(int) if n > 1 else [0] * k\n        arrs = []; pss = []\n        for p in picks:\n            fp, ps = files[min(p, n - 1)]\n            try:\n                arrs.append(read_px(fp)); pss.append(ps)\n            except Exception:\n                arrs.append(None); pss.append(med_ps)\n        valid = [a for a in arrs if a is not None]\n        if valid:\n            allpx = np.concatenate([a.ravel() for a in valid])\n            loq, hiq = np.percentile(allpx, [2.0, 98.0])\n        else:\n            loq, hiq = 0.0, 1.0\n        for a, ps in zip(arrs, pss):\n            if idx >= MAXS: break\n            if a is None: idx += 1; continue\n            aw = np.clip((a - loq) / (hiq - loq + 1e-6), 0, 1)\n            aw = mm_crop_resize(aw, ps if ps > 0 else med_ps)\n            vol[idx] = (aw * 255).astype(np.uint8); idx += 1\n        if idx >= MAXS: break\n    mask = (vol.reshape(MAXS, -1).sum(1) > 0).astype(np.uint8)\n    return vol, mask\n\n\n# ============================================================================\n# Test-root discovery + weights + main\n# ============================================================================\ndef find_test_root():\n    cands = [\"/kaggle/input/competitions/rsna-knee-abnormality-detection\",\n             \"/kaggle/input/rsna-knee-abnormality-detection\"]\n    for b in cands:\n        if os.path.exists(b + \"/test.csv\"):\n            return b\n    for d, _, f in os.walk(\"/kaggle/input\"):\n        if \"test.csv\" in f and (os.path.isdir(d + \"/test_series\") or os.path.isdir(d + \"/test_images\")):\n            return d\n    for d, _, f in os.walk(\"/kaggle/input\"):\n        if \"test.csv\" in f:\n            return d\n    raise RuntimeError(\"no test root under /kaggle/input\")\n\n\ndef find_weight_file(fname):\n    # direct dataset mounts first; NEVER recursive-glob the competitions DICOM tree.\n    direct = [f\"/kaggle/input/raptor-knee-arms/{fname}\",\n              f\"/kaggle/input/raptor-knee-arms/1/{fname}\",\n              f\"/kaggle/input/raptor-cnn336/{fname}\"]\n    for p in direct:\n        if os.path.exists(p):\n            return p\n    for d in sorted(glob.glob(\"/kaggle/input/*/\")):\n        if \"competition\" in d.lower():\n            continue\n        hits = glob.glob(os.path.join(d, \"**\", fname), recursive=True)\n        if hits:\n            return hits[0]\n    raise RuntimeError(f\"{fname} not found under /kaggle/input\")\n\n\ndef main():\n    import pandas as pd\n    t0 = time.time()\n    dev = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    ngpu = torch.cuda.device_count()\n    print(f\"device {dev} | gpus {ngpu} | torch {torch.__version__}\", flush=True)\n\n    ROOT = find_test_root()\n    tsdir = ROOT + \"/test_series\"\n    if not os.path.isdir(tsdir):\n        tsdir = ROOT + \"/test_images\"\n    print(\"test root:\", ROOT, \"| series dir:\", tsdir, flush=True)\n\n    test = pd.read_csv(ROOT + \"/test.csv\"); test[\"StudyInstanceUID\"] = test[\"StudyInstanceUID\"].astype(str)\n    test_ids = test[\"StudyInstanceUID\"].tolist()\n    tser = pd.read_csv(ROOT + \"/test_series.csv\")\n    tser[\"StudyInstanceUID\"] = tser[\"StudyInstanceUID\"].astype(str)\n    tser[\"SeriesInstanceUID\"] = tser[\"SeriesInstanceUID\"].astype(str)\n    SER = {k: v.to_dict(\"records\") for k, v in tser.groupby(\"StudyInstanceUID\")}\n    print(f\"test studies {len(test_ids)} | test series {len(tser)}\", flush=True)\n\n    sub_cols = [\"StudyInstanceUID\"] + LAB\n    ssub = os.path.join(ROOT, \"sample_submission.csv\")\n    if os.path.exists(ssub):\n        sub_cols = list(pd.read_csv(ssub, nrows=1).columns)\n\n    reader = _make_reader()\n    N = len(test_ids); A = len(ARMS)\n    arm_probs = [np.full((N, len(LAB)), 0.5, np.float32) for _ in range(A)]\n\n    # SEQUENTIAL ARMS (the OOM fix): only ONE model is resident at a time, so peak system RAM ==\n    # one model == the single-arm champion's footprint (which graded fine at 0.879). Holding both\n    # arms simultaneously OOM'd system RAM on the full hidden test. Each study is re-preprocessed\n    # per arm (build_study is cheap vs inference) and every per-study buffer is freed. Same models,\n    # same windowing, same rank-mean blend -> identical 0.8893 result, just serialized.\n    for a, arm in enumerate(ARMS):\n        wp = find_weight_file(arm[\"file\"])\n        model, res = load_model(wp, arm[\"arch\"], arm[\"res\"], dev)\n        print(f\"[arm {a}] loaded {arm['file']} | res {res} | {time.time()-t0:.0f}s\", flush=True)\n        for i, sid in enumerate(test_ids):\n            try:\n                vol, mask = build_study(sid, SER, tsdir, reader)\n                xw = eval_windows(vol, mask, k=K_EVAL, res=res, norm=NORM)\n                arm_probs[a][i] = infer_probs(model, xw, dev)\n                del vol, mask, xw\n            except Exception as e:\n                print(f\"  [arm {a}] study {i} {sid[:16]} FALLBACK ({type(e).__name__}: {e})\", flush=True)\n            if (i + 1) % 100 == 0 or i + 1 == N:\n                print(f\"  [arm {a}] {i+1}/{N} | {time.time()-t0:.0f}s\", flush=True)\n        del model\n        gc.collect()\n        if str(dev).startswith(\"cuda\"):\n            torch.cuda.empty_cache()\n        print(f\"[arm {a}] done + freed | {time.time()-t0:.0f}s\", flush=True)\n\n    # WEIGHTED rank-mean blend across the test set, per finding (the offline recipe).\n    # Weights come from ARMS[*][\"w\"] and are normalised here, so dropping/adding an arm can\n    # never silently change the scale. Falls back to equal weights if none are declared.\n    _w = np.array([float(a.get(\"w\", 1.0)) for a in ARMS], dtype=np.float64)\n    _w = _w / _w.sum()\n    print(f\"[blend] weighted rank-mean w={dict(zip([a['file'] for a in ARMS], _w.round(4)))}\", flush=True)\n    ranks = np.tensordot(_w, np.stack([rankpct(np.clip(p, 0, 1)) for p in arm_probs]),\n                         axes=(0, 0))                                          # (N,12) in [0,1]\n    if not np.isfinite(ranks).all():\n        ranks[~np.isfinite(ranks)] = 0.5\n\n    sub = pd.DataFrame(ranks.astype(np.float32), columns=LAB)\n    sub.insert(0, \"StudyInstanceUID\", test_ids)\n    sub = sub[sub_cols]\n    assert list(sub.columns) == sub_cols, \"column order drift\"\n    assert sub[\"StudyInstanceUID\"].tolist() == test_ids, \"row identity drift\"\n    assert np.isfinite(sub[LAB].values).all()\n    out = \"/kaggle/working/coatnet_submission.csv\"\n    sub.to_csv(out, index=False)\n    print(\"wrote\", out, \"|\", len(sub), \"rows x\", len(sub.columns), \"cols\", flush=True)\n    print(sub.head().to_string(index=False), flush=True)\n    print(f\"DONE {time.time()-t0:.0f}s\", flush=True)\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"title":"[code]","trusted":true,"execution":{"iopub.status.busy":"2026-08-30T13:46:40.539548Z","iopub.execute_input":"2026-08-30T13:46:40.54027Z","iopub.status.idle":"2026-08-30T13:47:01.625779Z","shell.execute_reply.started":"2026-08-30T13:46:40.540214Z","shell.execute_reply":"2026-08-30T13:47:01.625131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"Paste/run as one Kaggle notebook cell after the existing 0.932 pipeline.\n\nRequired Kaggle input:\nhttps://www.kaggle.com/datasets/dreaddevelopment/raptor-knee-maxspan\n\nThis is the standalone MaxSpan branch only. It never reads or overwrites\n``/kaggle/working/submission.csv``. It writes raw probabilities and the exact\nordinal percentile ranks used by the public MaxSpan inference implementation.\n\"\"\"\n\nimport gc as ms_gc\nimport glob as ms_glob\nimport os as ms_os\nimport time as ms_time\nfrom pathlib import Path as MSPath\n\nimport cv2 as ms_cv2\nimport numpy as ms_np\nimport pandas as ms_pd\nimport pydicom as ms_pydicom\nimport timm as ms_timm\nimport torch as ms_torch\nimport torch.nn as ms_nn\nimport torch.nn.functional as ms_F\n\ntry:\n    from pydicom.pixels import apply_modality_lut as ms_apply_modality_lut\nexcept ImportError:  # pydicom < 3\n    from pydicom.pixel_data_handlers.util import apply_modality_lut as ms_apply_modality_lut\n\n\nMS_TARGETS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n    \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\", \"Synovitis\",\n    \"Baker's\", \"Contusion\", \"Fracture\",\n]\nMS_SLOTS = [\n    (\"Sagittal\", 1, 18),\n    (\"Sagittal\", 0, 14),\n    (\"Coronal\", 1, 12),\n    (\"Coronal\", 0, 8),\n    (\"Axial\", -1, 12),\n]\nMS_CACHE_SIZE = 336\nMS_CROP_MM = 140.0\nMS_N_SLICES = 64\nMS_N_WINDOWS = 62\nMS_WEIGHT_NAME = \"raptor_ft_coatnet_v5_full_swa.pt\"\nMS_ARCH = \"coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k\"\nMS_MODEL_SIZE = 384\nMS_MEAN = ms_torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)\nMS_STD = ms_torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)\n\nassert sum(slot[2] for slot in MS_SLOTS) == MS_N_SLICES\nassert MS_N_WINDOWS == MS_N_SLICES - 2\n\n\nclass MSRaptorClassifier(ms_nn.Module):\n    def __init__(self, backbone, feature_dim, n_targets=12, dropout=0.2):\n        super().__init__()\n        self.backbone = backbone\n        self.norm = ms_nn.LayerNorm(feature_dim)\n        self.att = ms_nn.Sequential(\n            ms_nn.Linear(feature_dim, 256),\n            ms_nn.Tanh(),\n            ms_nn.Dropout(dropout),\n            ms_nn.Linear(256, n_targets),\n        )\n        self.clsW = ms_nn.Parameter(ms_torch.zeros(n_targets, feature_dim))\n        self.clsb = ms_nn.Parameter(ms_torch.zeros(n_targets))\n\n    def head(self, features):\n        hidden = self.norm(features)\n        attention = ms_torch.softmax(self.att(hidden), dim=1)\n        pooled = ms_torch.einsum(\"bkn,bkf->bnf\", attention, hidden)\n        return (pooled * self.clsW).sum(-1) + self.clsb\n\n\ndef ms_find_competition_root():\n    candidates = [\n        MSPath(\"/kaggle/input/competitions/rsna-knee-abnormality-detection\"),\n        MSPath(\"/kaggle/input/rsna-knee-abnormality-detection\"),\n    ]\n    for path in candidates:\n        if (path / \"test.csv\").is_file() and (path / \"test_series.csv\").is_file():\n            return path\n    for test_csv in MSPath(\"/kaggle/input\").rglob(\"test.csv\"):\n        path = test_csv.parent\n        if (path / \"test_series.csv\").is_file() and (\n            (path / \"test_series\").is_dir() or (path / \"test_images\").is_dir()\n        ):\n            return path\n    raise FileNotFoundError(\"MaxSpan: competition test data not found under /kaggle/input\")\n\n\ndef ms_find_checkpoint():\n    candidates = [\n        MSPath(\"/kaggle/input/raptor-knee-maxspan\") / MS_WEIGHT_NAME,\n        MSPath(\"/kaggle/input/datasets/dreaddevelopment/raptor-knee-maxspan\") / MS_WEIGHT_NAME,\n        MSPath(\"/kaggle/input/raptor-knee-arms\") / MS_WEIGHT_NAME,\n    ]\n    for path in candidates:\n        if path.is_file():\n            return path\n    hits = [\n        MSPath(path) for path in ms_glob.glob(\n            f\"/kaggle/input/**/{MS_WEIGHT_NAME}\", recursive=True\n        )\n    ]\n    if len(hits) != 1:\n        raise FileNotFoundError(\n            f\"MaxSpan: expected exactly one {MS_WEIGHT_NAME}, found {len(hits)}: {hits}\"\n        )\n    return hits[0]\n\n\ndef ms_load_model(checkpoint_path, device):\n    checkpoint = ms_torch.load(checkpoint_path, map_location=\"cpu\", weights_only=False)\n    if \"model\" not in checkpoint:\n        raise RuntimeError(\"MaxSpan checkpoint contract drift: missing 'model' state_dict\")\n    arch = checkpoint.get(\"arch\", MS_ARCH)\n    resolution = int(checkpoint.get(\"res\", MS_MODEL_SIZE))\n    if arch != MS_ARCH or resolution != MS_MODEL_SIZE:\n        raise RuntimeError(\n            f\"Unexpected MaxSpan contract: arch={arch!r}, resolution={resolution}\"\n        )\n    backbone = ms_timm.create_model(\n        arch, pretrained=False, num_classes=0, in_chans=3, global_pool=\"avg\"\n    )\n    model = MSRaptorClassifier(backbone, backbone.num_features, len(MS_TARGETS))\n    model.load_state_dict(checkpoint[\"model\"], strict=True)\n    model.eval().to(device)\n    del checkpoint\n    return model\n\n\ndef ms_float(value, default):\n    try:\n        value = float(value)\n        return value if ms_np.isfinite(value) else default\n    except (TypeError, ValueError):\n        return default\n\n\ndef ms_order_series(series_dir):\n    records, spacing = [], []\n    for file_path in ms_glob.glob(ms_os.path.join(series_dir, \"*.dcm\")):\n        try:\n            header = ms_pydicom.dcmread(file_path, stop_before_pixels=True, force=True)\n            iop = getattr(header, \"ImageOrientationPatient\", None)\n            ipp = getattr(header, \"ImagePositionPatient\", None)\n            if iop is not None and ipp is not None and len(iop) == 6:\n                row = ms_np.asarray(iop[:3], dtype=float)\n                col = ms_np.asarray(iop[3:], dtype=float)\n                position = float(ms_np.dot(ms_np.asarray(ipp, dtype=float), ms_np.cross(row, col)))\n            else:\n                position = ms_float(getattr(header, \"InstanceNumber\", 0), 0.0)\n            raw_spacing = getattr(header, \"PixelSpacing\", None)\n            pixel_spacing = ms_float(raw_spacing[0], 0.5) if raw_spacing is not None else 0.5\n            spacing.append(pixel_spacing)\n            records.append((position, file_path, pixel_spacing))\n        except Exception:\n            records.append((0.0, file_path, 0.5))\n    records.sort(key=lambda item: (item[0], item[1]))\n    median_spacing = float(ms_np.median(spacing)) if spacing else 0.5\n    return [(path, ps) for _, path, ps in records], median_spacing\n\n\ndef ms_read_pixels(file_path):\n    dicom = ms_pydicom.dcmread(file_path, force=True)\n    pixels = ms_apply_modality_lut(dicom.pixel_array, dicom).astype(ms_np.float32)\n    if str(getattr(dicom, \"PhotometricInterpretation\", \"\")) == \"MONOCHROME1\":\n        pixels = pixels.max() - pixels\n    return pixels\n\n\ndef ms_crop_resize(pixels, pixel_spacing):\n    height, width = pixels.shape\n    crop_pixels = int(round(MS_CROP_MM / max(float(pixel_spacing), 0.001)))\n    crop_pixels = min(crop_pixels, height, width)\n    y0, x0 = (height - crop_pixels) // 2, (width - crop_pixels) // 2\n    pixels = pixels[y0:y0 + crop_pixels, x0:x0 + crop_pixels]\n    return ms_cv2.resize(\n        pixels, (MS_CACHE_SIZE, MS_CACHE_SIZE), interpolation=ms_cv2.INTER_AREA\n    )\n\n\ndef ms_fluid_flag(row):\n    value = ms_float(row.get(\"Fluid_Sensitive\", 0), 0.0)\n    return int(value > 0.5)\n\n\ndef ms_pick_series(rows, plane, fluid, used):\n    candidates = [\n        row for row in rows\n        if row.get(\"Anatomical_Plane\") == plane\n        and str(row.get(\"SeriesInstanceUID\")) not in used\n    ]\n    if fluid in (0, 1):\n        preferred = [row for row in candidates if ms_fluid_flag(row) == fluid]\n        if preferred:\n            return preferred[0]\n    return candidates[0] if candidates else None\n\n\ndef ms_build_study(study_id, series_by_study, image_root):\n    volume = ms_np.zeros((MS_N_SLICES, MS_CACHE_SIZE, MS_CACHE_SIZE), ms_np.uint8)\n    rows, used, offset = series_by_study.get(study_id, []), set(), 0\n    for plane, fluid, n_slices in MS_SLOTS:\n        row = ms_pick_series(rows, plane, fluid, used)\n        if row is None:\n            offset += n_slices\n            continue\n        series_id = str(row[\"SeriesInstanceUID\"])\n        used.add(series_id)\n        files, median_spacing = ms_order_series(str(image_root / study_id / series_id))\n        if not files:\n            offset += n_slices\n            continue\n\n        # Exact MaxSpan coverage: 2--98% of the ordered stack.\n        n_files = len(files)\n        lo = int(n_files * 0.02)\n        hi = max(lo, int(n_files * 0.98) - 1)\n        picks = (\n            ms_np.linspace(lo, hi, n_slices).round().astype(int)\n            if n_files > 1 else ms_np.zeros(n_slices, dtype=int)\n        )\n        arrays, spacings = [], []\n        for pick in picks:\n            file_path, pixel_spacing = files[min(int(pick), n_files - 1)]\n            try:\n                arrays.append(ms_read_pixels(file_path))\n            except Exception:\n                arrays.append(None)\n            spacings.append(pixel_spacing)\n\n        valid = [array for array in arrays if array is not None]\n        if valid:\n            pooled = ms_np.concatenate([array.ravel() for array in valid])\n            low, high = ms_np.percentile(pooled, [2.0, 98.0])\n        else:\n            low, high = 0.0, 1.0\n        for array, pixel_spacing in zip(arrays, spacings):\n            if array is not None:\n                scaled = ms_np.clip((array - low) / (high - low + 1e-6), 0, 1)\n                scaled = ms_crop_resize(\n                    scaled, pixel_spacing if pixel_spacing > 0 else median_spacing\n                )\n                volume[offset] = (scaled * 255).astype(ms_np.uint8)\n            offset += 1\n\n    mask = (volume.reshape(MS_N_SLICES, -1).sum(1) > 0).astype(ms_np.uint8)\n    return volume, mask\n\n\ndef ms_make_windows(volume, mask, resolution=MS_MODEL_SIZE):\n    valid = ms_np.flatnonzero(mask > 0)\n    if len(valid) < 3:\n        valid = ms_np.arange(min(3, MS_N_SLICES))\n    lo, hi = int(valid.min()), int(valid.max())\n    centers = [center for center in range(lo + 1, hi) if center + 1 <= hi]\n    if not centers:\n        centers = [max(1, min((lo + hi) // 2, MS_N_SLICES - 2))]\n    indices = ms_np.linspace(0, len(centers) - 1, MS_N_WINDOWS).round().astype(int)\n    centers = [centers[index] for index in indices]\n\n    windows = ms_np.empty((MS_N_WINDOWS, 3, resolution, resolution), ms_np.float32)\n    for index, center in enumerate(centers):\n        triplet = volume[center - 1:center + 2].astype(ms_np.float32) / 255.0\n        tensor = ms_torch.from_numpy(triplet)\n        if tensor.shape[-2:] != (resolution, resolution):\n            tensor = ms_F.interpolate(\n                tensor[None], size=(resolution, resolution),\n                mode=\"bilinear\", align_corners=False,\n            )[0]\n        windows[index] = tensor.numpy()\n    tensor = ms_torch.from_numpy(windows)\n    return (tensor - MS_MEAN) / MS_STD\n\n\n@ms_torch.inference_mode()\ndef ms_predict(model, windows, device, encode_batch=16):\n    # Chunk only the backbone. Attention still sees all 62 windows simultaneously,\n    # so this is functionally equivalent to the public implementation with lower VRAM.\n    features = []\n    use_amp = device.type == \"cuda\"\n    for start in range(0, len(windows), encode_batch):\n        batch = windows[start:start + encode_batch].to(device, non_blocking=True)\n        with ms_torch.autocast(\"cuda\", dtype=ms_torch.float16, enabled=use_amp):\n            features.append(model.backbone(batch))\n    features = ms_torch.cat(features, dim=0).unsqueeze(0)\n    with ms_torch.autocast(\"cuda\", dtype=ms_torch.float16, enabled=use_amp):\n        probability = ms_torch.sigmoid(model.head(features))\n    return probability[0].float().cpu().numpy()\n\n\ndef ms_ordinal_rank(values):\n    # Exact rank transform used by the public MaxSpan code (ties are ordinal).\n    order = values.argsort(axis=0).argsort(axis=0)\n    return order.astype(ms_np.float64) / max(1, len(values) - 1)\n\n\ndef run_coatnet_maxspan(\n    output_dir=\"/kaggle/working\",\n    device=\"cuda:0\",\n    encode_batch=16,\n    max_failed_studies=0,\n):\n    started = ms_time.time()\n    root = ms_find_competition_root()\n    checkpoint_path = ms_find_checkpoint()\n    image_root = root / (\"test_series\" if (root / \"test_series\").is_dir() else \"test_images\")\n    test = ms_pd.read_csv(root / \"test.csv\", dtype={\"StudyInstanceUID\": str})\n    series = ms_pd.read_csv(\n        root / \"test_series.csv\",\n        dtype={\"StudyInstanceUID\": str, \"SeriesInstanceUID\": str},\n    )\n    test_ids = test[\"StudyInstanceUID\"].astype(str).tolist()\n    if len(test_ids) != len(set(test_ids)):\n        raise RuntimeError(\"MaxSpan: duplicate StudyInstanceUID in test.csv\")\n    series_by_study = {\n        str(study_id): frame.to_dict(\"records\")\n        for study_id, frame in series.groupby(\"StudyInstanceUID\", sort=False)\n    }\n\n    torch_device = ms_torch.device(device if ms_torch.cuda.is_available() else \"cpu\")\n    ms_gc.collect()\n    if ms_torch.cuda.is_available():\n        ms_torch.cuda.empty_cache()\n    model = ms_load_model(checkpoint_path, torch_device)\n    predictions = ms_np.full((len(test_ids), len(MS_TARGETS)), 0.5, ms_np.float32)\n    failures = []\n    print(\n        f\"[MaxSpan] checkpoint={checkpoint_path} | device={torch_device} | \"\n        f\"studies={len(test_ids)} | encode_batch={encode_batch}\", flush=True,\n    )\n\n    for index, study_id in enumerate(test_ids):\n        try:\n            volume, mask = ms_build_study(study_id, series_by_study, image_root)\n            windows = ms_make_windows(volume, mask)\n            predictions[index] = ms_predict(model, windows, torch_device, encode_batch)\n            del volume, mask, windows\n        except Exception as error:\n            failures.append((study_id, type(error).__name__, str(error)[:200]))\n            print(f\"[MaxSpan] FAILED {index} {study_id}: {failures[-1][1:]}\", flush=True)\n        if (index + 1) % 100 == 0 or index + 1 == len(test_ids):\n            print(\n                f\"[MaxSpan] {index + 1}/{len(test_ids)} | \"\n                f\"failures={len(failures)} | {ms_time.time() - started:.0f}s\",\n                flush=True,\n            )\n\n    del model\n    ms_gc.collect()\n    if ms_torch.cuda.is_available():\n        ms_torch.cuda.empty_cache()\n    if len(failures) > max_failed_studies:\n        raise RuntimeError(\n            f\"MaxSpan integrity check failed: {len(failures)} studies failed; \"\n            f\"allowed={max_failed_studies}; first={failures[:3]}\"\n        )\n    if not ms_np.isfinite(predictions).all():\n        raise RuntimeError(\"MaxSpan produced non-finite probabilities\")\n\n    raw = ms_pd.DataFrame(predictions, columns=MS_TARGETS)\n    raw.insert(0, \"StudyInstanceUID\", test_ids)\n    ranked = ms_pd.DataFrame(ms_ordinal_rank(predictions), columns=MS_TARGETS)\n    ranked.insert(0, \"StudyInstanceUID\", test_ids)\n    expected_columns = [\"StudyInstanceUID\", *MS_TARGETS]\n    if raw.columns.tolist() != expected_columns or ranked.columns.tolist() != expected_columns:\n        raise RuntimeError(\"MaxSpan output schema drift\")\n\n    output_dir = MSPath(output_dir)\n    output_dir.mkdir(parents=True, exist_ok=True)\n    raw_path = output_dir / \"pred_coatnet_maxspan_raw.csv\"\n    rank_path = output_dir / \"pred_coatnet_maxspan_rank.csv\"\n    raw.to_csv(raw_path, index=False)\n    ranked.to_csv(rank_path, index=False)\n    print(\n        f\"[MaxSpan] wrote {raw_path.name} and {rank_path.name}; \"\n        f\"shape={raw.shape}; elapsed={ms_time.time() - started:.0f}s\",\n        flush=True,\n    )\n    return raw, ranked\n\n\n# Run explicitly in the notebook after the 0.932 pipeline has released its models:\nmaxspan_raw, maxspan_rank = run_coatnet_maxspan(encode_batch=16)","metadata":{"papermill":{"duration":18.80371,"end_time":"2026-08-26T08:51:28.192961+00:00","exception":false,"start_time":"2026-08-26T08:51:09.389251+00:00","status":"completed"},"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T13:47:01.626665Z","iopub.execute_input":"2026-08-30T13:47:01.626957Z","iopub.status.idle":"2026-08-30T13:47:21.736738Z","shell.execute_reply.started":"2026-08-30T13:47:01.626934Z","shell.execute_reply":"2026-08-30T13:47:21.736087Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ConvNeXt Spatial MIL — five-fold inference\n\nCell này tái tạo đúng kiến trúc và preprocessing đã dùng khi train: năm slot,\ntám lát/slot, cửa sổ 2.5D, crop vật lý 140 mm, chuẩn hoá percentile 2–98 theo\nseries và chuẩn hoá laterality. Dataset checkpoint phải chứa đủ năm file `best`\ncùng `manifest.json`; cell tự dò đường dẫn và kiểm SHA-256 trước khi inference.","metadata":{}},{"cell_type":"code","source":"import gc as cv_gc\nimport hashlib as cv_hashlib\nimport json as cv_json\nimport math as cv_math\nimport os as cv_os\nfrom pathlib import Path as CVPath\n\nimport cv2 as cv_cv2\nimport numpy as cv_np\nimport pandas as cv_pd\nimport pydicom as cv_pydicom\nimport timm as cv_timm\nimport torch as cv_torch\nimport torch.nn as cv_nn\nimport torch.nn.functional as cv_F\n\ntry:\n    from pydicom.pixels import apply_modality_lut as cv_apply_modality_lut\nexcept ImportError:\n    from pydicom.pixel_data_handlers.util import apply_modality_lut as cv_apply_modality_lut\n\n\nCV_TARGETS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n    \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\", \"Synovitis\",\n    \"Baker's\", \"Contusion\", \"Fracture\",\n]\nCV_SLOTS = [\n    (\"SAG_FLUID\", \"Sagittal\", 1),\n    (\"SAG_STRUCT\", \"Sagittal\", 0),\n    (\"COR_FLUID\", \"Coronal\", 1),\n    (\"COR_STRUCT\", \"Coronal\", 0),\n    (\"AXIAL\", \"Axial\", None),\n]\nCV_MODEL_NAME = \"convnext_tiny.fb_in22k_ft_in1k\"\nCV_MODEL_SIZE = 224\nCV_GRID_SIZE = 3\nCV_CROP_MM = 140.0\nCV_SLICE_LOW = 0.02\nCV_SLICE_HIGH = 0.98\nCV_SLICES_PER_SLOT = 8\nCV_TOPK_TOKENS = 6\nCV_EXPECTED_FILES = [f\"convnext_spatial_mil_fold{fold}_best.pt\" for fold in range(5)]\nCV_MEAN = cv_torch.tensor([0.485, 0.456, 0.406], dtype=cv_torch.float32).view(1, 1, 3, 1, 1)\nCV_STD = cv_torch.tensor([0.229, 0.224, 0.225], dtype=cv_torch.float32).view(1, 1, 3, 1, 1)\n\n\ndef cv_find_competition_root():\n    direct = [\n        CVPath(\"/kaggle/input/competitions/rsna-knee-abnormality-detection\"),\n        CVPath(\"/kaggle/input/rsna-knee-abnormality-detection\"),\n    ]\n    for root in direct:\n        if (root / \"test.csv\").is_file() and (root / \"test_series.csv\").is_file():\n            return root\n    hits = []\n    for path in CVPath(\"/kaggle/input\").rglob(\"test_series.csv\"):\n        if (path.parent / \"test.csv\").is_file():\n            hits.append(path.parent)\n    if len(hits) != 1:\n        raise FileNotFoundError(f\"Expected one competition root, found: {hits}\")\n    return hits[0]\n\n\ndef cv_sha256(path, chunk_size=8 * 1024 * 1024):\n    digest = cv_hashlib.sha256()\n    with path.open(\"rb\") as stream:\n        while True:\n            chunk = stream.read(chunk_size)\n            if not chunk:\n                break\n            digest.update(chunk)\n    return digest.hexdigest()\n\n\ndef cv_find_checkpoint_dir():\n    valid = []\n    for manifest_path in CVPath(\"/kaggle/input\").rglob(\"manifest.json\"):\n        try:\n            manifest = cv_json.loads(manifest_path.read_text(encoding=\"utf-8\"))\n        except Exception:\n            continue\n        if manifest.get(\"model_family\") != \"convnext_tiny.fb_in22k_ft_in1k spatial MIL\":\n            continue\n        if manifest.get(\"targets\") != CV_TARGETS:\n            continue\n        if all((manifest_path.parent / name).is_file() for name in CV_EXPECTED_FILES):\n            valid.append((manifest_path.parent, manifest))\n    if len(valid) != 1:\n        raise FileNotFoundError(\n            \"Attach exactly one private Kaggle dataset made from checkpoint-conv; \"\n            f\"found {len(valid)} valid directories: {[str(item[0]) for item in valid]}\"\n        )\n    root, manifest = valid[0]\n    entries = {int(item[\"fold\"]): item for item in manifest.get(\"checkpoints\", [])}\n    if set(entries) != set(range(5)):\n        raise RuntimeError(f\"Checkpoint manifest fold drift: {sorted(entries)}\")\n    for fold, filename in enumerate(CV_EXPECTED_FILES):\n        path = root / filename\n        item = entries[fold]\n        if item.get(\"file\") != filename or path.stat().st_size != int(item.get(\"bytes\", -1)):\n            raise RuntimeError(f\"Checkpoint size/filename mismatch: {path}\")\n        observed = cv_sha256(path)\n        if observed != item.get(\"sha256\"):\n            raise RuntimeError(f\"Checkpoint SHA-256 mismatch: {path}\")\n    return root\n\n\ndef cv_safe_float(value, default):\n    try:\n        value = float(value)\n        return value if cv_np.isfinite(value) else default\n    except (TypeError, ValueError):\n        return default\n\n\ndef cv_fluid_value(row):\n    return int(cv_safe_float(row.get(\"Fluid_Sensitive\", 0), 0.0) > 0.5)\n\n\ndef cv_list_dicoms(series_dir):\n    try:\n        return [\n            entry.path for entry in cv_os.scandir(series_dir)\n            if entry.is_file() and entry.name.lower().endswith(\".dcm\")\n        ]\n    except OSError:\n        return []\n\n\ndef cv_choose_slots(rows, study_id, image_root):\n    enriched = []\n    for row in rows:\n        series_id = str(row[\"SeriesInstanceUID\"])\n        path = image_root / study_id / series_id\n        enriched.append((row, series_id, path, len(cv_list_dicoms(path))))\n    chosen, used = [], set()\n    for _, plane, fluid in CV_SLOTS:\n        candidates = [\n            item for item in enriched\n            if str(item[0].get(\"Anatomical_Plane\")) == plane and item[1] not in used\n        ]\n        if fluid is not None:\n            preferred = [item for item in candidates if cv_fluid_value(item[0]) == fluid]\n            if preferred:\n                candidates = preferred\n        elif plane == \"Axial\":\n            preferred = [item for item in candidates if cv_fluid_value(item[0]) == 1]\n            if preferred:\n                candidates = preferred\n        selected = max(candidates, key=lambda item: (item[3], item[1])) if candidates else None\n        chosen.append(selected)\n        if selected is not None:\n            used.add(selected[1])\n    return chosen\n\n\ndef cv_ordered_series(series_dir):\n    records, spacings, lateralities, patient_x = [], [], [], []\n    for file_path in cv_list_dicoms(series_dir):\n        try:\n            header = cv_pydicom.dcmread(file_path, stop_before_pixels=True, force=True)\n            orientation = cv_np.asarray(getattr(header, \"ImageOrientationPatient\", []), dtype=cv_np.float64)\n            position = cv_np.asarray(getattr(header, \"ImagePositionPatient\", []), dtype=cv_np.float64)\n            if orientation.size == 6 and position.size == 3:\n                coordinate = float(cv_np.dot(position, cv_np.cross(orientation[:3], orientation[3:])))\n                patient_x.append(float(position[0]))\n            else:\n                coordinate = cv_safe_float(getattr(header, \"InstanceNumber\", 0), 0.0)\n            raw_spacing = getattr(header, \"PixelSpacing\", None)\n            spacing = cv_safe_float(raw_spacing[0], 0.5) if raw_spacing is not None else 0.5\n            laterality = str(getattr(header, \"Laterality\", \"\")).strip().upper()[:1]\n            if laterality in {\"L\", \"R\"}:\n                lateralities.append(laterality)\n            spacings.append(spacing)\n            records.append((coordinate, file_path, spacing))\n        except Exception:\n            records.append((0.0, file_path, 0.5))\n    records.sort(key=lambda item: (item[0], item[1]))\n    side = max(set(lateralities), key=lateralities.count) if lateralities else None\n    if side is None and patient_x:\n        median_x = float(cv_np.median(patient_x))\n        side = \"R\" if median_x < -20.0 else \"L\" if median_x > 20.0 else None\n    return records, (float(cv_np.median(spacings)) if spacings else 0.5), side\n\n\ndef cv_decode_pixels(file_path):\n    dicom = cv_pydicom.dcmread(file_path, force=True)\n    pixels = cv_apply_modality_lut(dicom.pixel_array, dicom).astype(cv_np.float32)\n    if str(getattr(dicom, \"PhotometricInterpretation\", \"\")) == \"MONOCHROME1\":\n        pixels = pixels.max() - pixels\n    return pixels\n\n\ndef cv_physical_crop_resize(pixels, spacing):\n    height, width = pixels.shape\n    crop = int(round(CV_CROP_MM / max(float(spacing), 1e-3)))\n    crop = min(max(crop, 16), height, width)\n    y0, x0 = (height - crop) // 2, (width - crop) // 2\n    pixels = pixels[y0:y0 + crop, x0:x0 + crop]\n    return cv_cv2.resize(pixels, (CV_MODEL_SIZE, CV_MODEL_SIZE), interpolation=cv_cv2.INTER_AREA)\n\n\ndef cv_normalize_laterality(pixels, plane, side):\n    if side != \"R\":\n        return pixels\n    return pixels[::-1] if plane == \"Sagittal\" else pixels[:, ::-1]\n\n\ndef cv_build_study(study_id, rows, image_root):\n    volume = cv_np.zeros(\n        (len(CV_SLOTS), CV_SLICES_PER_SLOT, CV_MODEL_SIZE, CV_MODEL_SIZE), dtype=cv_np.uint8\n    )\n    slice_mask = cv_np.zeros((len(CV_SLOTS), CV_SLICES_PER_SLOT), dtype=cv_np.uint8)\n    decode_failures = 0\n    selected_slots = cv_choose_slots(rows, study_id, image_root)\n    for slot_index, ((_, plane, _), selected) in enumerate(zip(CV_SLOTS, selected_slots)):\n        if selected is None or selected[3] == 0:\n            continue\n        records, median_spacing, side = cv_ordered_series(selected[2])\n        if not records:\n            continue\n        count = len(records)\n        low = int(count * CV_SLICE_LOW)\n        high = max(low, int(count * CV_SLICE_HIGH) - 1)\n        picks = cv_np.linspace(low, high, CV_SLICES_PER_SLOT).round().astype(int)\n        decoded, spacings = [], []\n        for pick in picks:\n            _, file_path, spacing = records[min(int(pick), count - 1)]\n            try:\n                decoded.append(cv_decode_pixels(file_path))\n            except Exception:\n                decoded.append(None)\n                decode_failures += 1\n            spacings.append(spacing if spacing > 0 else median_spacing)\n        valid = [pixels for pixels in decoded if pixels is not None]\n        if not valid:\n            continue\n        pooled = cv_np.concatenate([pixels.ravel() for pixels in valid])\n        lower, upper = cv_np.percentile(pooled, [2.0, 98.0])\n        for slice_index, (pixels, spacing) in enumerate(zip(decoded, spacings)):\n            if pixels is None:\n                continue\n            pixels = cv_np.clip((pixels - lower) / max(upper - lower, 1e-6), 0.0, 1.0)\n            pixels = cv_physical_crop_resize(pixels, spacing)\n            pixels = cv_normalize_laterality(pixels, plane, side)\n            volume[slot_index, slice_index] = (pixels * 255).astype(cv_np.uint8)\n            slice_mask[slot_index, slice_index] = 1\n    return volume, slice_mask, decode_failures\n\n\ndef cv_make_windows(volume, slice_mask, device):\n    windows, slots, positions, valid_windows = [], [], [], []\n    for slot in range(len(CV_SLOTS)):\n        for center in range(1, CV_SLICES_PER_SLOT - 1):\n            windows.append(volume[slot, center - 1:center + 2])\n            slots.append(slot)\n            positions.append(center / (CV_SLICES_PER_SLOT - 1))\n            valid_windows.append(slice_mask[slot, center - 1:center + 2].all())\n    windows = cv_torch.from_numpy(cv_np.stack(windows).astype(cv_np.float32) / 255.0)[None]\n    windows = (windows - CV_MEAN) / CV_STD\n    return (\n        windows.to(device, non_blocking=True),\n        cv_torch.tensor(slots, dtype=cv_torch.long, device=device)[None],\n        cv_torch.tensor(positions, dtype=cv_torch.float32, device=device)[None],\n        cv_torch.tensor(valid_windows, dtype=cv_torch.bool, device=device)[None],\n    )\n\n\nclass CVConvNeXtSpatialMIL(cv_nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = cv_timm.create_model(\n            CV_MODEL_NAME, pretrained=False, num_classes=0, global_pool=\"\"\n        )\n        dimension = int(self.backbone.num_features)\n        self.dimension = dimension\n        self.slot_embedding = cv_nn.Embedding(len(CV_SLOTS), dimension)\n        self.position_embedding = cv_nn.Sequential(\n            cv_nn.Linear(1, dimension), cv_nn.GELU(), cv_nn.Linear(dimension, dimension)\n        )\n        self.token_norm = cv_nn.LayerNorm(dimension)\n        self.queries = cv_nn.Parameter(cv_torch.randn(len(CV_TARGETS), dimension) * 0.02)\n        self.fusion = cv_nn.Sequential(\n            cv_nn.LayerNorm(4 * dimension), cv_nn.Linear(4 * dimension, dimension),\n            cv_nn.GELU(), cv_nn.Dropout(0.20),\n        )\n        self.classifier_weight = cv_nn.Parameter(cv_torch.randn(len(CV_TARGETS), dimension) * 0.02)\n        self.classifier_bias = cv_nn.Parameter(cv_torch.zeros(len(CV_TARGETS)))\n\n    def spatial_features(self, images):\n        features = self.backbone.forward_features(images)\n        if features.ndim != 4:\n            raise RuntimeError(f\"Expected ConvNeXt feature map, got {features.shape}\")\n        if features.shape[1] != self.dimension and features.shape[-1] == self.dimension:\n            features = features.permute(0, 3, 1, 2).contiguous()\n        if features.shape[1] != self.dimension:\n            raise RuntimeError(f\"ConvNeXt channel contract drift: {features.shape}\")\n        features = cv_F.adaptive_avg_pool2d(features, (CV_GRID_SIZE, CV_GRID_SIZE))\n        return features.flatten(2).transpose(1, 2)\n\n    def forward(self, windows, slots, positions, window_mask):\n        batch, n_windows = windows.shape[:2]\n        features = self.spatial_features(windows.flatten(0, 1))\n        n_spatial = features.shape[1]\n        features = features.view(batch, n_windows, n_spatial, self.dimension)\n        features = features + self.slot_embedding(slots)[:, :, None, :]\n        features = features + self.position_embedding(positions[..., None])[:, :, None, :]\n        tokens = self.token_norm(features.flatten(1, 2))\n        token_mask = window_mask[:, :, None].expand(-1, -1, n_spatial).flatten(1)\n        empty = ~token_mask.any(dim=1)\n        if empty.any():\n            token_mask = token_mask.clone()\n            token_mask[empty, 0] = True\n        scores = cv_torch.einsum(\"btd,nd->bnt\", tokens, self.queries) / cv_math.sqrt(self.dimension)\n        scores = scores.masked_fill(~token_mask[:, None, :], -1e4)\n        attention = cv_torch.softmax(scores, dim=-1)\n        attended = cv_torch.einsum(\"bnt,btd->bnd\", attention, tokens)\n        mask_float = token_mask.to(tokens.dtype).unsqueeze(-1)\n        mean = (tokens * mask_float).sum(1) / mask_float.sum(1).clamp_min(1.0)\n        maximum = tokens.masked_fill(~token_mask[..., None], -1e4).amax(1)\n        maximum = cv_torch.where(empty[:, None], cv_torch.zeros_like(maximum), maximum)\n        mean = mean[:, None, :].expand(-1, len(CV_TARGETS), -1)\n        maximum = maximum[:, None, :].expand(-1, len(CV_TARGETS), -1)\n        topk = min(CV_TOPK_TOKENS, tokens.shape[1])\n        top_indices = scores.topk(topk, dim=-1).indices\n        expanded = tokens[:, None, :, :].expand(-1, len(CV_TARGETS), -1, -1)\n        gathered = cv_torch.gather(\n            expanded, 2, top_indices[..., None].expand(-1, -1, -1, self.dimension)\n        )\n        top_features = gathered.mean(2)\n        fused = self.fusion(cv_torch.cat([attended, mean, maximum, top_features], dim=-1))\n        return (fused * self.classifier_weight[None]).sum(-1) + self.classifier_bias\n\n\ndef cv_validate_payload(payload, fold, fold_hash):\n    required = {\"schema_version\", \"model_state\", \"fold\", \"targets\", \"config\", \"cache_contract\", \"fold_assignment_sha256\"}\n    missing = required - set(payload)\n    if missing:\n        raise RuntimeError(f\"fold {fold}: missing checkpoint keys {sorted(missing)}\")\n    if int(payload[\"fold\"]) != fold or payload[\"targets\"] != CV_TARGETS:\n        raise RuntimeError(f\"fold {fold}: fold/target contract mismatch\")\n    expected_config = {\n        \"MODEL_NAME\": CV_MODEL_NAME, \"MODEL_SIZE\": CV_MODEL_SIZE,\n        \"GRID_SIZE\": CV_GRID_SIZE, \"CROP_MM\": CV_CROP_MM,\n        \"SLICE_LOW\": CV_SLICE_LOW, \"SLICE_HIGH\": CV_SLICE_HIGH,\n        \"SLICES_PER_SLOT\": CV_SLICES_PER_SLOT, \"TOPK_TOKENS\": CV_TOPK_TOKENS,\n    }\n    drift = {key: (payload[\"config\"].get(key), value) for key, value in expected_config.items() if payload[\"config\"].get(key) != value}\n    if drift:\n        raise RuntimeError(f\"fold {fold}: config drift {drift}\")\n    contract = payload[\"cache_contract\"]\n    expected_contract = {\n        \"slots\": CV_SLOTS,\n        \"crop_mm\": CV_CROP_MM,\n        \"slice_band\": [CV_SLICE_LOW, CV_SLICE_HIGH],\n        \"laterality\": \"R: sagittal vertical flip; coronal/axial horizontal flip\",\n        \"normalization\": \"per-series percentile 2-98 to uint8\",\n    }\n    for key, value in expected_contract.items():\n        if contract.get(key) != value:\n            raise RuntimeError(f\"fold {fold}: preprocessing drift for {key}: {contract.get(key)}\")\n    observed_hash = payload[\"fold_assignment_sha256\"]\n    if fold_hash is not None and observed_hash != fold_hash:\n        raise RuntimeError(f\"fold {fold}: fold-assignment hash mismatch\")\n    return observed_hash\n\n\ndef cv_run_inference():\n    device = cv_torch.device(\"cuda:0\" if cv_torch.cuda.is_available() else \"cpu\")\n    if device.type != \"cuda\":\n        raise RuntimeError(\"ConvNeXt five-fold inference requires a GPU runtime\")\n    cv_torch.backends.cudnn.benchmark = True\n    cv_torch.backends.cuda.matmul.allow_tf32 = True\n    root = cv_find_competition_root()\n    checkpoint_dir = cv_find_checkpoint_dir()\n    image_root = root / \"test_series\"\n    if not image_root.is_dir():\n        raise FileNotFoundError(f\"Missing test DICOM directory: {image_root}\")\n    test = cv_pd.read_csv(root / \"test.csv\", dtype={\"StudyInstanceUID\": str})\n    series = cv_pd.read_csv(\n        root / \"test_series.csv\",\n        dtype={\"StudyInstanceUID\": str, \"SeriesInstanceUID\": str},\n    )\n    if test.StudyInstanceUID.duplicated().any():\n        raise RuntimeError(\"Duplicate StudyInstanceUID in test.csv\")\n    series_by_study = {\n        str(study_id): frame.to_dict(\"records\")\n        for study_id, frame in series.groupby(\"StudyInstanceUID\", sort=False)\n    }\n\n    cv_gc.collect()\n    cv_torch.cuda.empty_cache()\n    models, fold_hash = [], None\n    for fold, filename in enumerate(CV_EXPECTED_FILES):\n        path = checkpoint_dir / filename\n        payload = cv_torch.load(path, map_location=\"cpu\", weights_only=False)\n        fold_hash = cv_validate_payload(payload, fold, fold_hash)\n        model = CVConvNeXtSpatialMIL()\n        incompatible = model.load_state_dict(payload[\"model_state\"], strict=True)\n        if incompatible.missing_keys or incompatible.unexpected_keys:\n            raise RuntimeError(f\"fold {fold}: state_dict incompatibility {incompatible}\")\n        model.eval().to(device)\n        models.append(model)\n        del payload\n        cv_gc.collect()\n        print(f\"Loaded ConvNeXt fold {fold}: {path.name}\")\n\n    predictions = cv_np.zeros((len(test), len(CV_TARGETS)), dtype=cv_np.float32)\n    total_decode_failures = 0\n    for row_index, study_id in enumerate(test.StudyInstanceUID.astype(str)):\n        volume, slice_mask, failures = cv_build_study(\n            study_id, series_by_study.get(study_id, []), image_root\n        )\n        total_decode_failures += failures\n        windows, slots, positions, window_mask = cv_make_windows(volume, slice_mask, device)\n        if not bool(window_mask.any().item()):\n            raise RuntimeError(f\"{study_id}: no valid ConvNeXt window\")\n        fold_predictions = []\n        with cv_torch.inference_mode(), cv_torch.autocast(device_type=\"cuda\", dtype=cv_torch.float16):\n            for model in models:\n                fold_predictions.append(cv_torch.sigmoid(model(windows, slots, positions, window_mask)).float())\n        predictions[row_index] = cv_torch.stack(fold_predictions).mean(0)[0].cpu().numpy()\n        if (row_index + 1) % 25 == 0 or row_index + 1 == len(test):\n            print(f\"ConvNeXt inference {row_index + 1}/{len(test)}; decode failures={total_decode_failures}\")\n        del volume, slice_mask, windows, slots, positions, window_mask, fold_predictions\n\n    if not cv_np.isfinite(predictions).all():\n        raise RuntimeError(\"ConvNeXt produced non-finite predictions\")\n    raw = cv_pd.DataFrame(predictions, columns=CV_TARGETS)\n    raw.insert(0, \"StudyInstanceUID\", test.StudyInstanceUID.astype(str).values)\n    ranked = raw.copy()\n    ranked[CV_TARGETS] = raw[CV_TARGETS].rank(method=\"average\", pct=True)\n    raw_path = CVPath(\"/kaggle/working/pred_convnext_spatial_mil_raw.csv\")\n    rank_path = CVPath(\"/kaggle/working/pred_convnext_spatial_mil_rank.csv\")\n    raw.to_csv(raw_path, index=False)\n    ranked.to_csv(rank_path, index=False)\n    print(f\"Saved {raw_path} and {rank_path}; checkpoint_dir={checkpoint_dir}\")\n    return raw, ranked\n\n\nconvnext_raw, convnext_rank = cv_run_inference()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T13:47:21.737963Z","iopub.execute_input":"2026-08-30T13:47:21.738254Z","iopub.status.idle":"2026-08-30T13:53:36.702793Z","shell.execute_reply.started":"2026-08-30T13:47:21.73823Z","shell.execute_reply":"2026-08-30T13:53:36.701985Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nfrom itertools import combinations\n\nTARGETS = [\n      \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n      \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\",\n      \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\"\n  ]\n\npaths = {\n      \"dino_rad\": \"/kaggle/working/submission.csv\",\n      \"coatnet_v4\": \"/kaggle/working/coatnet_submission.csv\",\n      \"maxspan_v5\": \"/kaggle/working/pred_coatnet_maxspan_raw.csv\",\n      \"convnext_spatial_mil\": \"/kaggle/working/pred_convnext_spatial_mil_raw.csv\",\n  }\n\nbranches = {\n      name: pd.read_csv(path, dtype={\"StudyInstanceUID\": str})\n                 .set_index(\"StudyInstanceUID\")\n      for name, path in paths.items()\n  }\n\n# Căn chỉnh study và kiểm tra toàn vẹn.\nreference_ids = branches[\"dino_rad\"].index\nfor name, frame in branches.items():\n      if set(frame.index) != set(reference_ids):\n          raise RuntimeError(f\"{name}: tập StudyInstanceUID không khớp\")\n      branches[name] = frame.loc[reference_ids]\n\n      values = branches[name][TARGETS].to_numpy(float)\n      if not np.isfinite(values).all():\n          raise RuntimeError(f\"{name}: có prediction không hữu hạn\")\n\n# Spearman từng target cho ba cặp model.\nrows = []\nfor target in TARGETS:\n      row = {\"target\": target}\n\n      for left, right in combinations(branches, 2):\n          row[f\"{left}__{right}\"] = branches[left][target].corr(\n              branches[right][target],\n              method=\"spearman\",\n          )\n\n      rows.append(row)\n\ncorrelations = pd.DataFrame(rows)\n\n# So sánh trực tiếp ConvNeXt với baseline 0.934 hiện tại (DINO/Rad + MaxSpan 50/50).\nbaseline_rank = 0.50 * branches[\"dino_rad\"][TARGETS].rank(method=\"average\", pct=True)\nbaseline_rank += 0.50 * branches[\"maxspan_v5\"][TARGETS].rank(method=\"average\", pct=True)\ncorrelations[\"convnext_spatial_mil__baseline_50_50\"] = [\n      branches[\"convnext_spatial_mil\"][target].corr(baseline_rank[target], method=\"spearman\")\n      for target in TARGETS\n  ]\nprint(correlations.to_string(index=False))\n\nprint(\"\\nMean Spearman:\")\nprint(\n      correlations.drop(columns=\"target\")\n                  .mean()\n                  .sort_values()\n                  .to_string()\n  )\ncorrelations.to_csv(\"/kaggle/working/convnext_branch_spearman.csv\", index=False)","metadata":{"papermill":{"duration":0.060259,"end_time":"2026-08-26T08:51:28.263836+00:00","exception":false,"start_time":"2026-08-26T08:51:28.203577+00:00","status":"completed"},"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T13:53:36.703958Z","iopub.execute_input":"2026-08-30T13:53:36.704294Z","iopub.status.idle":"2026-08-30T13:53:36.795103Z","shell.execute_reply.started":"2026-08-30T13:53:36.70427Z","shell.execute_reply":"2026-08-30T13:53:36.794274Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ranks = {\n      name: frame[TARGETS].rank(method=\"average\", pct=True)\n      for name, frame in branches.items()\n}\n\nCONVNEXT_WEIGHT = 0.05\nbaseline = 0.50 * ranks[\"dino_rad\"] + 0.50 * ranks[\"maxspan_v5\"]\nensemble = (1.0 - CONVNEXT_WEIGHT) * baseline + CONVNEXT_WEIGHT * ranks[\"convnext_spatial_mil\"]\n\nbaseline_submission = baseline.reset_index()\nsubmission = ensemble.reset_index()\nexpected_columns = [\"StudyInstanceUID\", *TARGETS]\nfor name, frame in {\"baseline\": baseline_submission, \"candidate\": submission}.items():\n      if frame.columns.tolist() != expected_columns:\n          raise RuntimeError(f\"{name}: schema submission sai: {frame.columns.tolist()}\")\n      values = frame[TARGETS].to_numpy(float)\n      if not np.isfinite(values).all() or not ((0.0 <= values) & (values <= 1.0)).all():\n          raise RuntimeError(f\"{name}: prediction không hữu hạn hoặc ngoài [0, 1]\")\n\nbaseline_submission.to_csv(\"/kaggle/working/submission_baseline_50_50.csv\", index=False)\nsubmission.to_csv(\"/kaggle/working/submission_convnext_005.csv\", index=False)\nsubmission.to_csv(\"/kaggle/working/submission.csv\", index=False)\nprint(\"Final weights: DINO/Rad=0.475, MaxSpan=0.475, ConvNeXt=0.050\")\nprint(\"Saved baseline, ConvNeXt candidate, and submission.csv\")\n\nsubmission.head()","metadata":{"papermill":{"duration":0.036871,"end_time":"2026-08-26T08:51:28.310836+00:00","exception":false,"start_time":"2026-08-26T08:51:28.273965+00:00","status":"completed"},"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T13:53:36.796221Z","iopub.execute_input":"2026-08-30T13:53:36.796555Z","iopub.status.idle":"2026-08-30T13:53:36.829572Z","shell.execute_reply.started":"2026-08-30T13:53:36.796532Z","shell.execute_reply":"2026-08-30T13:53:36.828994Z"}},"outputs":[],"execution_count":null}]}