{"cells":[{"cell_type":"markdown","id":"rad-title","metadata":{},"source":"# RSNA Knee — RadImageNet ResNet-50 only\n\n元NotebookのStage 3からRadImageNet ResNet-50の推論だけを抽出し、`/kaggle/working/submission.csv`を生成します。前段Transformer提出との融合と、その予測を特徴量に含むV18キャリブレーターは除外しています。"},{"cell_type":"code","execution_count":null,"id":"rad-dicom-preprocessor","metadata":{},"outputs":[],"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/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\n"},{"id":"86a37ee7","cell_type":"markdown","source":"## Stage 3: Dual-Layout RadImageNet ResNet-50 Arm & V18 Calibration (Score: 0.920)\n- Extracts 2048-dim GAP features from RadImageNet ResNet-50.\n- Evaluates dual pulse sequence layouts: Layout 1 (E13: Sag-FS, Cor-FS, Ax-FS, Sag-noFS) + Layout 2 (E11: Sag-noFS, Cor-noFS, Ax-noFS, Sag-FS).\n- Applies 88-feature protocol-aware transformer linear calibration to correct scanner metadata variation.\n- Preserves focal gating: $\\alpha_{\\text{Baker's}} = 0.0$ and $\\alpha_{\\text{Fracture}} = 0.0$.\n","metadata":{}},{"cell_type":"code","execution_count":null,"id":"rad-code","metadata":{},"outputs":[],"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_REFERENCE_HEADS_SHA256 = '0f465649799ecfbccaac1767844639e7ced44e1bc9babde6e4bac7c5d9b89eaa'\n_RAD_ENCODER_SHA256 = '08629f7e7bd3e29b8ee9522ca3f65ce4d010a7ddf74f0ea3c7e3f3d0bbab0734'\n_RAD_E13_HEADS_SHA256 = 'ad9f19af73bfdf4e49263c0e45060dc3cb239e1195039b26dc8c0a3a6bcd1a8a'\n_RAD_E13_MEMBER_WEIGHT = 0.5\n_RAD_V48_SECOND_ALPHA = 0.15\n_RAD_TOKEN_DIM, _RAD_HEAD_DIM = (2048, 512)\n_RAD_E11_SLOTS = [('SAG_NOFS', 'Sagittal', None, False), ('COR_NOFS', 'Coronal', None, False), ('AX_NOFS', 'Axial', None, False), ('SAG_FS', 'Sagittal', None, True)]\n_RAD_E11_CROP_MM = 130.0\n_RAD_E13_SLOTS = [('SAG_FS', 'Sagittal', None, True), ('COR_FS', 'Coronal', None, True), ('AX_FS', 'Axial', None, True), ('SAG_NOFS', 'Sagittal', None, False)]\n_RAD_E13_CROP_MM = 130.0\n_RAD_E13_CACHE_SLICES = 8\n_RAD_E13_IMG = 224\nSLOTS = [('SAG_FS', 'Sagittal', None, True), ('COR_FS', 'Coronal', None, True), ('AX_FS', 'Axial', None, True)]\nN_SLOT = len(SLOTS)\nCACHE_SLICES = 8\n\ndef _rad_sha256(path, chunk=8 << 20):\n    digest = _rad_hashlib.sha256()\n    with open(path, 'rb') as handle:\n        for block in iter(lambda: handle.read(chunk), b''):\n            digest.update(block)\n    return digest.hexdigest()\n\ndef _rad_find_file(name, expected_sha=None, explicit_env=None):\n    files = {\n        _RAD_ENCODER_SHA256: ASSET / 'resnet-50-radimagenet-marwan/ResNet50.pt',\n        _RAD_REFERENCE_HEADS_SHA256: ASSET / 'rsna-knee-e9-radimagenet-heads-v15/v52_radimagenet_heads.pt',\n        _RAD_E13_HEADS_SHA256: ASSET / 'kernel-sources/rsna-knee-e13-train/rsna_rad_e11/v52_e11_heads.pt',\n    }\n    path = files.get(expected_sha)\n    if path is not None and path.is_file() and _rad_sha256(path) == expected_sha:\n        return path\n    for p in Path('/kaggle/input').glob(f'**/{name}'):\n        if expected_sha is None or _rad_sha256(p) == expected_sha:\n            return p\n    for p in Path('/kaggle/input').glob('**/*.pt'):\n        if p.name == name or (expected_sha is not None and _rad_sha256(p) == expected_sha):\n            return p\n    raise FileNotFoundError(f'RadImageNet asset {name} ({expected_sha}) not found under /kaggle/input')\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\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    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\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\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, reference_probability, features, token_mask, pixels, masks\n    _rad_gc.collect()\n    _rad_torch.cuda.empty_cache()\n\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, features, token_mask, pixels, masks\n    _rad_gc.collect()\n    _rad_torch.cuda.empty_cache()\n\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\n    # Standalone output: no baseline Transformer predictions and no 88-feature\n    # calibrator, because both consume outputs from other model arms.\n    rad_rank = _rad_rank_columns(\n        (1.0 - _RAD_V48_SECOND_ALPHA) * reference_rank\n        + _RAD_V48_SECOND_ALPHA * pass2_rank\n    )\n    final = _rad_pd.DataFrame(rad_rank, columns=_RAD_LABELS)\n    final.insert(0, 'StudyInstanceUID', expected_ids)\n    _rad_validate(final, expected_ids)\n    final.to_csv(primary, index=False)\n    print(f'wrote {primary} from RadImageNet ResNet-50 only; {final.shape}', flush=True)\n\n_rad_main()\n"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.12"}},"nbformat":4,"nbformat_minor":5}