{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.13"},"rsna_one_dataset_reproduction":{"artifact_role":"documented inference notebook","diagnostic_outputs":[],"prediction_recipe_changed":false,"reason":"Inference-only replica with direct paths and the stable output-effective prediction path.","runtime_members_removed":5,"source_cells_sha256":"aefc642d72502d69c040a02f7c67f255dcef09083cf326ea36f9368acc6cc5dc","source_file_sha256":"30f1f71b0498b39f0dffd64060d5c8033ed2d424f11ef2476b65d725e90fb08f","source_notebook":"mattiaangeli/bend-the-knee-to-dinov3-the-original","source_script_version_id":342992625,"source_version_number":78},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":154281},{"sourceType":"datasetVersion","sourceId":19120128},{"sourceType":"datasetVersion","sourceId":19125270},{"sourceType":"datasetVersion","sourceId":19003959},{"sourceType":"modelInstanceVersion","sourceId":4533}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"dinosaurs":{"name":"RSNA Knee | DINOsaur V5.4 SpatialEvidence","default_submission":"submission.csv","anchor":"exact V4.5-style 0.936 fusion","strategy":"same DINOv2 + DINOv3 + RadImageNet + dual-Raptor inference; harvest max/top2 per-window Raptor evidence and learned slot-attention agreement with zero extra backbone passes","runtime_change":"head-only evidence statistics; no additional backbone inference passes","variants":["submission_v45_anchor.csv","submission_v54_safe.csv","submission_v54_main.csv","submission_v54_probe.csv"]}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"f2ada68e-30bc-4e78-b0cf-ab2697791d51","cell_type":"markdown","source":"# RSNA Knee | DINOsaur V4 🦖\n\nExact anchor + zero-extra-backbone Raptor MIL local-evidence residual.\nFocal max/top2 evidence is used only when the two Raptor checkpoints agree spatially.\n","metadata":{}},{"id":"85411088-1aff-417c-a229-71e2e9b86340","cell_type":"code","source":"from __future__ import annotations\nimport os as _os\ndef _comp_root():\n    # An API-attached competition mounts at /kaggle/input/competitions/<slug>/; only a UI-added one\n    # uses the short /kaggle/input/<slug>/ path. This notebook already resolves BOTH for the test\n    # root but hardcodes ROOT/COMP to the short form, so an API-pushed fork dies at cell 1 with\n    # FileNotFoundError on train.csv. Resolve it the same way instead of assuming.\n    for _c in (\"/kaggle/input/competitions/rsna-knee-abnormality-detection\",\n               \"/kaggle/input/rsna-knee-abnormality-detection\"):\n        if _os.path.isdir(_c):\n            return _c\n    raise RuntimeError(\"competition data not found under /kaggle/input\")\n_COMP_ROOT = _comp_root()\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(_COMP_ROOT)\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","metadata":{},"outputs":[],"execution_count":null},{"id":"b7782b99-2393-4aae-8810-73bac7f9f79e","cell_type":"code","source":"_A5_SAVED = dict(globals())\nimport gc, os, time, warnings\nfrom concurrent.futures import ProcessPoolExecutor, as_completed\nfrom pathlib import Path\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nwarnings.filterwarnings('ignore')\ncv2.setNumThreads(1)\nCROP_MM = 130.0\nSIZE = 336\nSLICE_BAND = (0.12, 0.88)\nN_SLICE = 16\nINTENSITY = 'slice'\nSLOTS = [('Sagittal', 1), ('Sagittal', 0), ('Coronal', 1), ('Coronal', 0), ('Axial', 1), ('Axial', 0)]\nN_SLOT = len(SLOTS)\nLABELS = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture']\nCOMP = Path(_COMP_ROOT)\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","metadata":{},"outputs":[],"execution_count":null},{"id":"51c55d1f-d956-4edd-acdd-d0e006132270","cell_type":"code","source":"from __future__ import annotations\nimport contextlib as _rad_contextlib\nimport base64 as _rad_b64\nimport zlib as _rad_zlib\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_CAL_PAYLOAD = 'eNrtmk1vI8cRhv9KsJdcKKE/q6tzc4z4ZCMBcjQWhrCRDSG2ZEjaIEGQ/57n7RlRQ3KG4jqLJAcDS4o709NdXR9vvVU9/3z30+3N/bvffRuuawgxlm7eq2ePeffrpV8v/V9euvLraD3k6ilW6zn126vYd+U6eKmx9RiaF7OSx+X1weE67K7SdSgpVSauuaecUxr3rtp1LK0HyyV7y9Gmy/E6pBhTL61Zt2gWx+WtSew6JmOoM5AHutXp+ro8+ZrlsnvJOfDdfRL+Kl4zd2bqllNmga7rvrva2OzGNFuynGxpmnxrS/069hCtpt6T1dLLOQVsTbJ1fWunG2qXAX8NiU+7VK8rdqvNU89uwQwjpWyeLdTCnVy84ua9pthSjznHXLjDpZxbLe4WPRAQCT+LrSVcGGdLzNX0XKhesC2eV5uZ67lQPDlmSzUhS2+61ELtoTOyWGjdPudc73fvnj7c/Hg7ElpyzYXjdwIZ32m7X3INZ0crtvtc833ua6fyoUYUGf4H13rJpX+GK5aa5VbeWO9ye/03bLi97i8SpaSCl49nc4gdxI2p1dZyVmQnfL56jBH/j70GqSq6dyMROLHQGvCpcSnkYsyeohNErnVjbamVjs47wOpBa7BAsa61SXulDfSIwAbAkQqDmSK38WwImKK1kFJ3d11CpFqR1mPHlj1pWVAHeDfyU2hk0vGogz0GqCBbKmFMx9xIWEAbQ8I4PUvmAplQCeuij7GNXGMk3yhBsFIfEnvVE6HFYMHDtFuLzKwpya9gisaZduAkcyAvmw1ROjmd7OnBYytpOBqJjTyK4KzT0Pi0tQJYslc0nGqcNZV7dEaRvfcSewg9h4p41iwO+4CWMkNG1326VBmHFUoB7VoakzE+JAtYN/NknS9lLJ+9S5qR56Q2kLA4qgTuplGNmUJAkbXiGbrU0HCIsgtYaGOUVXbayNeykLfpGv8lulBfgiFEmx61xm3LTlrOPoa51Ng7IyJiT4tWxrE0ZmnJRnhawZsyLgjXIC9MRu0tRAgImkJ7bUyHKk1eUgIewtRjXGDJqO1mNJrHfNUty1HRW2SSoV3FAVIA//hrUXKANxhbRbEkgqAFKiCC43qOEXsPTaKtKCdkqx3/koMYsuP+mh4HrHoQtqSIw/h4DM7EpRL5zRirAcHiUDj2SUlr9fFAsElHBX/zWgJCQlw+83Qksw8Pt9+Ty0hmOBQBh9XQAr7fxjRQ2IkGTT/w8cAiBKwRc1jwcMh+WKtMgFShNaDGJWRAechNSOMUldXTynNIUPAaEKKkCv1cLFzxaowBOmX8vU8Pj1voWa5JTPFcGN4QXm+PpZPjA8QQYYp9Qab9rZiIAuLX8AA8f/jX4kHgI+BXCtBEeOTFvNPapIyC92YAEOmHypJcCSiFOpQi712M73IL8KiBX8Hz62pDNsAH8MK/rMR+tNSRIRJwAkIIY8Rn/fD2kB2V9I7BMX+KSJ/XJtNAbwFglmQZbu/90CZcMhAeZKudKAvlSBLwNwP1aA88BYpPdJRkbADGeBzWN2IOvySVELxyU0Lx0JHQHsFpEQRsCB9iO1qT3IOHGUjMDqeczX4BiIbvCjkiLD6t6G239p+gAML1KQuAdaLLdidDV6Y6uPR+9+3LT8KrlYgzAsjg7ok+VJakimvgAjn5qUFmtTGJKXGxSadg2UuLkWok4EvGJ0Hdsr+DmpVlGsUMZkxtuTI8vDAVaIsnSIDbK0DhVSpAFlu4iRtgrLYWRaArADOcFDAfErEdwBS8dWpNqHupW3o7nIukTmxlmASq4L/51ESrznqJTSkgCzslVSMncekb8659ljeyYttxK0oX9k50n3h2kDqoapToWt8S9/yO3tTXyteLu+19rkPViKkIblLnzJhJUhYENUJCtdfWKkEFe9WTsWin6azpqJFIfEwNyLNes+PZXFV3h+tIOXjb5mzIRnCLaWKuhnPtjsCVPAOwdkhQhBSMfWLVruTgwEM/WXt1FYgv5AN4IsyBiHrOSscCBjIomksgZCPFn9grqDrE9UXDQCSPfs5RXyeuQl3gwcEIdn5JzADhpG34s4uF5NQ24eh1GWISkEmYA6iFiNezYAbkKNGJhIk+l9XNXK3r7AgLT4YnRUuHicJS0NRLLG0KDrF1IbkaMzj2MoeKzAH/kTo9KsnWJe8gDZUuxs1iPi0CayABMdLIT4qQCbfENfCiAsujIsj9KMUQ1kXwXUBFZspHZm9Zj/dB/SrVy1qOJg9ApKAlmSQNpr6SjibnNgVoqZQCr8gllgsTiFAOCFSeDcYWuKOChZQXi5dPxouF5Oo6ogVTaxDwnVE8KfQFI6Zomwgb/oI4VG2Ae1MdsTQCwQuZp5ZROZPcdudWHrZRfwe4cPE1cv9L/FDuwRwGNYOttGllMRItS1K3sFDQDAwQMvgpfIzyIeQj2yTFOhomGLsBVPtAASGLqiXKnChASW/4ILUTGlftQlUVl8FDjifbKYIzLKXmco4anNPK6ZVl8MhVBs+D5YgTo8HXuIEDQDJF4wHEYL5PvEExn4PIt7x0JumDeUgjxHeFIZDt+6xN5mALCc6pMjTZ7oy8EomCUA09MTRvvuCLxyFC+sEWCGiqjE+HIAJ4g39Bi5sadSN9MJZaFWLbpubBNBvlmwrPAONJWS1DpaIgbyKuZSbKg7qZzgNKsqnwIX8RAOvINdgfRg0jPNJBWE+5xAQ9qDaAPh4nKInagjAL9otj26U+sAnbUERTF0QYjHV8j8OEA+mohtH9SLLuoUIh6uAWHI6yAw54uEsqL6zufMAAt3yW+6xSFLmKwEmYgIPlXHYbPAo3AvKovCk/0KR5m3fmKkpFz2W3ng58h7KzqVVHqVYUpfGlAJHvQ2dN7QKmXoQITggHUz9HGZhSty5smYHeSL4pFdvbXlXZcSJV7GCHKPpsePEmVAOm41WxbuW3Y+qkrhYhisbrSBqbGiGdAkBwd5gzNehxGUVhzu5JCKa2VAyL2pfYgYgUtcTwu3RKZzO2JublhfCDbO1Qqypu8SdTWO0501Z6iAAH68g21mECvjX8bZZ7yvpViqq9BlHCnhj7DFuiXm8qEwgrAkK9hpfbuU/nN/i4l7I/qQEnpR4qVHUxgDjblWuM2UZFhOpzp+apb8i4ua8gnyHrqPMHs839gK7gaaYDDtILqNTWOF8kxVYdMxlVDwE68/F8DR2CgpCSkkzkEvKY3x9GoRqjpn8ZLj62TyEoRVV1dxrV2GtNOMosPAjqSC6pm6yRKRJFM5wnixbhh6uLl9FdjElsUvX7WxR61T1gGohOkQaej27MxiR+DQjAndQ5SgL4Yb+VMKT2UncGK4se5hfeR7Jr6nygbTJcPGK7QSpXOoBJkuMXlFTYQX0aRsObYEQn22BNrDQdKwmZ1DDaZNcANxumdkGs3pZIJcaCcxBuPSlGziTf+UNZIgqmLMTWS1/XYNBxFowH5FchZb7uT5EQqervUK+Rc15pP55mBpKra6Wju5MCbZwiKOSS/HFuZyVRBFeeIoVrc35oNnVt1ARpKDaHF2BO1z3DLILOrOVZdhIG52ojyFxsVa1k6Bq+sOS7AA0upIqpiCt9Om8+1SsIoI4/ZYFKZq9tdwnhCzrVRpoCxpqSo+0ua1DNLCOMc3X19vGSZZ9znfEAMTUpymSGqYW/kpXU6FWV6+pah9OGLveaTkopOV1d3U9oWfAhMiyK46fRgPeL9LTSDAtUjj5cFN4D8S5zmaCjmaTjJ/Kvxb3MaFrHVI27QS2vRfcMCmeqcUXY+NX6652okyj19OEC3SaFrdWyeyJMxo7gPrANn17GjQ5RgtqvregNjzr1hYOOxATeos4b/IItGRxEFZC4aNruVL1ZbGx8ojpE5moLiGNavazB+baxPoXr/gK56zjI1VGbjtG8nA3Xg3SpFl2gtqzqsM/oGgb5Axj4w+XD1g48MBFrMAnR+TRxMci1+8h+uCC1e1nmC7XEWhf2zEdIw5SsCe3Tcc18XibsU/PKhJhd7a380s67RLEvQQWUw9KKjmps7qefVeub/IZQoLbK0qxnHRFP7qlOy8BURC7Fy2LDPk7N1F1VEm1lrZbSmSWVSoNk63DQXgvIohAjYcjXU7b/zI8u7RVvPiRChSqVjkiU6YQjrT1ImVAAN/WoKEp62V2aQGAdAuQKhSW5qouyIdQ4jNe5NS6I8v10hLIeqIIMqplDn71o0dYl4fEFkRZ4LmhFJNUSQqZCNlBT2fIWnKzJ9XWwG48bOzGZWEIffcVx/q23xAziBV83gf1cE4/+mvI8/EFV3QVmPW0a6rB7ZD314ploTllOSlFYqX3X4u7TJg6jmxLRWxit1FaPtArKoHfHeb3jwPmNGDqfvS+Iv9Ns14Vv6p9rfyft+HGiFqmJIQSU++rlbegPBqDumli9zoOGSptec8O6GAWqtaAgkOgqZ3P1WhQQ08aj2KGONuEIqdZetzuLFB6qQniid8HkZTnlyGsPGnA4lY5Qu1456WnbokG9IupwU0NoPhHTwRRUqynZU0HObV8QUyZuo/lb6tFhsdxb7bCoU9qWe8vLxBwj2laVzgS+bIZGMfee1Qaldl87gBZ5jjr3Y4rq2Xaf3LatohHwxzbqBKtv4d8bHnh1eUqfVRx07jc4zfzySrj8IGV9NMFX4mBKrq6FXc4Jz154/3737u7++fbxw+3Pz9N7eg0X0mHCeHfHbHrNhrCIpdXRA/LRRC4WR9sU/ZqJB46XJsJouwU1rPT+il6uKHqNAdbJ2InUJp0wVWBV7w8NuldU7uMyajZqh2N6UfeoM6ym1zLGizF6T6AnNQZiHO/KZMijZMOV2/yKQFKv0TSXq0Hb5teE9F4HLp50QKJ3OX64edZ7ie+++PLrd7t339z+5e7mx9/88Qt+f82dx5f//Omr6e8fvv/+49Pdwz0/f3/z19vH3z7x68uH++fpKhP+/Pjw/PDh4cfv+Hz86f5Jk99/93T7eHersfff/fnmh/H3y4fH8feLv9/x96ub5+l7vq9f0wj9msf8+HH6fhnDr3kMvzRGG3p8+PizVt3vaXy/ysjox5sPzx8fbxn+7cuWv7m9v3v68PFpsfH9pcWwLc1oyEI3f/7H/cPf7p7vnhZ6ev/+X/8GYIe3xg=='\n_RAD_CAL_W = 0.40\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\n\ndef _v18_cal_protocol(uids):\n    frame = _rad_pd.read_csv(\n        ROOT / 'test_series.csv',\n        dtype={\n            'StudyInstanceUID': str,\n            'SeriesInstanceUID': str,\n        },\n    )\n\n    frame['StudyInstanceUID'] = (\n        frame['StudyInstanceUID'].astype(str)\n    )\n\n    index = _rad_pd.Index(\n        [str(uid) for uid in uids],\n        name='StudyInstanceUID',\n    )\n\n    table = _rad_pd.DataFrame(index=index)\n\n    table['n_series'] = (\n        frame.groupby('StudyInstanceUID')\n        .size()\n        .reindex(index)\n        .fillna(0)\n    )\n\n    for plane in (\n        'Sagittal',\n        'Coronal',\n        'Axial',\n    ):\n        part = frame[\n            frame['Anatomical_Plane']\n            .astype(str)\n            .eq(plane)\n        ]\n\n        table[f'n_{plane[:3]}'] = (\n            part.groupby('StudyInstanceUID')\n            .size()\n            .reindex(index)\n            .fillna(0)\n        )\n\n    for flag in (\n        'Fat_Suppression',\n        'Fluid_Sensitive',\n    ):\n        marked = frame[\n            _rad_pd.to_numeric(\n                frame[flag],\n                errors='coerce',\n            )\n            .fillna(0)\n            > 0\n        ]\n\n        prefix = flag[:3]\n\n        table[prefix] = (\n            marked.groupby('StudyInstanceUID')\n            .size()\n            .reindex(index)\n            .fillna(0)\n        )\n\n        for plane in (\n            'Sagittal',\n            'Coronal',\n            'Axial',\n        ):\n            part = marked[\n                marked['Anatomical_Plane']\n                .astype(str)\n                .eq(plane)\n            ]\n\n            table[f'{prefix}_{plane[:3]}'] = (\n                part.groupby('StudyInstanceUID')\n                .size()\n                .reindex(index)\n                .fillna(0)\n            )\n\n    return table\n\n\ndef _v18_calibrate_transformer(\n    branch,\n    baseline_rank,\n    public_rank,\n    pass2_rank,\n    expected_ids,\n):\n    payload = _rad_json.loads(\n        _rad_zlib.decompress(\n            _rad_b64.b64decode(\n                _RAD_CAL_PAYLOAD\n            )\n        ).decode()\n    )\n\n    gate = set(payload['gate'])\n\n    protocol = _v18_cal_protocol(\n        expected_ids\n    )\n\n    if (\n        protocol.columns.tolist()\n        != list(\n            payload[\n                'protocol_columns'\n            ]\n        )\n    ):\n        raise RuntimeError(\n            'V18 calibration protocol '\n            'layout mismatch'\n        )\n\n    mean_rank = (\n        baseline_rank\n        + public_rank\n        + pass2_rank\n    ) / 3.0\n\n    blocks = [\n        baseline_rank,\n        public_rank,\n        pass2_rank,\n        public_rank - baseline_rank,\n        pass2_rank - baseline_rank,\n        mean_rank,\n    ]\n\n    for group in payload['groups']:\n        columns = [\n            _RAD_LABELS.index(target)\n            for target in group\n        ]\n\n        blocks.append(\n            mean_rank[\n                :,\n                columns,\n            ].mean(\n                axis=1,\n                keepdims=True,\n            )\n        )\n\n    blocks.append(\n        protocol.to_numpy(\n            _rad_np.float64\n        )\n    )\n\n    x = _rad_np.concatenate(\n        blocks,\n        axis=1,\n    )\n\n    centre = _rad_np.asarray(\n        payload['mean'],\n        _rad_np.float64,\n    )\n    spread = _rad_np.asarray(\n        payload['scale'],\n        _rad_np.float64,\n    )\n    coef = _rad_np.asarray(\n        payload['coef'],\n        _rad_np.float64,\n    )\n    bias = _rad_np.asarray(\n        payload['intercept'],\n        _rad_np.float64,\n    )\n\n    if (\n        x.shape[1] != 88\n        or coef.shape != (\n            len(_RAD_LABELS),\n            88,\n        )\n    ):\n        raise RuntimeError(\n            f'V18 calibration feature drift: '\n            f'x={x.shape}, coef={coef.shape}'\n        )\n\n    spread = _rad_np.where(\n        _rad_np.abs(spread) > 1e-8,\n        spread,\n        1.0,\n    )\n\n    adjusted = _rad_rank_columns(\n        (\n            (\n                x - centre\n            )\n            / spread\n        )\n        @ coef.T\n        + bias\n    )\n\n    output = branch.copy()\n\n    values = output[\n        _RAD_LABELS\n    ].to_numpy(\n        _rad_np.float64\n    ).copy()\n\n    for index, target in enumerate(\n        _RAD_LABELS\n    ):\n        if target in gate:\n            values[\n                :,\n                index,\n            ] = (\n                (\n                    1.0\n                    - _RAD_CAL_W\n                )\n                * values[\n                    :,\n                    index,\n                ]\n                + _RAD_CAL_W\n                * adjusted[\n                    :,\n                    index,\n                ]\n            )\n\n    output[\n        _RAD_LABELS\n    ] = _rad_rank_columns(\n        values\n    )\n\n    _rad_validate(\n        output,\n        expected_ids,\n    )\n\n    return output, gate\n\n\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] = (\n        (1.0 - _RAD_V48_SECOND_ALPHA)\n        * _rad_rank_columns(\n            e10[_RAD_LABELS].to_numpy()\n        )\n        + _RAD_V48_SECOND_ALPHA\n        * pass2_rank\n    )\n\n    final[_RAD_LABELS] = _rad_rank_columns(\n        final[_RAD_LABELS].to_numpy()\n    )\n\n    _rad_validate(\n        final,\n        expected_ids,\n    )\n\n    globals()['V18_TRANSFORMER_RAW'] = (\n        final.copy()\n    )\n    globals()['V18_CALIBRATOR_APPLIED'] = False\n    globals()['V18_CAL_GATE'] = tuple()\n\n    try:\n        calibrated, gate = (\n            _v18_calibrate_transformer(\n                final,\n                baseline_rank,\n                reference_rank,\n                pass2_rank,\n                expected_ids,\n            )\n        )\n\n        final = calibrated\n\n        globals()[\n            'V18_TRANSFORMER_CAL'\n        ] = final.copy()\n\n        globals()[\n            'V18_CALIBRATOR_APPLIED'\n        ] = True\n\n        globals()[\n            'V18_CAL_GATE'\n        ] = tuple(\n            sorted(gate)\n        )\n\n        print(\n            '[V18] 88-feature transformer '\n            'calibration applied to: '\n            + ', '.join(\n                sorted(gate)\n            ),\n            flush=True,\n        )\n\n    except Exception as exc:\n        print(\n            '[V18] calibration skipped '\n            'safely; raw transformer kept: '\n            f'{type(exc).__name__}: {exc}',\n            flush=True,\n        )\n\n    _rad_validate(\n        final,\n        expected_ids,\n    )\n\n    final.to_csv(\n        primary,\n        index=False,\n    )\n_rad_main()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"43ccc361-0b48-4d97-99dd-c723fd06620d","cell_type":"code","source":"# External Raptor checkpoints retain their original provenance and training recipe.\nimport os, sys, glob, time, json, gc\nfrom concurrent.futures import ThreadPoolExecutor\nos.environ.setdefault(\"HF_HUB_OFFLINE\", \"1\")\nos.environ.setdefault(\"TRANSFORMERS_OFFLINE\", \"1\")\nos.environ.setdefault(\"HF_HUB_DISABLE_TELEMETRY\", \"1\")\nimport numpy as np\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport timm\n# T4 (Turing) cuDNN v9 has fp16/fp32 conv engines but NOT bf16 for these shapes\n# (\"GET was unable to find an engine...\"); benchmark lets it pick a valid algo for\n# the fixed (1,24,3,res,res) input.\ntorch.backends.cudnn.benchmark = True\ntorch.backends.cuda.matmul.allow_tf32 = True\n\n# ---- fixed config (must match training exactly) -----------------------------\nIMG = 336\nCROP_MM = 140.0\n# 64 slices per study instead of 44, same proportions. Must match the corpus the weights\n# were trained on (knee_corpus_v4.py).\nSLOTS = [(\"Sagittal\", 1, 18), (\"Sagittal\", 0, 14), (\"Coronal\", 1, 12),\n         (\"Coronal\", 0, 8), (\"Axial\", -1, 12)]\nMAXS = sum(s[2] for s in SLOTS)                     # 64\nK_EVAL = 62   # every window position the volume holds, not an evenly spaced subset\nNORM = \"imagenet\"\nLAB = [\"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\", \"Medial OA\", \"Lateral OA\",\n       \"PF OA\", \"Effusion\", \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\"]\n_MEAN = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)\n_STD = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)\n\n# Three arms: (weights filename, fallback arch, fallback res). ck carries arch+res too.\n# Selected 2026-08-19 by greedy forward selection AND exhaustive subset search over a 7-arm\n# panel on the 45-study gold set (phase2/blend_panel.py); both agree on this exact set.\n# Singles: coatnet384 0.9025 | swinbase384 0.8825 | effv2l480 0.8716.\n# Blend {coatnet+swin+effv2l} = 0.9068 (2-arm {coatnet+swin} = 0.9059, coatnet alone 0.9025).\n# Dropped as redundant: cnn336 (0.8833, the former champion), cnbase384 (0.8754),\n# cnlarge384 (0.8752), maxvit384 (0.8438).\n#\n# SINGLE ARM: coatnet_rmlp_2_rw_384 retrained on the EXPANDED 4,349-study corpus.\n#\n# Why one arm and not the 3-arm blend: on the live leaderboard CoAtNet alone scored 0.914 while\n# every blend scored 0.914-0.915, so ensembling is worth ~+0.001 there -- the ~+0.010 it showed\n# on the old 45-study gold set was gold-set noise. One arm is also 1/3 the kernel runtime.\n#\n# Corpus expansion: the corpus previously held 3,200 of the 4,349 labelled studies and only 45\n# of the 58 gold studies. Rebuilt to 4,407 studies (+37.8% training data, 58-study gate).\n#\n# Measured on the 58-study gate (the incumbent re-scored on the SAME gate for a fair compare):\n#   incumbent CoAtNet (3,155-study corpus) 0.8923\n#   this model       (4,349-study corpus) 0.9054   (+0.0131, better in 92.7% of 2000 bootstraps)\n# Biggest gains land on the findings that were capping us: Lateral Meniscus +0.071,\n# Fracture +0.057, Lateral OA +0.048, Medial Meniscus +0.035, ACL +0.028.\nARMS = [\n    {\n        \"file\": \"raptor_ft_coatnet_v5_full_swa.pt\",\n        \"arch\": \"coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k\",\n        \"res\": 384,\n        \"w\": 1.0,\n    },\n]\n\nLEGACY_ARM = {\n    \"file\": \"raptor_ft_coatnet_v4_full.pt\",\n    \"arch\": \"coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k\",\n    \"res\": 384,\n    \"k_eval\": 42,\n    \"span_lo\": 0.06,\n    \"span_hi\": 0.94,\n}\n\nPRIMARY_SPAN_LO = 0.02\nPRIMARY_SPAN_HI = 0.98\n\n\n# ============================================================================\n# Model -- verbatim from finetune_raptor.py\n# ============================================================================\ndef build_backbone(arch, pretrained=False):\n    # maxvit/maxxvit/coatnet are conv-attention hybrids: NO CLS token, NO interpolatable\n    # pos-embed -> avg pool. The \"vit\" substring in \"coatnet\"/\"maxvit\" must NOT route them\n    # down the ViT path (mirrors finetune_raptor.py exactly).\n    hybrid = arch.startswith((\"maxvit\", \"maxxvit\", \"coatnet\", \"coat_\", \"convnext\"))\n    is_vit = (not hybrid) and any(k in arch for k in (\"vit\", \"deit\", \"dinov2\", \"eva\", \"beit\"))\n    kw = dict(pretrained=pretrained, num_classes=0, in_chans=3)\n    if is_vit:\n        kw.update(global_pool=\"token\", dynamic_img_size=True)\n    else:\n        kw.update(global_pool=\"avg\")\n    return timm.create_model(arch, **kw)\n\n\nclass RaptorClassifier(nn.Module):\n    def __init__(self, backbone, F_dim=768, n=12, drop=0.2):\n        super().__init__()\n        self.backbone = backbone\n        self.norm = nn.LayerNorm(F_dim)\n        self.att = nn.Sequential(nn.Linear(F_dim, 256), nn.Tanh(), nn.Dropout(drop),\n                                 nn.Linear(256, n))\n        self.clsW = nn.Parameter(torch.zeros(n, F_dim))\n        self.clsb = nn.Parameter(torch.zeros(n))\n        nn.init.trunc_normal_(self.clsW, std=0.02)\n        self.n = n\n\n    def encode(self, x):\n        B, K = x.shape[:2]\n        f = self.backbone(x.flatten(0, 1))\n        return f.view(B, K, -1)\n\n    def head(self, feats):\n        h = self.norm(feats)\n        a = self.att(h)\n        a = torch.softmax(a, dim=1)\n        pooled = torch.einsum(\"bkn,bkf->bnf\", a, h)\n        logits = (pooled * self.clsW).sum(-1) + self.clsb\n        return logits\n\n    def forward(self, x):\n        return self.head(self.encode(x))\n\n\ndef load_model(pt_path, arch_default, res_default, device, ngpu=1):\n    ck = torch.load(pt_path, map_location=\"cpu\", weights_only=False)\n    arch = ck.get(\"arch\", arch_default)\n    ck_res = int(ck.get(\"res\", res_default))\n    bb = build_backbone(arch, pretrained=False)\n    model = RaptorClassifier(bb, F_dim=bb.num_features)\n    model.load_state_dict(ck[\"model\"], strict=True)\n    model.eval().to(device)\n    # NOTE: DataParallel removed on purpose. On the full hidden test it drove a system-RAM OOM\n    # (per-forward module replication over many studies); a single T4 handles K_EVAL=24 windows\n    # fine. Arms are also run SEQUENTIALLY (see main) so peak RAM == one model, not two.\n    del ck\n    gc.collect()\n    return model, ck_res\n\n\n# ============================================================================\n# Eval windowing -- verbatim from finetune_raptor.py StudyWindows (train=False)\n# ============================================================================\ndef _eval_centers(mask, D, k):\n    valid = np.where(mask > 0)[0]\n    if len(valid) < 3:\n        valid = np.arange(min(3, D))\n    lo, hi = int(valid.min()), int(valid.max())\n    cs = [c for c in range(lo + 1, hi) if c - 1 >= lo and c + 1 <= hi]\n    if not cs:\n        cs = [max(1, min((lo + hi) // 2, D - 2))]\n    idx = np.linspace(0, len(cs) - 1, k).round().astype(int)\n    return [cs[i] for i in idx]\n\n\ndef eval_windows(vol, mask, k, res, norm=NORM):\n    D = vol.shape[0]\n    cs = _eval_centers(mask, D, k)\n    wins = np.empty((len(cs), 3, res, res), np.float32)\n    for j, c in enumerate(cs):\n        c = max(1, min(c, D - 2))\n        tri = np.stack([vol[c - 1], vol[c], vol[c + 1]], 0).astype(np.float32) / 255.0\n        t = torch.from_numpy(tri)\n        if t.shape[-1] != res:\n            t = F.interpolate(t[None], size=(res, res), mode=\"bilinear\",\n                              align_corners=False)[0]\n        wins[j] = t.numpy()\n    x = torch.from_numpy(wins)\n    if norm == \"imagenet\":\n        x = (x - _MEAN) / _STD\n    return x\n\n\n@torch.no_grad()\ndef infer_probs(model, xwins, device):\n    x = xwins.unsqueeze(0).to(\n        device,\n        non_blocking=True,\n    )\n\n    use_cuda = (\n        str(device).startswith(\"cuda\")\n    )\n\n    def _forward():\n        return torch.sigmoid(\n            model(x).float()\n        )[0].cpu().numpy()\n\n    if use_cuda:\n        try:\n            with torch.autocast(\n                \"cuda\",\n                dtype=torch.float16,\n            ):\n                return _forward()\n\n        except RuntimeError as error:\n            try:\n                with torch.cuda.device(device):\n                    torch.cuda.empty_cache()\n            except Exception:\n                pass\n\n            print(\n                \"[DINOsaur V5.4] \"\n                f\"{device} fp16 retry in fp32: \"\n                f\"{type(error).__name__}\",\n                flush=True,\n            )\n\n            return _forward()\n\n    return _forward()\n\n\n_V54_MAX_TARGETS = {\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\",\n}\n_V54_TOP2_TARGETS = {\n    \"ACL\",\n    \"MCL\",\n}\n_V54_FOCAL_TARGETS = _V54_MAX_TARGETS | _V54_TOP2_TARGETS\n\n\ndef _v54_slot_ids(centers):\n    edges = np.cumsum([int(slot[2]) for slot in SLOTS])\n    ids = np.searchsorted(edges, np.asarray(centers, dtype=np.int64), side=\"right\")\n    return np.clip(ids, 0, len(SLOTS) - 1).astype(np.int64)\n\n\n@torch.no_grad()\ndef infer_probs_with_evidence(model, xwins, centers, device):\n    \"\"\"Exact Raptor head prediction plus zero-extra-backbone focal evidence.\"\"\"\n    x = xwins.unsqueeze(0).to(\n        device,\n        non_blocking=True,\n    )\n    center_array = np.asarray(centers, dtype=np.int64)\n    if center_array.ndim != 1 or len(center_array) != int(xwins.shape[0]):\n        raise ValueError(\"Raptor evidence centers/windows mismatch\")\n    slot_ids = _v54_slot_ids(center_array)\n\n    use_cuda = str(device).startswith(\"cuda\")\n\n    def _forward():\n        feats = model.encode(x)\n        h = model.norm(feats)\n        att = torch.softmax(model.att(h), dim=1)\n        pooled = torch.einsum(\"bkn,bkf->bnf\", att, h)\n        logits = (pooled * model.clsW).sum(-1) + model.clsb\n        probs = torch.sigmoid(logits)\n\n        # Per-window target evidence before MIL attention pooling. This costs only\n        # head-level tensor ops because the backbone features are already computed.\n        local = torch.einsum(\"bkf,nf->bkn\", h, model.clsW)\n        local = local + model.clsb.view(1, 1, -1)\n        local0 = local[0].float()\n\n        k = int(local0.shape[0])\n        top2 = torch.topk(local0, k=min(2, k), dim=0).values.mean(dim=0)\n        local_max = local0.max(dim=0).values\n        local_pool = logits[0].float().clone()\n        for target_index, target in enumerate(LAB):\n            if target in _V54_MAX_TARGETS:\n                local_pool[target_index] = local_max[target_index]\n            elif target in _V54_TOP2_TARGETS:\n                local_pool[target_index] = top2[target_index]\n\n        a = att[0].float()  # K x targets\n        slot_mass = torch.zeros(\n            (len(LAB), len(SLOTS)),\n            device=a.device,\n            dtype=torch.float32,\n        )\n        for window_index, slot_index in enumerate(slot_ids.tolist()):\n            slot_mass[:, slot_index] += a[window_index]\n\n        peak = a.max(dim=0).values\n        entropy = -(\n            a.clamp_min(1e-7) * a.clamp_min(1e-7).log()\n        ).sum(dim=0)\n        entropy = entropy / max(float(np.log(max(k, 2))), 1e-6)\n\n        return (\n            probs[0].float().cpu().numpy(),\n            local_pool.cpu().numpy(),\n            slot_mass.cpu().numpy(),\n            peak.cpu().numpy(),\n            entropy.cpu().numpy(),\n        )\n\n    if use_cuda:\n        try:\n            with torch.autocast(\"cuda\", dtype=torch.float16):\n                return _forward()\n        except RuntimeError as error:\n            try:\n                with torch.cuda.device(device):\n                    torch.cuda.empty_cache()\n            except Exception:\n                pass\n            print(\n                \"[DINOsaur V5.4] \"\n                f\"{device} evidence fp16 retry in fp32: \"\n                f\"{type(error).__name__}\",\n                flush=True,\n            )\n            return _forward()\n\n    return _forward()\n\ndef rankpct(x):                                   # per-column percentile rank in [0,1]\n    order = x.argsort(0).argsort(0).astype(np.float64)\n    return order / max(1, (x.shape[0] - 1))\n\n\n# ============================================================================\n# Preprocessing -- verbatim from kprep2/dino_preprocess.py, retargeted to TEST\n# ============================================================================\ndef _make_reader():\n    import pydicom, cv2\n    from pydicom.pixel_data_handlers.util import apply_modality_lut\n\n    def order_and_meta(sdir):\n        fs = glob.glob(sdir + \"/*.dcm\"); recs = []; ps_list = []\n        for f in fs:\n            try:\n                h = pydicom.dcmread(f, stop_before_pixels=True)\n                iop = getattr(h, 'ImageOrientationPatient', None)\n                ipp = getattr(h, 'ImagePositionPatient', None)\n                if iop is not None and ipp is not None and len(iop) == 6:\n                    r = np.array(iop[:3], float); c = np.array(iop[3:], float)\n                    n = np.cross(r, c); pos = float(np.dot(np.array(ipp, float), n))\n                else:\n                    pos = float(getattr(h, 'InstanceNumber', 0) or 0)\n                ps = getattr(h, 'PixelSpacing', None); ps = float(ps[0]) if ps is not None else 0.5\n                ps_list.append(ps); recs.append((pos, f, ps))\n            except Exception:\n                recs.append((0.0, f, 0.5))\n        recs.sort(key=lambda x: x[0])\n        med_ps = float(np.median(ps_list)) if ps_list else 0.5\n        return [(f, ps) for _, f, ps in recs], med_ps\n\n    def read_px(f):\n        d = pydicom.dcmread(f)\n        a = apply_modality_lut(d.pixel_array, d).astype(np.float32)\n        if str(getattr(d, 'PhotometricInterpretation', '')) == 'MONOCHROME1':\n            a = a.max() - a\n        return a\n\n    def mm_crop_resize(a, ps):\n        h, w = a.shape; cpx = int(round(CROP_MM / max(ps, 1e-3)))\n        cpx = min(cpx, min(h, w)); y0 = (h - cpx) // 2; x0 = (w - cpx) // 2\n        a = a[y0:y0 + cpx, x0:x0 + cpx]\n        return cv2.resize(a, (IMG, IMG), interpolation=cv2.INTER_AREA)\n\n    return order_and_meta, read_px, mm_crop_resize\n\n\ndef _pick_series_for_slot(rows, plane, fluid, used):\n    cands = [r for r in rows if r['Anatomical_Plane'] == plane and r['SeriesInstanceUID'] not in used]\n    if fluid in (0, 1):\n        pref = [r for r in cands if int(r.get('Fluid_Sensitive', 0) or 0) == fluid]\n        if pref:\n            return pref[0]\n    return cands[0] if cands else None\n\n\ndef _fill_variant_volume(\n    target_volume,\n    offset,\n    picks,\n    pixel_cache,\n    files,\n    med_ps,\n    mm_crop_resize,\n):\n    arrays = []\n    spacings = []\n\n    for position in picks:\n        position = min(\n            int(position),\n            len(files) - 1,\n        )\n        file_path, spacing = files[\n            position\n        ]\n        arrays.append(\n            pixel_cache.get(\n                position\n            )\n        )\n        spacings.append(\n            spacing\n            if spacing > 0\n            else med_ps\n        )\n\n    valid = [\n        array\n        for array in arrays\n        if array is not None\n    ]\n\n    if valid:\n        all_pixels = np.concatenate(\n            [\n                array.ravel()\n                for array in valid\n            ]\n        )\n        low, high = np.percentile(\n            all_pixels,\n            [2.0, 98.0],\n        )\n    else:\n        low, high = 0.0, 1.0\n\n    for local_index, (\n        array,\n        spacing,\n    ) in enumerate(\n        zip(\n            arrays,\n            spacings,\n        )\n    ):\n        output_index = (\n            offset\n            + local_index\n        )\n\n        if (\n            output_index\n            >= MAXS\n        ):\n            break\n\n        if array is None:\n            continue\n\n        normalized = np.clip(\n            (\n                array - low\n            )\n            / (\n                high - low\n                + 1e-6\n            ),\n            0,\n            1,\n        )\n\n        normalized = mm_crop_resize(\n            normalized,\n            spacing,\n        )\n\n        target_volume[\n            output_index\n        ] = (\n            normalized\n            * 255\n        ).astype(\n            np.uint8\n        )\n\n\ndef build_study_pair(\n    sid,\n    ser_records,\n    tsdir,\n    reader,\n):\n    \"\"\"\n    Produce exact MaxSpan and legacy-span volumes while reading every required\n    DICOM only once. Both checkpoints keep their own percentile normalization.\n    \"\"\"\n    (\n        order_and_meta,\n        read_px,\n        mm_crop_resize,\n    ) = reader\n\n    rows = ser_records.get(\n        sid,\n        [],\n    )\n\n    primary_volume = np.zeros(\n        (\n            MAXS,\n            IMG,\n            IMG,\n        ),\n        np.uint8,\n    )\n    legacy_volume = np.zeros_like(\n        primary_volume\n    )\n\n    used = set()\n    offset = 0\n\n    for plane, fluid, count in SLOTS:\n        record = _pick_series_for_slot(\n            rows,\n            plane,\n            fluid,\n            used,\n        )\n\n        if record is None:\n            offset += count\n            continue\n\n        used.add(\n            record[\n                \"SeriesInstanceUID\"\n            ]\n        )\n\n        files, med_ps = order_and_meta(\n            f\"{tsdir}/{sid}/\"\n            f\"{record['SeriesInstanceUID']}\"\n        )\n\n        if not files:\n            offset += count\n            continue\n\n        number = len(files)\n\n        primary_low = int(\n            number\n            * PRIMARY_SPAN_LO\n        )\n        primary_high = int(\n            number\n            * PRIMARY_SPAN_HI\n        ) - 1\n        primary_high = max(\n            primary_high,\n            primary_low,\n        )\n\n        legacy_low = int(\n            number\n            * float(\n                LEGACY_ARM[\n                    \"span_lo\"\n                ]\n            )\n        )\n        legacy_high = int(\n            number\n            * float(\n                LEGACY_ARM[\n                    \"span_hi\"\n                ]\n            )\n        ) - 1\n        legacy_high = max(\n            legacy_high,\n            legacy_low,\n        )\n\n        if number > 1:\n            primary_picks = np.linspace(\n                primary_low,\n                primary_high,\n                count,\n            ).round().astype(int)\n\n            legacy_picks = np.linspace(\n                legacy_low,\n                legacy_high,\n                count,\n            ).round().astype(int)\n        else:\n            primary_picks = np.zeros(\n                count,\n                dtype=int,\n            )\n            legacy_picks = np.zeros(\n                count,\n                dtype=int,\n            )\n\n        required_positions = sorted(\n            set(\n                primary_picks.tolist()\n                + legacy_picks.tolist()\n            )\n        )\n\n        pixel_cache = {}\n\n        for position in required_positions:\n            position = min(\n                int(position),\n                number - 1,\n            )\n\n            file_path, _ = files[\n                position\n            ]\n\n            try:\n                pixel_cache[\n                    position\n                ] = read_px(\n                    file_path\n                )\n            except Exception:\n                pixel_cache[\n                    position\n                ] = None\n\n        _fill_variant_volume(\n            primary_volume,\n            offset,\n            primary_picks,\n            pixel_cache,\n            files,\n            med_ps,\n            mm_crop_resize,\n        )\n\n        _fill_variant_volume(\n            legacy_volume,\n            offset,\n            legacy_picks,\n            pixel_cache,\n            files,\n            med_ps,\n            mm_crop_resize,\n        )\n\n        offset += count\n\n        if offset >= MAXS:\n            break\n\n    primary_mask = (\n        primary_volume.reshape(\n            MAXS,\n            -1,\n        ).sum(1)\n        > 0\n    ).astype(\n        np.uint8\n    )\n\n    legacy_mask = (\n        legacy_volume.reshape(\n            MAXS,\n            -1,\n        ).sum(1)\n        > 0\n    ).astype(\n        np.uint8\n    )\n\n    return (\n        primary_volume,\n        primary_mask,\n        legacy_volume,\n        legacy_mask,\n    )\n\n\n# ============================================================================\n# Test-root discovery + weights + main\n# ============================================================================\ndef find_test_root():\n    cands = [\"/kaggle/input/competitions/rsna-knee-abnormality-detection\",\n             \"/kaggle/input/rsna-knee-abnormality-detection\"]\n    for b in cands:\n        if os.path.exists(b + \"/test.csv\"):\n            return b\n    for d, _, f in os.walk(\"/kaggle/input\"):\n        if \"test.csv\" in f and (os.path.isdir(d + \"/test_series\") or os.path.isdir(d + \"/test_images\")):\n            return d\n    for d, _, f in os.walk(\"/kaggle/input\"):\n        if \"test.csv\" in f:\n            return d\n    raise RuntimeError(\"no test root under /kaggle/input\")\n\n\ndef find_weight_file(\n    fname,\n    required=True,\n):\n    direct = [\n        f\"/kaggle/input/raptor-knee-arms/{fname}\",\n        f\"/kaggle/input/raptor-knee-arms/1/{fname}\",\n        f\"/kaggle/input/raptor-cnn336/{fname}\",\n    ]\n\n    for path in direct:\n        if os.path.exists(path):\n            return path\n\n    for directory in sorted(\n        glob.glob(\n            \"/kaggle/input/*/\"\n        )\n    ):\n        if (\n            \"competition\"\n            in directory.lower()\n        ):\n            continue\n\n        hits = glob.glob(\n            os.path.join(\n                directory,\n                \"**\",\n                fname,\n            ),\n            recursive=True,\n        )\n\n        if hits:\n            return hits[0]\n\n    if required:\n        raise RuntimeError(\n            f\"{fname} not found \"\n            \"under /kaggle/input\"\n        )\n\n    return None\n\n\n\ndef _d45_find_swin_checkpoint():\n    roots = [\n        \"/kaggle/input/raptor-knee-arms\",\n        \"/kaggle/input/raptor-knee-arms/1\",\n        \"/kaggle/input/raptor-cnn336\",\n    ]\n\n    candidates = []\n\n    for root in roots:\n        if not os.path.isdir(root):\n            continue\n\n        for pattern in (\n            \"**/*swin*.pt\",\n            \"**/*swin*.pth\",\n            \"**/*swin*.bin\",\n        ):\n            candidates.extend(\n                glob.glob(\n                    os.path.join(\n                        root,\n                        pattern,\n                    ),\n                    recursive=True,\n                )\n            )\n\n    seen = set()\n\n    for path in sorted(candidates):\n        if path in seen:\n            continue\n\n        seen.add(path)\n\n        try:\n            checkpoint = torch.load(\n                path,\n                map_location=\"cpu\",\n                weights_only=False,\n            )\n\n            architecture = str(\n                checkpoint.get(\n                    \"arch\",\n                    \"\",\n                )\n            )\n\n            resolution = int(\n                checkpoint.get(\n                    \"res\",\n                    384,\n                )\n            )\n\n            has_model = (\n                isinstance(\n                    checkpoint,\n                    dict,\n                )\n                and \"model\" in checkpoint\n            )\n\n            del checkpoint\n\n            if (\n                has_model\n                and \"swin\" in architecture.lower()\n            ):\n                return {\n                    \"path\": path,\n                    \"arch\": architecture,\n                    \"res\": resolution,\n                }\n\n        except Exception:\n            continue\n\n    return None\n\n\ndef _d45_run_sparse_swin(\n    arm,\n    test_ids,\n    series_map,\n    series_dir,\n    reader,\n):\n    if arm is None:\n        return None, 0.0\n\n    if (\n        not torch.cuda.is_available()\n        or torch.cuda.device_count() < 1\n    ):\n        return None, 0.0\n\n    device = torch.device(\n        \"cuda:1\"\n        if torch.cuda.device_count() >= 2\n        else \"cuda:0\"\n    )\n\n    model = None\n\n    try:\n        model, resolution = load_model(\n            arm[\"path\"],\n            arm[\"arch\"],\n            arm[\"res\"],\n            device,\n        )\n\n        prediction = np.full(\n            (\n                len(test_ids),\n                len(LAB),\n            ),\n            0.5,\n            dtype=np.float32,\n        )\n\n        success = np.zeros(\n            len(test_ids),\n            dtype=np.bool_,\n        )\n\n        for study_index, study_id in enumerate(\n            test_ids\n        ):\n            try:\n                (\n                    _primary_volume,\n                    _primary_mask,\n                    legacy_volume,\n                    legacy_mask,\n                ) = build_study_pair(\n                    study_id,\n                    series_map,\n                    series_dir,\n                    reader,\n                )\n\n                windows = eval_windows(\n                    legacy_volume,\n                    legacy_mask,\n                    k=24,\n                    res=resolution,\n                    norm=NORM,\n                )\n\n                prediction[\n                    study_index\n                ] = infer_probs(\n                    model,\n                    windows,\n                    device,\n                )\n\n                success[\n                    study_index\n                ] = True\n\n                del (\n                    _primary_volume,\n                    _primary_mask,\n                    legacy_volume,\n                    legacy_mask,\n                    windows,\n                )\n\n            except Exception as error:\n                print(\n                    \"[DINOsaur V4.5] \"\n                    f\"Swin study {study_index} fallback: \"\n                    f\"{type(error).__name__}: {error}\",\n                    flush=True,\n                )\n\n            if (\n                (study_index + 1) % 150 == 0\n                or study_index + 1 == len(test_ids)\n            ):\n                print(\n                    \"[DINOsaur V4.5] \"\n                    f\"Swin {study_index+1}/{len(test_ids)}\",\n                    flush=True,\n                )\n\n        fraction = float(\n            success.mean()\n        )\n\n        if fraction < 0.97:\n            return None, fraction\n\n        return prediction, fraction\n\n    except Exception as error:\n        print(\n            \"[DINOsaur V4.5] \"\n            \"optional Swin disabled safely: \"\n            f\"{type(error).__name__}: {error}\",\n            flush=True,\n        )\n\n        return None, 0.0\n\n    finally:\n        if model is not None:\n            del model\n\n        gc.collect()\n\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n\ndef main():\n    import pandas as pd\n    t0 = time.time()\n    dev = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    ngpu = torch.cuda.device_count()\n    print(f\"device {dev} | gpus {ngpu} | torch {torch.__version__}\", flush=True)\n\n    ROOT = find_test_root()\n    tsdir = ROOT + \"/test_series\"\n    if not os.path.isdir(tsdir):\n        tsdir = ROOT + \"/test_images\"\n    print(\"test root:\", ROOT, \"| series dir:\", tsdir, flush=True)\n\n    test = pd.read_csv(ROOT + \"/test.csv\"); test[\"StudyInstanceUID\"] = test[\"StudyInstanceUID\"].astype(str)\n    test_ids = test[\"StudyInstanceUID\"].tolist()\n    tser = pd.read_csv(ROOT + \"/test_series.csv\")\n    tser[\"StudyInstanceUID\"] = tser[\"StudyInstanceUID\"].astype(str)\n    tser[\"SeriesInstanceUID\"] = tser[\"SeriesInstanceUID\"].astype(str)\n    SER = {k: v.to_dict(\"records\") for k, v in tser.groupby(\"StudyInstanceUID\")}\n    print(f\"test studies {len(test_ids)} | test series {len(tser)}\", flush=True)\n\n    sub_cols = [\"StudyInstanceUID\"] + LAB\n    ssub = os.path.join(ROOT, \"sample_submission.csv\")\n    if os.path.exists(ssub):\n        sub_cols = list(pd.read_csv(ssub, nrows=1).columns)\n\n    reader = _make_reader()\n\n    number_studies = len(\n        test_ids\n    )\n\n    primary_predictions = np.full(\n        (\n            number_studies,\n            len(LAB),\n        ),\n        0.5,\n        np.float32,\n    )\n\n    legacy_predictions = np.full_like(\n        primary_predictions,\n        0.5,\n    )\n\n    # V5.4 keeps the exact V4.5 probabilities while harvesting target-local\n    # evidence and learned spatial attention from the same backbone forward.\n    n_slots_v54 = len(SLOTS)\n    primary_local_evidence = np.full_like(primary_predictions, np.nan)\n    legacy_local_evidence = np.full_like(primary_predictions, np.nan)\n    primary_slot_mass = np.zeros(\n        (number_studies, len(LAB), n_slots_v54),\n        dtype=np.float32,\n    )\n    legacy_slot_mass = np.zeros_like(primary_slot_mass)\n    primary_att_peak = np.full_like(primary_predictions, np.nan)\n    legacy_att_peak = np.full_like(primary_predictions, np.nan)\n    primary_att_entropy = np.full_like(primary_predictions, np.nan)\n    legacy_att_entropy = np.full_like(primary_predictions, np.nan)\n    primary_evidence_success = np.zeros(number_studies, dtype=np.bool_)\n\n    legacy_success = np.zeros(\n        number_studies,\n        dtype=np.bool_,\n    )\n\n    if (\n        torch.cuda.is_available()\n        and torch.cuda.device_count() >= 1\n    ):\n        primary_device = torch.device(\n            \"cuda:0\"\n        )\n    else:\n        primary_device = torch.device(\n            \"cpu\"\n        )\n\n    legacy_path = find_weight_file(\n        LEGACY_ARM[\n            \"file\"\n        ],\n        required=False,\n    )\n\n    legacy_enabled = (\n        legacy_path is not None\n        and torch.cuda.is_available()\n        and torch.cuda.device_count() >= 2\n    )\n\n    primary_path = find_weight_file(\n        ARMS[0][\n            \"file\"\n        ],\n        required=True,\n    )\n\n    primary_model, primary_res = load_model(\n        primary_path,\n        ARMS[0][\n            \"arch\"\n        ],\n        ARMS[0][\n            \"res\"\n        ],\n        primary_device,\n    )\n\n    print(\n        \"[DINOsaur V4.5] primary \"\n        f\"{ARMS[0]['file']} \"\n        f\"on {primary_device}\",\n        flush=True,\n    )\n\n    legacy_model = None\n    legacy_device = None\n    legacy_res = None\n\n    if legacy_enabled:\n        legacy_device = torch.device(\n            \"cuda:1\"\n        )\n\n        try:\n            legacy_model, legacy_res = load_model(\n                legacy_path,\n                LEGACY_ARM[\n                    \"arch\"\n                ],\n                LEGACY_ARM[\n                    \"res\"\n                ],\n                legacy_device,\n            )\n\n            print(\n                \"[DINOsaur V4.5] complement \"\n                f\"{LEGACY_ARM['file']} \"\n                f\"on {legacy_device}\",\n                flush=True,\n            )\n\n        except Exception as error:\n            legacy_enabled = False\n            legacy_model = None\n\n            print(\n                \"[DINOsaur V4.5] \"\n                \"legacy checkpoint disabled \"\n                f\"safely: \"\n                f\"{type(error).__name__}: \"\n                f\"{error}\",\n                flush=True,\n            )\n\n    else:\n        print(\n            \"[DINOsaur V4.5] \"\n            \"legacy complement unavailable \"\n            \"or second GPU absent; \"\n            \"exact 0.935 Raptor retained\",\n            flush=True,\n        )\n\n    executor = (\n        ThreadPoolExecutor(\n            max_workers=2\n        )\n        if legacy_enabled\n        else None\n    )\n\n    for study_index, study_id in enumerate(\n        test_ids\n    ):\n        try:\n            (\n                primary_volume,\n                primary_mask,\n                legacy_volume,\n                legacy_mask,\n            ) = build_study_pair(\n                study_id,\n                SER,\n                tsdir,\n                reader,\n            )\n\n            primary_centers = np.asarray(\n                _eval_centers(\n                    primary_mask,\n                    primary_volume.shape[0],\n                    K_EVAL,\n                ),\n                dtype=np.int16,\n            )\n            primary_windows = eval_windows(\n                primary_volume,\n                primary_mask,\n                k=K_EVAL,\n                res=primary_res,\n                norm=NORM,\n            )\n\n            if legacy_enabled:\n                legacy_k = int(LEGACY_ARM[\"k_eval\"])\n                legacy_centers = np.asarray(\n                    _eval_centers(\n                        legacy_mask,\n                        legacy_volume.shape[0],\n                        legacy_k,\n                    ),\n                    dtype=np.int16,\n                )\n                legacy_windows = eval_windows(\n                    legacy_volume,\n                    legacy_mask,\n                    k=legacy_k,\n                    res=legacy_res,\n                    norm=NORM,\n                )\n\n                primary_future = executor.submit(\n                    infer_probs_with_evidence,\n                    primary_model,\n                    primary_windows,\n                    primary_centers,\n                    primary_device,\n                )\n\n                legacy_future = executor.submit(\n                    infer_probs_with_evidence,\n                    legacy_model,\n                    legacy_windows,\n                    legacy_centers,\n                    legacy_device,\n                )\n\n                (\n                    primary_prediction,\n                    primary_local,\n                    primary_mass,\n                    primary_peak,\n                    primary_entropy,\n                ) = primary_future.result()\n                primary_evidence_success[study_index] = True\n\n                try:\n                    (\n                        legacy_prediction,\n                        legacy_local,\n                        legacy_mass,\n                        legacy_peak,\n                        legacy_entropy,\n                    ) = legacy_future.result()\n                    legacy_success[study_index] = True\n                except Exception as legacy_error:\n                    legacy_prediction = primary_prediction.copy()\n                    legacy_local = primary_local.copy()\n                    legacy_mass = primary_mass.copy()\n                    legacy_peak = primary_peak.copy()\n                    legacy_entropy = primary_entropy.copy()\n\n                    print(\n                        \"[DINOsaur V5.4] \"\n                        f\"legacy study {study_index} fallback: \"\n                        f\"{type(legacy_error).__name__}: {legacy_error}\",\n                        flush=True,\n                    )\n\n                del legacy_windows, legacy_centers\n\n            else:\n                (\n                    primary_prediction,\n                    primary_local,\n                    primary_mass,\n                    primary_peak,\n                    primary_entropy,\n                ) = infer_probs_with_evidence(\n                    primary_model,\n                    primary_windows,\n                    primary_centers,\n                    primary_device,\n                )\n                primary_evidence_success[study_index] = True\n                legacy_prediction = primary_prediction.copy()\n                legacy_local = primary_local.copy()\n                legacy_mass = primary_mass.copy()\n                legacy_peak = primary_peak.copy()\n                legacy_entropy = primary_entropy.copy()\n\n            primary_predictions[study_index] = primary_prediction\n            legacy_predictions[study_index] = legacy_prediction\n            primary_local_evidence[study_index] = primary_local\n            legacy_local_evidence[study_index] = legacy_local\n            primary_slot_mass[study_index] = primary_mass\n            legacy_slot_mass[study_index] = legacy_mass\n            primary_att_peak[study_index] = primary_peak\n            legacy_att_peak[study_index] = legacy_peak\n            primary_att_entropy[study_index] = primary_entropy\n            legacy_att_entropy[study_index] = legacy_entropy\n\n            del (\n                primary_volume,\n                primary_mask,\n                legacy_volume,\n                legacy_mask,\n                primary_windows,\n                primary_centers,\n                primary_prediction,\n                legacy_prediction,\n                primary_local,\n                legacy_local,\n                primary_mass,\n                legacy_mass,\n                primary_peak,\n                legacy_peak,\n                primary_entropy,\n                legacy_entropy,\n            )\n\n        except Exception as error:\n            print(\n                \"[DINOsaur V4.5] \"\n                f\"study {study_index} \"\n                f\"{study_id[:16]} FALLBACK \"\n                f\"({type(error).__name__}: \"\n                f\"{error})\",\n                flush=True,\n            )\n\n        if (\n            (\n                study_index + 1\n            )\n            % 100\n            == 0\n            or study_index + 1\n            == number_studies\n        ):\n            print(\n                \"[DINOsaur V4.5] \"\n                f\"{study_index+1}/\"\n                f\"{number_studies} | \"\n                f\"{time.time()-t0:.0f}s\",\n                flush=True,\n            )\n\n    if executor is not None:\n        executor.shutdown(\n            wait=True\n        )\n\n    del primary_model\n\n    if legacy_model is not None:\n        del legacy_model\n\n    gc.collect()\n\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n    primary_rank = rankpct(\n        np.clip(\n            primary_predictions,\n            0,\n            1,\n        )\n    )\n\n    raptor_rank = (\n        primary_rank.copy()\n    )\n\n    legacy_fraction = float(\n        legacy_success.mean()\n    ) if legacy_enabled else 0.0\n\n    if (\n        legacy_enabled\n        and legacy_fraction >= 0.98\n    ):\n        legacy_rank = rankpct(\n            np.clip(\n                legacy_predictions,\n                0,\n                1,\n            )\n        )\n\n        # The expanded MaxSpan checkpoint's documented largest gains are ACL,\n        # both menisci, Lateral OA and Fracture. Keep it almost pure there.\n        # Use the previous 0.934 checkpoint only for the remaining findings,\n        # where it may restore complementary ordering.\n        complement_weight = {\n            \"ACL\": 0.025,\n            \"MCL\": 0.16,\n            \"Medial Meniscus\": 0.025,\n            \"Lateral Meniscus\": 0.020,\n            \"Medial OA\": 0.10,\n            \"Lateral OA\": 0.025,\n            \"PF OA\": 0.12,\n            \"Effusion\": 0.10,\n            \"Synovitis\": 0.16,\n            \"Baker's\": 0.12,\n            \"Contusion\": 0.12,\n            \"Fracture\": 0.020,\n        }\n\n        complement_log = []\n\n        for target_index, target in enumerate(\n            LAB\n        ):\n            weight = float(\n                complement_weight.get(\n                    target,\n                    0.0,\n                )\n            )\n\n            if weight <= 0:\n                continue\n\n            correlation = float(\n                np.corrcoef(\n                    primary_rank[\n                        :,\n                        target_index,\n                    ],\n                    legacy_rank[\n                        :,\n                        target_index,\n                    ],\n                )[0, 1]\n            )\n\n            if not np.isfinite(\n                correlation\n            ):\n                weight = 0.0\n            elif correlation > 0.992:\n                weight *= 0.50\n            elif correlation < 0.65:\n                weight *= 0.40\n\n            if weight <= 0:\n                continue\n\n            raptor_rank[\n                :,\n                target_index,\n            ] = (\n                (\n                    1.0\n                    - weight\n                )\n                * primary_rank[\n                    :,\n                    target_index,\n                ]\n                + weight\n                * legacy_rank[\n                    :,\n                    target_index,\n                ]\n            )\n\n            complement_log.append(\n                (\n                    target,\n                    weight,\n                    correlation,\n                )\n            )\n\n        raptor_rank = rankpct(\n            raptor_rank\n        )\n\n        print(\n            \"[DINOsaur V4.5] \"\n            \"legacy complement: \"\n            + \"; \".join(\n                f\"{target}=w{weight:.3f},\"\n                f\"corr={correlation:.3f}\"\n                for (\n                    target,\n                    weight,\n                    correlation,\n                )\n                in complement_log\n            ),\n            flush=True,\n        )\n\n    else:\n        print(\n            \"[DINOsaur V4.5] \"\n            f\"legacy success={legacy_fraction:.3f}; \"\n            \"exact primary Raptor used\",\n            flush=True,\n        )\n\n    ranks = raptor_rank\n\n    # Optional third architecture. It is intentionally only a small residual.\n    # If the checkpoint is absent, incompatible, too slow, or incomplete,\n    # this block becomes a strict no-op and preserves the 0.936 anchor.\n    swin_arm = _d45_find_swin_checkpoint()\n\n    if swin_arm is not None:\n        print(\n            \"[DINOsaur V4.5] optional Swin found: \"\n            f\"{os.path.basename(swin_arm['path'])} | \"\n            f\"{swin_arm['arch']} | res={swin_arm['res']}\",\n            flush=True,\n        )\n\n        swin_prediction, swin_fraction = _d45_run_sparse_swin(\n            swin_arm,\n            test_ids,\n            SER,\n            tsdir,\n            reader,\n        )\n\n        if (\n            swin_prediction is not None\n            and swin_fraction >= 0.97\n        ):\n            swin_rank = rankpct(\n                np.clip(\n                    swin_prediction,\n                    0,\n                    1,\n                )\n            )\n\n            diversity_weight = {\n                \"ACL\": 0.035,\n                \"MCL\": 0.075,\n                \"Medial Meniscus\": 0.045,\n                \"Lateral Meniscus\": 0.040,\n                \"Medial OA\": 0.065,\n                \"Lateral OA\": 0.045,\n                \"PF OA\": 0.070,\n                \"Effusion\": 0.060,\n                \"Synovitis\": 0.080,\n                \"Baker's\": 0.070,\n                \"Contusion\": 0.070,\n                \"Fracture\": 0.040,\n            }\n\n            swin_log = []\n\n            for target_index, target in enumerate(\n                LAB\n            ):\n                weight = float(\n                    diversity_weight[\n                        target\n                    ]\n                )\n\n                correlation = float(\n                    np.corrcoef(\n                        ranks[\n                            :,\n                            target_index,\n                        ],\n                        swin_rank[\n                            :,\n                            target_index,\n                        ],\n                    )[0, 1]\n                )\n\n                if not np.isfinite(\n                    correlation\n                ):\n                    weight = 0.0\n                elif correlation > 0.985:\n                    weight *= 0.45\n                elif correlation < 0.50:\n                    weight *= 0.35\n\n                if weight <= 0:\n                    continue\n\n                ranks[\n                    :,\n                    target_index,\n                ] = (\n                    (\n                        1.0\n                        - weight\n                    )\n                    * ranks[\n                        :,\n                        target_index,\n                    ]\n                    + weight\n                    * swin_rank[\n                        :,\n                        target_index,\n                    ]\n                )\n\n                swin_log.append(\n                    (\n                        target,\n                        weight,\n                        correlation,\n                    )\n                )\n\n            ranks = rankpct(\n                ranks\n            )\n\n            print(\n                \"[DINOsaur V4.5] Swin residual: \"\n                + \"; \".join(\n                    f\"{target}=w{weight:.3f},corr={correlation:.3f}\"\n                    for (\n                        target,\n                        weight,\n                        correlation,\n                    )\n                    in swin_log\n                ),\n                flush=True,\n            )\n\n    if not np.isfinite(ranks).all():\n        ranks[~np.isfinite(ranks)] = 0.5\n\n    # ------------------------------------------------------------------------\n    # V5.4: Raptor spatial-evidence residual.\n    #\n    # The trained Raptor head is a multiple-instance learner. Its normal output\n    # retains only the attention-weighted study logit. For focal findings the\n    # competition semantics are closer to \"present in at least one local region\".\n    # We therefore harvest max/top2 per-window evidence from the SAME encoded\n    # features, then require two independently trained Raptor checkpoints to\n    # agree on both the evidence ordering and the anatomical slot receiving\n    # attention. No additional backbone pass is performed.\n    # ------------------------------------------------------------------------\n    anchor_ranks = rankpct(np.asarray(ranks, np.float64))\n    global_raptor_reference = rankpct(np.asarray(raptor_rank, np.float64))\n\n    def _v54_rank_safe(values, fallback):\n        frame = pd.DataFrame(np.asarray(values, np.float64))\n        ranked = frame.rank(method=\"average\", pct=True).to_numpy(np.float64)\n        bad = ~np.isfinite(ranked)\n        if bad.any():\n            ranked[bad] = np.asarray(fallback, np.float64)[bad]\n        return ranked\n\n    primary_local_rank = _v54_rank_safe(\n        primary_local_evidence,\n        global_raptor_reference,\n    )\n    legacy_local_rank = _v54_rank_safe(\n        legacy_local_evidence,\n        global_raptor_reference,\n    )\n    local_consensus_rank = rankpct(\n        0.60 * primary_local_rank\n        + 0.40 * legacy_local_rank\n    )\n\n    slot_overlap = np.minimum(\n        np.clip(primary_slot_mass, 0.0, 1.0),\n        np.clip(legacy_slot_mass, 0.0, 1.0),\n    ).sum(axis=2)\n    local_model_agreement = 1.0 - np.abs(\n        primary_local_rank - legacy_local_rank\n    )\n\n    primary_peak_floor = 1.0 / max(int(K_EVAL), 2)\n    legacy_peak_floor = 1.0 / max(int(LEGACY_ARM.get(\"k_eval\", K_EVAL)), 2)\n    primary_concentration = np.clip(\n        (primary_att_peak - primary_peak_floor)\n        / max(1.0 - primary_peak_floor, 1e-6),\n        0.0,\n        1.0,\n    )\n    legacy_concentration = np.clip(\n        (legacy_att_peak - legacy_peak_floor)\n        / max(1.0 - legacy_peak_floor, 1e-6),\n        0.0,\n        1.0,\n    )\n    concentration = np.sqrt(\n        np.clip(primary_concentration, 0.0, 1.0)\n        * np.clip(legacy_concentration, 0.0, 1.0)\n    )\n\n    paired = (\n        primary_evidence_success\n        & legacy_success\n    )\n    paired_matrix = paired[:, None].astype(np.float64)\n\n    slot_gate = np.clip(\n        (slot_overlap - 0.40) / 0.40,\n        0.0,\n        1.0,\n    )\n    agreement_gate = np.clip(\n        (local_model_agreement - 0.68) / 0.27,\n        0.0,\n        1.0,\n    )\n    evidence_gate = (\n        paired_matrix\n        * slot_gate\n        * agreement_gate\n        * (0.75 + 0.25 * concentration)\n    )\n\n    def _v54_evidence_variant(strength):\n        candidate = anchor_ranks.copy()\n        diagnostic = {}\n        for target_index, target in enumerate(LAB):\n            if target not in _V54_FOCAL_TARGETS:\n                continue\n\n            local_column = local_consensus_rank[:, target_index]\n            global_column = global_raptor_reference[:, target_index]\n            corr = float(np.corrcoef(local_column, global_column)[0, 1])\n            gate_column = evidence_gate[:, target_index].copy()\n            active_fraction = float(np.mean(gate_column > 0.10))\n\n            # A useful residual must be related enough to be meaningful but not\n            # so correlated that it is merely a duplicate of the study logit.\n            usable = (\n                np.isfinite(corr)\n                and 0.58 <= corr <= 0.997\n                and active_fraction >= 0.08\n                and number_studies >= 40\n            )\n\n            residual = np.clip(\n                local_column - global_column,\n                -0.22,\n                0.22,\n            )\n            residual[np.abs(residual) < 0.025] = 0.0\n\n            if usable:\n                candidate[:, target_index] = (\n                    anchor_ranks[:, target_index]\n                    + float(strength)\n                    * gate_column\n                    * residual\n                )\n\n            diagnostic[target] = {\n                \"pool\": (\n                    \"top2\"\n                    if target in _V54_TOP2_TARGETS\n                    else \"max\"\n                ),\n                \"corr_local_to_global\": corr,\n                \"mean_slot_overlap\": float(\n                    np.mean(slot_overlap[paired, target_index])\n                ) if paired.any() else 0.0,\n                \"active_fraction\": active_fraction,\n                \"mean_abs_residual\": float(np.mean(np.abs(residual))),\n                \"usable\": bool(usable),\n            }\n\n        return rankpct(candidate), diagnostic\n\n    safe_ranks, safe_diag = _v54_evidence_variant(0.08)\n    main_ranks, main_diag = _v54_evidence_variant(0.14)\n    probe_ranks, probe_diag = _v54_evidence_variant(0.22)\n\n    evidence_active = any(\n        bool(item.get(\"usable\"))\n        for item in main_diag.values()\n    )\n    globals()[\"V54_EVIDENCE_ACTIVE\"] = bool(evidence_active)\n\n    if evidence_active:\n        print(\n            \"[DINOsaur V5.4] spatial-evidence active: \"\n            + \"; \".join(\n                f\"{target}:{info['pool']},\"\n                f\"corr={info['corr_local_to_global']:.3f},\"\n                f\"slot={info['mean_slot_overlap']:.3f},\"\n                f\"active={info['active_fraction']:.2%}\"\n                for target, info in main_diag.items()\n                if info[\"usable\"]\n            ),\n            flush=True,\n        )\n    else:\n        print(\n            \"[DINOsaur V5.4] spatial-evidence guard rejected all focal targets; \"\n            \"exact V4.5 Raptor anchor retained\",\n            flush=True,\n        )\n\n    def _v54_write_raptor(values, filename):\n        frame = pd.DataFrame(\n            np.asarray(values, np.float32),\n            columns=LAB,\n        )\n        frame.insert(0, \"StudyInstanceUID\", test_ids)\n        frame = frame[sub_cols]\n        assert list(frame.columns) == sub_cols, \"column order drift\"\n        assert frame[\"StudyInstanceUID\"].tolist() == test_ids, \"row identity drift\"\n        assert np.isfinite(frame[LAB].values).all()\n        path = os.path.join(\"/kaggle/working\", filename)\n        frame.to_csv(path, index=False)\n        return path\n\n    _v54_write_raptor(\n        anchor_ranks,\n        \"submission_raptor_v45_anchor.csv\",\n    )\n    _v54_write_raptor(\n        safe_ranks,\n        \"submission_raptor_v54_safe.csv\",\n    )\n    _v54_write_raptor(\n        main_ranks,\n        \"submission_raptor_v54_main.csv\",\n    )\n    _v54_write_raptor(\n        probe_ranks,\n        \"submission_raptor_v54_probe.csv\",\n    )\n    out = _v54_write_raptor(\n        main_ranks if evidence_active else anchor_ranks,\n        \"submission_coatnet.csv\",\n    )\n\n    diagnostics = {\n        \"version\": \"DINOsaur V5.4 SpatialEvidence\",\n        \"evidence_active\": bool(evidence_active),\n        \"paired_raptor_fraction\": float(np.mean(paired)),\n        \"strength\": {\n            \"safe\": 0.08,\n            \"main\": 0.14,\n            \"probe\": 0.22,\n        },\n        \"targets\": main_diag,\n        \"runtime\": \"same backbone passes as V4.5; head-only evidence extraction\",\n    }\n    with open(\n        \"/kaggle/working/dinosaur_v54_spatial_evidence.json\",\n        \"w\",\n    ) as handle:\n        json.dump(diagnostics, handle, indent=2, sort_keys=True)\n\n    print(\n        \"wrote\", out, \"|\", len(test_ids),\n        \"rows x\", len(sub_cols), \"cols\",\n        flush=True,\n    )\n    print(f\"DONE {time.time()-t0:.0f}s\", flush=True)\n\n\nif __name__ == \"__main__\":\n    try:\n        main()\n    except Exception as _coat_exc:\n        import traceback as _coat_traceback\n        print(f\"CoAtNet branch failed; retaining transformer submission: {type(_coat_exc).__name__}: {_coat_exc}\", flush=True)\n        _coat_traceback.print_exc()\n\n\n\n# Hidden-rerun fail-safe final fusion.\nfrom pathlib import Path as _D42Path\nimport shutil as _d42_shutil\n\n_d42_primary = _D42Path('/kaggle/working/submission.csv')\n_d42_backup = _D42Path('/kaggle/working/.v54_transformer_backup.csv')\n\nif _d42_primary.is_file():\n    _d42_shutil.copy2(_d42_primary, _d42_backup)\n\ntry:\n    from pathlib import Path as _BlendPath\n    import numpy as _blend_np\n    import pandas as _blend_pd\n\n    _blend_work = _BlendPath('/kaggle/working')\n    _blend_transformer_path = _blend_work / 'submission.csv'\n    _blend_transformer = _blend_pd.read_csv(\n        _blend_transformer_path,\n        dtype={'StudyInstanceUID': str},\n    )\n    _blend_labels = [\n        c for c in _blend_transformer.columns\n        if c != 'StudyInstanceUID'\n    ]\n    _blend_tr = _blend_transformer[\n        _blend_labels\n    ].rank(method='average', pct=True)\n\n    def _v54_fuse(raptor_path):\n        raptor = _blend_pd.read_csv(\n            raptor_path,\n            dtype={'StudyInstanceUID': str},\n        )\n        if raptor.columns.tolist() != _blend_transformer.columns.tolist():\n            raise RuntimeError(\n                f'Raptor/transformer schema mismatch: {raptor_path.name}'\n            )\n        if (\n            raptor['StudyInstanceUID'].tolist()\n            != _blend_transformer['StudyInstanceUID'].tolist()\n        ):\n            raise RuntimeError(\n                f'Raptor/transformer study order mismatch: {raptor_path.name}'\n            )\n\n        rr = raptor[_blend_labels].rank(\n            method='average',\n            pct=True,\n        )\n        output = _blend_transformer.copy()\n\n        # EXACT V4.5 transformer/Raptor fusion rule. V5.4 changes only the\n        # Raptor ordering supplied to this already validated fusion.\n        weight = {\n            label: 0.50\n            for label in _blend_labels\n        }\n        if globals().get(\n            'V18_CALIBRATOR_APPLIED',\n            False,\n        ):\n            weight.update(\n                {\n                    'ACL': 0.54,\n                    'Medial Meniscus': 0.565,\n                    'Lateral Meniscus': 0.625,\n                    'Lateral OA': 0.56,\n                    'Fracture': 0.625,\n                }\n            )\n\n        base_0935 = {\n            label: 0.50\n            for label in _blend_labels\n        }\n        base_0935.update(\n            {\n                'Medial Meniscus': 0.52,\n                'Lateral Meniscus': 0.54,\n                'Fracture': 0.54,\n            }\n        )\n\n        for label in _blend_labels:\n            correlation = float(\n                _blend_np.corrcoef(\n                    _blend_tr[label].to_numpy(_blend_np.float64),\n                    rr[label].to_numpy(_blend_np.float64),\n                )[0, 1]\n            )\n            if not _blend_np.isfinite(correlation):\n                weight[label] = base_0935[label]\n            elif correlation > 0.992:\n                weight[label] = (\n                    0.65 * weight[label]\n                    + 0.35 * base_0935[label]\n                )\n            elif correlation < 0.60:\n                weight[label] = (\n                    0.50 * weight[label]\n                    + 0.50 * base_0935[label]\n                )\n\n            cw = float(weight[label])\n            output[label] = (\n                (1.0 - cw) * _blend_tr[label]\n                + cw * rr[label]\n            )\n\n        output[_blend_labels] = output[\n            _blend_labels\n        ].rank(method='average', pct=True)\n\n        values = output[_blend_labels].to_numpy(_blend_np.float64)\n        if (\n            not _blend_np.isfinite(values).all()\n            or values.min() < 0\n            or values.max() > 1\n        ):\n            raise RuntimeError(\n                f'invalid V5.4 fused values: {raptor_path.name}'\n            )\n        return output\n\n    variants = {\n        'v45_anchor': _blend_work / 'submission_raptor_v45_anchor.csv',\n        'v54_safe': _blend_work / 'submission_raptor_v54_safe.csv',\n        'v54_main': _blend_work / 'submission_raptor_v54_main.csv',\n        'v54_probe': _blend_work / 'submission_raptor_v54_probe.csv',\n    }\n\n    fused = {}\n    for name, path in variants.items():\n        if not path.is_file():\n            continue\n        frame = _v54_fuse(path)\n        fused[name] = frame\n        frame.to_csv(\n            _blend_work / f'submission_{name}.csv',\n            index=False,\n        )\n\n    evidence_active = bool(\n        globals().get('V54_EVIDENCE_ACTIVE', False)\n    )\n    if evidence_active and 'v54_main' in fused:\n        fused['v54_main'].to_csv(\n            _blend_transformer_path,\n            index=False,\n        )\n        print(\n            'final submission.csv = DINOsaur V5.4 SpatialEvidence; '\n            f\"{fused['v54_main'].shape}\",\n            flush=True,\n        )\n    elif 'v45_anchor' in fused:\n        fused['v45_anchor'].to_csv(\n            _blend_transformer_path,\n            index=False,\n        )\n        print(\n            '[DINOsaur V5.4] evidence unavailable/rejected; '\n            'exact V4.5 fused anchor restored',\n            flush=True,\n        )\n    else:\n        print(\n            '[DINOsaur V5.4] Raptor outputs unavailable; '\n            'transformer submission retained',\n            flush=True,\n        )\n\n    for temp_path in (\n        _blend_work / 'submission_coatnet.csv',\n        _blend_work / 'submission_transformer_0920.csv',\n    ):\n        try:\n            if temp_path.is_file():\n                temp_path.unlink()\n        except OSError:\n            pass\n\nexcept Exception as _d42_error:\n    import traceback as _d42_traceback\n    print(\n        '[DINOsaur V5.4] final fusion failed; '\n        'restoring calibrated transformer submission: '\n        f'{type(_d42_error).__name__}: {_d42_error}',\n        flush=True,\n    )\n    _d42_traceback.print_exc()\n    if _d42_backup.is_file():\n        _d42_shutil.copy2(\n            _d42_backup,\n            _d42_primary,\n        )\n\nfinally:\n    try:\n        if _d42_backup.is_file():\n            _d42_backup.unlink()\n    except OSError:\n        pass\n\nif not _d42_primary.is_file():\n    raise RuntimeError(\n        'submission.csv missing after V5.4 fail-safe'\n    )\n","metadata":{},"outputs":[],"execution_count":null}]}