{"cells":[{"cell_type":"markdown","id":"ensemble-description","metadata":{},"source":"# RSNA Knee DINO-RadImageNet 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\nThe notebook writes one competition artifact: `submission.csv`. Its inputs are the competition data, one consolidated checkpoint dataset, and the external DINOv2-small model.\n"},{"cell_type":"code","id":"dinov2-inference","execution_count":null,"metadata":{},"outputs":[],"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()\n"},{"cell_type":"code","id":"dinov3-inference","execution_count":null,"metadata":{},"outputs":[],"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)\n"},{"cell_type":"code","id":"rad-inference","execution_count":null,"metadata":{},"outputs":[],"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()\n"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"import subprocess as _master_subprocess\nimport sys as _master_sys\nimport os\nimport json\nimport time\nimport pandas as pd\nimport numpy as np\nimport torch\nfrom pathlib import Path\n\nMASTER_FINAL_MODE = os.environ.get('MASTER_FINAL_MODE', 'master').lower()\nMASTER_LEGACY_ALPHA = float(os.environ.get('MASTER_LEGACY_ALPHA', '0.06'))\nMASTER_B3_ALPHA = float(os.environ.get('MASTER_B3_ALPHA', '0.05'))\nMASTER_FORCE_B3 = os.environ.get('MASTER_FORCE_B3', '0') == '1'\nMASTER_B3_TARGET_ALPHAS = {\n    'ACL': 0.00, 'MCL': 0.10, 'Medial Meniscus': 0.00,\n    'Lateral Meniscus': 0.35, 'Medial OA': 0.15, 'Lateral OA': 0.35,\n    'PF OA': 0.35, 'Effusion': 0.25, 'Synovitis': 0.35,\n    \"Baker's\": 0.35, 'Contusion': 0.00, 'Fracture': 0.00,\n}\n\ndef _master_validate(frame, test_df, tag):\n    frame = frame.copy()\n    expected = ['StudyInstanceUID'] + TARGETS\n    if frame.columns.tolist() != expected:\n        raise RuntimeError(f'{tag}: schema mismatch')\n    frame['StudyInstanceUID'] = frame['StudyInstanceUID'].astype(str)\n    test = test_df[['StudyInstanceUID']].copy()\n    test['StudyInstanceUID'] = test['StudyInstanceUID'].astype(str)\n    if frame['StudyInstanceUID'].duplicated().any():\n        raise RuntimeError(f'{tag}: duplicate StudyInstanceUID')\n    if set(frame['StudyInstanceUID']) != set(test['StudyInstanceUID']):\n        raise RuntimeError(f'{tag}: StudyInstanceUID set mismatch')\n    frame = test.merge(frame, on='StudyInstanceUID', how='left')\n    arr = frame[TARGETS].to_numpy(np.float64)\n    if not np.isfinite(arr).all():\n        raise RuntimeError(f'{tag}: non-finite values')\n    return frame\n\ndef _master_find_b3_package():\n    explicit = os.environ.get('KNEE_B3_DIR', '').strip()\n    candidates = [Path(explicit)] if explicit else []\n    candidates.append(Path('/kaggle/input/rsna-knee-b3-v47-folds-0-3'))\n    base = Path('/kaggle/input')\n    if base.is_dir():\n        for p in base.iterdir():\n            if p.is_dir():\n                if (p / 'source/efficientnet_b3_public_repro_v1_infer.py').is_file():\n                    candidates.append(p)\n                for sub_p in p.iterdir():\n                    if sub_p.is_dir() and (sub_p / 'source/efficientnet_b3_public_repro_v1_infer.py').is_file():\n                        candidates.append(sub_p)\n    seen = set()\n    for root in candidates:\n        key = str(root)\n        if key in seen:\n            continue\n        seen.add(key)\n        infer_py = root / 'source/efficientnet_b3_public_repro_v1_infer.py'\n        module_py = root / 'source/efficientnet_b3_public_repro_v4_t4.py'\n        folds = [root / f'fold{i}/fold{i}_final.pt' for i in range(5)]\n        if infer_py.is_file() and module_py.is_file() and all(p.is_file() for p in folds):\n            return root\n    return None\n\ndef _master_b3_audit(root):\n    p = root / 'audit/audit.json'\n    if not p.is_file():\n        return False, 'audit/audit.json absent'\n    try:\n        a = json.loads(p.read_text())\n        nested = float(a['selection']['global_nested_macro_auc'])\n        base = float(a['arms']['exact_public_macro_auc'])\n        return nested > base, f'nested OOF {nested:.5f} vs reference {base:.5f}'\n    except Exception as exc:\n        return False, f'audit parse failed: {type(exc).__name__}: {exc}'\n\ndef _master_run_b3_raw(test_df):\n    root = _master_find_b3_package()\n    if root is None:\n        print('master B3: full five-fold package not found')\n        return None, False\n    supports, msg = _master_b3_audit(root)\n    print(f'master B3 package: {root} | {msg}')\n    if not torch.cuda.is_available():\n        print('master B3 skipped: CUDA unavailable')\n        return None, supports\n    left = TIME_BUDGET - (time.time() - T0)\n    if left < 15 * 60:\n        print(f'master B3 skipped: only {left/60:.1f} min remain')\n        return None, supports\n    outdir = Path('/kaggle/working/rsna_b3_master_inference')\n    outdir.mkdir(parents=True, exist_ok=True)\n    infer_py = root / 'source/efficientnet_b3_public_repro_v1_infer.py'\n    module_py = root / 'source/efficientnet_b3_public_repro_v4_t4.py'\n    folds = [root / f'fold{i}/fold{i}_final.pt' for i in range(5)]\n    budget_hours = min(1.75, max(0.25, 0.90 * left / 3600.0))\n    cmd = [\n        _master_sys.executable, str(infer_py),\n        '--module', str(module_py),\n        '--test-csv', str(ROOT / 'test.csv'),\n        '--series-csv', str(ROOT / 'test_series.csv'),\n        '--image-root', str(ROOT / 'test_series'),\n        '--checkpoints', *map(str, folds),\n        '--output-dir', str(outdir),\n        '--budget-hours', f'{budget_hours:.6f}',\n        '--checkpoint-every', '10',\n    ]\n    print(f'master B3 inference budget: {budget_hours:.2f}h')\n    try:\n        res = _master_subprocess.run(cmd, check=False, timeout=max(60.0, min(left * 0.96, budget_hours * 3600 + 10 * 60)))\n    except Exception as exc:\n        print(f'master B3 failed: {type(exc).__name__}: {exc}')\n        return None, supports\n    p = outdir / 'submission.csv'\n    if res.returncode != 0 or not p.is_file():\n        print(f'master B3 unavailable after inference, exit={res.returncode}')\n        return None, supports\n    b3 = _master_validate(pd.read_csv(p, dtype={'StudyInstanceUID': str}), test_df, 'B3 raw')\n    b3.to_csv('/kaggle/working/submission_b3_raw.csv', index=False)\n    return b3, supports\n\n_test_df_master = pd.read_csv(ROOT / 'test.csv', dtype={'StudyInstanceUID': str})\n_parent = _master_validate(pd.read_csv('/kaggle/working/submission.csv', dtype={'StudyInstanceUID': str}), _test_df_master, 'target parent')\n_parent.to_csv('/kaggle/working/submission_target920_recipe.csv', index=False)\n\n_b3, _b3_audit_ok = _master_run_b3_raw(_test_df_master)\nif _b3 is not None:\n    _tr = _parent[TARGETS].rank(method='average', pct=True)\n    _br = _b3[TARGETS].rank(method='average', pct=True)\n    _b3_10 = _parent.copy()\n    _b3_10[TARGETS] = 0.90 * _tr + 0.10 * _br\n    _b3_10.to_csv('/kaggle/working/submission_target_plus_b3_10.csv', index=False)\n    _b3_target = _parent.copy()\n    for _t in TARGETS:\n        _a = float(MASTER_B3_TARGET_ALPHAS[_t])\n        _b3_target[_t] = (1.0 - _a) * _tr[_t] + _a * _br[_t]\n    _b3_target.to_csv('/kaggle/working/submission_target_plus_b3_targetwise.csv', index=False)\n\n_parent_rank = _parent[TARGETS].rank(method='average', pct=True).to_numpy(np.float64)\n_legacy_path = Path('/kaggle/working/submission_legacy_fold_blend.csv')\n_legacy = None\nif _legacy_path.is_file() and MASTER_LEGACY_ALPHA > 0:\n    _legacy = _master_validate(pd.read_csv(_legacy_path, dtype={'StudyInstanceUID': str}), _test_df_master, 'legacy fold blend')\n    _legacy_rank = _legacy[TARGETS].rank(method='average', pct=True).to_numpy(np.float64)\nelse:\n    _legacy_rank = None\n\n_legacy_alpha = MASTER_LEGACY_ALPHA if _legacy_rank is not None else 0.0\n_b3_alpha = MASTER_B3_ALPHA if (_b3 is not None and (_b3_audit_ok or MASTER_FORCE_B3)) else 0.0\nif _legacy_alpha < 0 or _b3_alpha < 0 or _legacy_alpha + _b3_alpha >= 0.5:\n    raise ValueError('master blend weights are unsafe; require nonnegative diversity weights summing to < 0.5')\n_parent_alpha = 1.0 - _legacy_alpha - _b3_alpha\n_master_arr = _parent_alpha * _parent_rank\nif _legacy_rank is not None:\n    _master_arr += _legacy_alpha * _legacy_rank\nif _b3_alpha > 0:\n    _b3_rank = _b3[TARGETS].rank(method='average', pct=True).to_numpy(np.float64)\n    _master_arr += _b3_alpha * _b3_rank\n\n_master_sub = _parent[['StudyInstanceUID']].copy()\n_master_sub[TARGETS] = _master_arr\n_master_sub = _master_validate(_master_sub, _test_df_master, 'master')\n_master_sub.to_csv('/kaggle/working/submission_master.csv', index=False)\n\n_diag=[]\n_b3_rank_diag = _b3[TARGETS].rank(method='average', pct=True).to_numpy(np.float64) if _b3 is not None else None\nfor _j, _t in enumerate(TARGETS):\n    row={'target':_t}\n    if _legacy is not None:\n        row['target_vs_legacy_spearman'] = float(pd.Series(_parent_rank[:, _j]).corr(pd.Series(_legacy_rank[:, _j]), method='spearman'))\n    if _b3_rank_diag is not None:\n        row['target_vs_b3_spearman'] = float(pd.Series(_parent_rank[:, _j]).corr(pd.Series(_b3_rank_diag[:, _j]), method='spearman'))\n    _diag.append(row)\npd.DataFrame(_diag).to_csv('/kaggle/working/master_rank_correlation.csv', index=False)\n\nif MASTER_FINAL_MODE == 'target':\n    _parent.to_csv('/kaggle/working/submission.csv', index=False)\n    print('MASTER_FINAL_MODE=target -> exact target parent restored as submission.csv')\nelif MASTER_FINAL_MODE == 'master':\n    _master_sub.to_csv('/kaggle/working/submission.csv', index=False)\n    print(f'MASTER final = {_parent_alpha:.3f} target + {_legacy_alpha:.3f} legacy-fold + {_b3_alpha:.3f} B3 (rank space)')\nelse:\n    raise ValueError(\"MASTER_FINAL_MODE must be 'master' or 'target'\")\n\ndisplay(pd.DataFrame(_diag).head(12))\ndisplay(pd.read_csv('/kaggle/working/submission.csv').head())\n\n"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"from pathlib import Path\n\nwork_dir = Path(\"/kaggle/working\")\nmaster = work_dir / \"submission_master.csv\"\nfinal = work_dir / \"submission.csv\"\n\n# Make sure master exists\nif master.exists():\n    # Delete all other submission CSV files\n    for f in work_dir.glob(\"submission*.csv\"):\n        if f.name != \"submission_master.csv\":\n            f.unlink()\n            print(\"Deleted:\", f.name)\n\n    # Rename master -> submission.csv\n    master.rename(final)\n\nprint(\"\\n✅ FINAL FILE READY:\")\nprint(final)\nprint(\"\\nRemaining submission files:\")\nprint([f.name for f in work_dir.glob(\"submission*.csv\")])\n\n"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12"},"rsna_one_dataset_reproduction":{"source_notebook":"mattiaangeli/bend-the-knee-to-dinov3-the-original","source_script_version_id":342992625,"source_version_number":78,"source_file_sha256":"30f1f71b0498b39f0dffd64060d5c8033ed2d424f11ef2476b65d725e90fb08f","source_cells_sha256":"aefc642d72502d69c040a02f7c67f255dcef09083cf326ea36f9368acc6cc5dc","prediction_recipe_changed":false,"runtime_members_removed":5,"reason":"Inference-only replica with direct paths and the stable output-effective prediction path.","artifact_role":"documented inference notebook","diagnostic_outputs":[]},"kaggle":{"isGpuEnabled":true}},"nbformat":4,"nbformat_minor":5}