{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11"},"kaggle":{"title":"RSNA Knee target 0.82 cache 2h hidden inference"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from __future__ import annotations\n\nimport gc\nimport json\nimport os\nimport re\nimport time\nimport traceback\nfrom concurrent.futures import ThreadPoolExecutor\nfrom pathlib import Path\n\nfor _v in (\"OMP_NUM_THREADS\", \"OPENBLAS_NUM_THREADS\", \"MKL_NUM_THREADS\"):\n    os.environ.setdefault(_v, \"4\")\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nT0 = time.time()\nSEED = 20260806\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\nTARGETS = [\"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\", \"Medial OA\",\n           \"Lateral OA\", \"PF OA\", \"Effusion\", \"Synovitis\", \"Baker's\",\n           \"Contusion\", \"Fracture\"]\n\nIMG = 224\nCROP_MM = 160.0\nGROUP = 3\nN_GROUP = 3\nCACHE_SLICES = GROUP * N_GROUP\nHDR_THREADS = 16\nPIX_THREADS = 12\nEVAL_BATCH = 12\nTIME_BUDGET = 8.0 * 3600\nBACKBONE_VARIANT = \"small\"\nUNFREEZE_LAST = 6\nMODEL_FILE = \"rsna_20260807_v4.pt\"\n\nLAT_FALLBACK        = \"auto\"\nLAT_MIN_AGREEMENT   = 0.85\nLAT_MIN_OFFSET_MM   = 5.0\n\nSLOTS_RECOVERED = [\n    (\"SAG_FLUID_FS\", \"Sagittal\", True, True),\n    (\"COR_FLUID_FS\", \"Coronal\", True, True),\n    (\"AX_FLUID_FS\", \"Axial\", True, True),\n    (\"SAG_FLUID_NOFS\", \"Sagittal\", True, False),\n    (\"COR_T1\", \"Coronal\", False, False),\n    (\"SAG_T1\", \"Sagittal\", False, False),\n]\n\nSLOTS_PUBLIC = [\n    (\"SAG_FLUID\", \"Sagittal\", None, True),\n    (\"COR_FLUID\", \"Coronal\", None, True),\n    (\"AX_FLUID\", \"Axial\", None, True),\n    (\"SAG_STRUCT\", \"Sagittal\", None, False),\n    (\"COR_STRUCT\", \"Coronal\", None, False),\n    (\"AX_STRUCT\", \"Axial\", None, False),\n]\n\nSLOT_SCHEME = os.environ.get(\"SLOT_SCHEME\", \"recovered\")\nSLOTS = SLOTS_PUBLIC if SLOT_SCHEME == \"public\" else SLOTS_RECOVERED\nN_SLOT = len(SLOTS)\n\nFATSAT_OPTS = {\"FS\", \"FATSAT\", \"FAT_SAT\", \"FSAT\"}\n_SEP = re.compile(r\"[_\\-.]\")\n_FATSAT_RX = re.compile(r\"\\bfs\\b|fatsat|fat sat|\\bstir\\b|\\bspair\\b|\\bspir\\b|\\bwe\\b|\"\n                        r\"water excit|\\btirm\\b|\\bsting\\b|\\bfatsup\\b\")\n_T1_RX = re.compile(r\"\\bt1\\b|\\bt1w\\b\")\n_T2_RX = re.compile(r\"\\bt2\\b|\\bt2w\\b\")\n_PD_RX = re.compile(r\"\\bpd\\b|\\bpdw\\b|proton|\\bdp\\b|dens\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def log(msg):\n    print(f\"[{time.time() - T0:7.1f}s] {msg}\", flush=True)\n\n\ndef find_root():\n    for c in [Path(\"/kaggle/input/competitions/rsna-knee-abnormality-detection\"),\n              Path(\"/kaggle/input/rsna-knee-abnormality-detection\"),\n              Path(\"data\"), Path(\".\")]:\n        if (c / \"test.csv\").is_file() and (c / \"test_series\").is_dir():\n            return c\n    base = Path(\"/kaggle/input\")\n    if base.is_dir():\n        for depth1 in sorted(p for p in base.iterdir() if p.is_dir()):\n            for cand in [depth1] + sorted(p for p in depth1.iterdir() if p.is_dir()):\n                if (cand / \"test.csv\").is_file() and (cand / \"test_series\").is_dir():\n                    return cand\n    raise FileNotFoundError(\"competition mount not found\")\n\n\ndef find_dinov2(variant=\"small\"):\n    base = Path(\"/kaggle/input\")\n    if not base.is_dir():\n        return None\n    hits = []\n    for root, dirs, files in os.walk(base):\n        dirs[:] = [d for d in dirs if d not in (\"train_series\", \"test_series\")]\n        if \"config.json\" in files and \"dinov2\" in root.lower():\n            hits.append(Path(root))\n    for h in hits:\n        if variant in str(h).lower():\n            return h\n    return hits[0] if hits else None\n\n\ndef find_model_path():\n    direct = [\n        Path(\"/kaggle/input/datasets/tonylica/rsna2026-models\") / MODEL_FILE,\n        Path(MODEL_FILE),\n    ]\n    for p in direct:\n        if p.is_file():\n            return p\n    base = Path(\"/kaggle/input\")\n    if base.is_dir():\n        for p in base.rglob(MODEL_FILE):\n            if p.is_file():\n                return p\n    raise FileNotFoundError(f\"{MODEL_FILE} is not attached\")\n\n\nROOT = find_root()\nlog(f\"input root: {ROOT}\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"HDR_TAGS = [\"SeriesDescription\", \"SequenceName\", \"ScanOptions\", \"ScanningSequence\",\n            \"RepetitionTime\", \"EchoTime\", \"Laterality\", \"ImageLaterality\",\n            \"ImagePositionPatient\", \"PixelSpacing\", \"Rows\",\n            \"Columns\", \"RescaleSlope\", \"RescaleIntercept\"]\n\n\ndef probe(item):\n    split, study, series, path = item\n    row = {\"split\": split, \"StudyInstanceUID\": study, \"SeriesInstanceUID\": series,\n           \"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]),\n                             stop_before_pixels=True, force=True)\n        for t in HDR_TAGS:\n            v = getattr(ds, t, None)\n            if v is None:\n                row[t] = None\n            elif isinstance(v, (list, tuple)) or type(v).__name__ == \"MultiValue\":\n                row[t] = \"|\".join(str(x) for x in v)\n            else:\n                row[t] = str(v)\n    except Exception as exc:\n        row[\"err\"] = str(exc)[:120]\n    return row\n\n\ndef walk(split):\n    base = ROOT / split\n    items = []\n    if not base.is_dir():\n        return pd.DataFrame()\n    for study in os.scandir(base):\n        if study.is_dir():\n            for series in os.scandir(study.path):\n                if series.is_dir():\n                    items.append((split, study.name, series.name, series.path))\n    with ThreadPoolExecutor(max_workers=HDR_THREADS) as pool:\n        rows = list(pool.map(probe, items))\n    return pd.DataFrame(rows)\n\n\ndef annotate(df):\n    \"\"\"Recover fat suppression and pulse-sequence weighting from the header.\"\"\"\n    desc = (df[\"SeriesDescription\"].fillna(\"\") + \" \" + df[\"SequenceName\"].fillna(\"\"))\n    desc = desc.str.lower().str.replace(_SEP, \" \", regex=True)\n\n    opts = df[\"ScanOptions\"].fillna(\"\").str.upper().str.split(\"|\")\n    # GE writes SAT_GEMS for spatial saturation, so ScanOptions must be matched as\n    # exact tokens; a substring test on \"SAT\" fires on non-fat-sat series.\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\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\n    df[\"weight\"] = np.where(t1 & ~t2 & ~pdw, \"T1\",\n                     np.where(t2 & ~pdw, \"T2\",\n                       np.where(pdw, \"PD\",\n                         np.where(gre, \"GRE\",\n                           np.where(tr < 800, \"T1\",\n                             np.where(te > 60, \"T2\",\n                               np.where(tr >= 800, \"PD\", \"UNK\")))))))\n    df[\"fluid\"] = np.isin(df[\"weight\"], [\"PD\", \"T2\"])\n    df[\"px\"] = pd.to_numeric(\n        df[\"PixelSpacing\"].fillna(\"\").str.split(\"|\").str[0].replace(\"\", np.nan),\n        errors=\"coerce\")\n    return df\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- laterality resolution -----------------------------------------------------------\n# The tag is authoritative where it exists. Where it does not, the patient x-coordinate\n# can stand in, but only if it agrees with the tag on the studies that have both: a wrong\n# flip is worse than no flip, and knee coils often place the joint near isocentre, where\n# the sign carries no information at all.\n\ndef _tag_side(g):\n    v = [str(x).strip().upper() for x in g[\"Laterality\"].dropna()]\n    if \"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    return v[0] if v else None\n\n\ndef _position_side(g):\n    xs = []\n    for s in g.get(\"ImagePositionPatient\", pd.Series(dtype=object)).dropna():\n        try:\n            xs.append(float(str(s).split(\"|\")[0]))\n        except Exception:\n            pass\n    if not xs:\n        return None\n    x = float(np.median(xs))\n    if abs(x) < LAT_MIN_OFFSET_MM:\n        return None                      # centred in the coil: the sign means nothing\n    return \"R\" if x < 0 else \"L\"         # LPS: the right knee sits at negative x\n\n\ndef laterality_maps(h):\n    \"\"\"Return (side_by_study, diagnostics). Uses the fallback only if it earns it.\"\"\"\n    tag, pos = {}, {}\n    for st, g in h.groupby(\"StudyInstanceUID\"):\n        tag[st] = _tag_side(g)\n        pos[st] = _position_side(g)\n\n    both = [st for st in tag if tag[st] and pos[st]]\n    agree = float(np.mean([tag[st] == pos[st] for st in both])) if both else np.nan\n    have_tag = float(np.mean([v is not None for v in tag.values()]))\n\n    if LAT_FALLBACK == \"on\":\n        use = True\n    elif LAT_FALLBACK == \"off\":\n        use = False\n    else:\n        use = bool(both) and np.isfinite(agree) and agree >= LAT_MIN_AGREEMENT\n\n    side = {st: (tag[st] or (pos[st] if use else None)) for st in tag}\n    covered = float(np.mean([v is not None for v in side.values()]))\n    info = {\"tag_coverage\": have_tag, \"agreement\": agree, \"n_compared\": len(both),\n            \"fallback_used\": use, \"final_coverage\": covered}\n    log(f\"laterality: tag on {have_tag:.1%} of studies, x-sign agrees with it on \"\n        f\"{agree:.1%} of {len(both)} comparable studies, fallback \"\n        f\"{'enabled' if use else 'disabled'} -> {covered:.1%} normalised\")\n    return side, info","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def pick_slots(series_df, plane_map):\n    \"\"\"One series per slot per study.\n\n    Ties are broken toward the stack with the most slices: a thicker stack samples the\n    joint more densely, and the three-slice sampler below benefits from the margin.\n    \"\"\"\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            # fluid=None means \"do not condition on weighting\" - the public scheme,\n            # where the single provided flag stands in for both axes at once.\n            if fluid is not None:\n                sel &= (g[\"fluid\"] == fluid)\n            cand = g[sel]\n            if len(cand) == 0 and fluid is False:\n                # T1 slots are the scarcest; fall back to any non-fat-sat, non-fluid\n                # series in the plane before giving up on the slot entirely.\n                cand = g[(g[\"plane\"] == plane) & (~g[\"fatsat\"])]\n            if len(cand):\n                chosen[name] = cand.sort_values(\"n_slices\", ascending=False).iloc[0]\n        out[study] = chosen\n    return out\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def _natural_key(name):\n    \"\"\"Stable fallback for numeric filenames such as 1.dcm, 2.dcm, 10.dcm.\"\"\"\n    return tuple(int(x) if x.isdigit() else x.lower()\n                 for x in re.split(r\"(\\d+)\", str(name)))\n\n\ndef spatially_sorted_files(rec):\n    \"\"\"Sort a DICOM stack by patient-space position, then InstanceNumber.\n\n    Filename order is not an anatomical guarantee and can scramble the three-slice\n    pseudo-RGB inputs. The varying ImagePositionPatient axis identifies the stack\n    direction without assuming a scanner-specific orientation. Header failures fall\n    back to InstanceNumber and finally natural filename order.\n    \"\"\"\n    files, directory = list(rec[\"files\"]), rec[\"dir\"]\n    rows = []\n    for pos, name in enumerate(files):\n        ipp, instance = None, None\n        try:\n            ds = pydicom.dcmread(\n                os.path.join(directory, name), stop_before_pixels=True, force=True,\n                specific_tags=[\"ImagePositionPatient\", \"InstanceNumber\"],\n            )\n            raw_ipp = getattr(ds, \"ImagePositionPatient\", None)\n            if raw_ipp is not None and len(raw_ipp) >= 3:\n                candidate = np.asarray(raw_ipp[:3], dtype=np.float64)\n                if np.isfinite(candidate).all():\n                    ipp = candidate\n            raw_instance = getattr(ds, \"InstanceNumber\", None)\n            if raw_instance is not None:\n                instance = float(raw_instance)\n        except Exception:\n            pass\n        rows.append((name, ipp, instance, pos))\n\n    positioned = [r for r in rows if r[1] is not None]\n    if len(positioned) >= max(2, int(0.8 * len(rows))):\n        xyz = np.stack([r[1] for r in positioned])\n        axis = int(np.argmax(np.ptp(xyz, axis=0)))\n        fallback = float(np.nanmedian(xyz[:, axis]))\n        rows.sort(key=lambda r: (\n            float(r[1][axis]) if r[1] is not None else fallback,\n            r[2] if r[2] is not None else float(\"inf\"),\n            r[3],\n        ))\n    elif sum(r[2] is not None for r in rows) >= max(2, int(0.8 * len(rows))):\n        rows.sort(key=lambda r: (\n            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]\n\n\ndef read_slot(rec, n_slice=None, out_size=None):\n    \"\"\"`n_slice` physically spread slices from one series, at `out_size` pixels.\n\n    Returns float32 [n_slice, out, out] normalised per-series to its 1st-99th\n    percentile. Percentiles rather than min/max because MR intensity has no absolute\n    scale and a single bright vessel would otherwise compress the whole dynamic range.\n\n    Reading is the expensive half of this pipeline, so the caller reads once at the\n    largest configuration it needs and derives the smaller ones from the returned buffer\n    rather than re-reading.\n    \"\"\"\n    # was `N_SLICE`, which is never defined: a latent NameError on the default path.\n    n_slice = CACHE_SLICES if n_slice is None else n_slice\n    out_size = IMG if out_size is None else out_size\n    files, d, px = spatially_sorted_files(rec), rec[\"dir\"], rec[\"px\"]\n    n = len(files)\n    if n == 0:\n        return None\n    # Spread the samples over the central 60% of the stack: the outer slices of a knee\n    # series are mostly soft tissue outside the joint.\n    lo, hi = int(0.20 * (n - 1)), int(0.80 * (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\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 = np.zeros((out_size, out_size), dtype=np.float32)\n        planes.append(a)\n\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\n    # constant physical extent, then resize: PixelSpacing varies 3.4x across the corpus\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\n    lo_v, hi_v = np.percentile(vol, [1, 99])\n    vol = np.clip((vol - lo_v) / max(hi_v - lo_v, 1e-6), 0, 1)\n\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    # uint8, not float32. These buffers queue up between the reader threads and the\n    # encoder, and at this size a float32 slot-series is several megabytes. Intensity is\n    # already normalised into [0, 1] here, so eight bits cost nothing that a bilinear\n    # resize has not already cost, and the queue is a quarter the size.\n    return (t.squeeze(0) * 255).round().clamp(0, 255).to(torch.uint8)\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def normalise_laterality(img, plane, lat):\n    \"\"\"Map every knee onto a left-knee convention.\n\n    Coronal and axial views mirror under a horizontal flip. Sagittal stacks are not\n    mirror images of each other - the slice order runs medial-to-lateral in opposite\n    directions - so the channel order is reversed instead.\n    \"\"\"\n    if lat != \"R\":\n        return img\n    if plane in (\"Coronal\", \"Axial\"):\n        return torch.flip(img, dims=[-1])\n    return torch.flip(img, dims=[0])\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_cache(slot_map, plane_map, lat_map, tag):\n    \"\"\"Decode every (study, slot) once into an in-memory uint8 array.\n\n    Fine-tuning revisits the same pixels every epoch. Reading them from the mount each\n    time would make the epoch count a function of I/O rather than of learning, so they\n    are decoded once and held as bytes: intensity has already been normalised into\n    [0, 1], and eight bits cost nothing a bilinear resize has not already cost.\n\n    CACHE_SLICES positions are kept per slot, which the training loop reads as N_GROUP\n    groups of GROUP consecutive channels.\n    \"\"\"\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\n    jobs = [(st, k, plane, slot_map[st][name])\n            for st in studies\n            for k, (name, plane, _, _) in enumerate(SLOTS)\n            if name in slot_map[st]]\n    log(f\"{tag}: decoding {len(jobs)} slot-series\")\n\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(\n                    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,\n                                                          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    gc.collect()\n    return studies, cache, mask\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SlotHead(nn.Module):\n    \"\"\"Per-diagnosis attention over the slot embeddings of one study.\n\n    Each finding is read on particular sequences - cruciates sagittally, collateral\n    ligaments and the meniscal body coronally, patellar cartilage axially - so pooling\n    the slots identically would dilute the one that carries the evidence with the rest.\n\n    The aggregation is deliberately this simple. Richer alternatives were measured\n    against it and lost: with a study-level label there is no signal telling the model\n    which part of a study matters, so extra attention parameters have nothing to learn\n    from and spend their capacity fitting noise.\n    \"\"\"\n\n    def __init__(self, dim, n_slot, n_out, hidden=256, p=0.2):\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\n        # A weak prior keeps noisy study-level supervision from learning obviously\n        # implausible view assignments. A value of 0.55 is only a 1.7x softmax tilt,\n        # so the learned query can override it when the data disagrees.\n        prior = torch.zeros(n_out, n_slot)\n        if SLOT_SCHEME == \"recovered\" and n_slot == 6 and n_out == len(TARGETS):\n            preferred = {\n                \"ACL\": (0, 3, 5), \"MCL\": (1, 4),\n                \"Medial Meniscus\": (0, 1, 3, 4),\n                \"Lateral Meniscus\": (0, 1, 3, 4),\n                \"Medial OA\": (1, 4, 5), \"Lateral OA\": (1, 4, 5),\n                \"PF OA\": (0, 2, 5), \"Effusion\": (0, 2),\n                \"Synovitis\": (0, 2), \"Baker's\": (0,),\n                \"Contusion\": (0, 1, 2), \"Fracture\": (0, 1, 2, 4, 5),\n            }\n            for target, slots in preferred.items():\n                prior[TARGETS.index(target), list(slots)] = 0.55\n        self.register_buffer(\"slot_prior\", prior)\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               + self.slot_prior.unsqueeze(0))\n        att = att.masked_fill(mask.unsqueeze(1) < 0.5, -1e4).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","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Model(nn.Module):\n    \"\"\"Encoder plus head, trained end to end.\n\n    A study arrives as a bag of slot images. The bag is flattened for the encoder and\n    folded back before the head, so the encoder never sees the study structure and the\n    head never sees pixels.\n    \"\"\"\n\n    def __init__(self, backbone, dim):\n        super().__init__()\n        self.backbone = backbone\n        self.head = SlotHead(dim, N_SLOT, len(TARGETS))\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):\n        B, S = imgs.shape[:2]\n        x = imgs.reshape(B * S, *imgs.shape[2:]).float().div_(255.0)\n        x = (x - self.mean) / self.std\n        out = self.backbone(pixel_values=x).last_hidden_state\n        patch = out[:, 1:]\n        # The mean is stable for diffuse disease; top-k pooling preserves focal high\n        # responses from tears, contusions and fractures that a global mean dilutes.\n        k = max(1, patch.shape[1] // 8)\n        focal = patch.topk(k, dim=1).values.mean(1)\n        feat = torch.cat([out[:, 0], patch.mean(1), focal], dim=1).reshape(B, S, -1)\n        return self.head(feat, mask)\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_model():\n    \"\"\"Load the encoder and open the last few blocks for training.\n\n    Only the last UNFREEZE_LAST blocks and the final norm are trainable. The early\n    blocks of a self-supervised transformer are generic edge and texture filters; there\n    is not enough supervision here to improve them and quite enough to damage them.\n    \"\"\"\n    from transformers import AutoModel\n    p = find_dinov2(BACKBONE_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 * 3\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 \"\n        f\"({trainable / 1e6:.1f}M params), feature dim {dim}\")\n    return Model(bb, dim)\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def take_group(cache_rows, g):\n    \"\"\"Slice GROUP consecutive channels out of the cached slices.\"\"\"\n    return cache_rows[:, :, g * GROUP:(g + 1) * GROUP]\n\n\n# `augment` is defined further down, after the reasoning for dropping the flip.\n\n\n@torch.no_grad()\ndef predict(model, cache, mask, idx, dev):\n    \"\"\"Average the logits over the groups of each slot.\n\n    Training sees one group at a time, which acts as augmentation along the stack;\n    inference averages over all of them, which is the same aggregation the frozen\n    pipeline used and the one that measured best.\n    \"\"\"\n    model.eval()\n    out = []\n    for b in range(0, len(idx), EVAL_BATCH):\n        sel = idx[b:b + EVAL_BATCH]\n        rows = torch.from_numpy(cache[sel]).to(dev)\n        m = torch.from_numpy(mask[sel]).to(dev)\n        acc = None\n        for g in range(N_GROUP):\n            with torch.autocast(\"cuda\", enabled=dev.type == \"cuda\"):\n                z = model(take_group(rows, g), m).float()\n            acc = z if acc is None else acc + z\n        out.append(torch.sigmoid(acc / N_GROUP).cpu().numpy())\n    return np.concatenate(out) if out else np.zeros((0, len(TARGETS)), np.float32)\n\n\ndef macro_auc(y, p):\n    from sklearn.metrics import roc_auc_score\n    return float(np.nanmean([roc_auc_score(y[:, j], p[:, j])\n                             if len(set(y[:, j])) > 1 else np.nan\n                             for j in range(y.shape[1])]))\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def require_cuda():\n    if not torch.cuda.is_available():\n        raise RuntimeError(\"CUDA GPU is required for this inference notebook.\")\n    name = torch.cuda.get_device_name(0)\n    cap = torch.cuda.get_device_capability(0)\n    arch = f\"sm_{cap[0]}{cap[1]}\"\n    supported = set(torch.cuda.get_arch_list())\n    log(f\"cuda device: {name}, capability {arch}; torch supports {sorted(supported)}\")\n    if arch not in supported:\n        raise RuntimeError(f\"Kaggle assigned {name} ({arch}), but this PyTorch build cannot run on it.\")\n    return torch.device(\"cuda\")\n\n\ndef load_bundle():\n    path = find_model_path()\n    log(f\"model bundle: {path}\")\n    try:\n        bundle = torch.load(path, map_location=\"cpu\", weights_only=False)\n    except TypeError:\n        bundle = torch.load(path, map_location=\"cpu\")\n    return bundle\n\n\ndef apply_bundle_config(bundle):\n    global TARGETS, SLOTS, N_SLOT, IMG, GROUP, N_GROUP, CACHE_SLICES, BACKBONE_VARIANT\n    TARGETS = list(bundle.get(\"targets\", TARGETS))\n    SLOTS = [tuple(x) for x in bundle.get(\"slots\", SLOTS)]\n    N_SLOT = len(SLOTS)\n    IMG = int(bundle.get(\"img\", IMG))\n    GROUP = int(bundle.get(\"group\", GROUP))\n    N_GROUP = int(bundle.get(\"n_group\", N_GROUP))\n    CACHE_SLICES = GROUP * N_GROUP\n    variant = str(bundle.get(\"model_variant\", \"dinov2-small\")).split(\"-\")[-1]\n    BACKBONE_VARIANT = \"base\" if variant == \"base\" else \"small\"\n    log(f\"bundle config: backbone={BACKBONE_VARIANT}, img={IMG}, groups={N_GROUP}, slots={N_SLOT}\")\n\n\ndef write_submission(rank_sum, n_models, st_te, test_df):\n    P = rank_sum / max(n_models, 1)\n    sub = pd.DataFrame(P, columns=TARGETS)\n    sub.insert(0, \"StudyInstanceUID\", st_te)\n    sub[\"StudyInstanceUID\"] = sub[\"StudyInstanceUID\"].astype(str)\n    out = test_df[[\"StudyInstanceUID\"]].merge(sub, on=\"StudyInstanceUID\", how=\"left\")\n    out[TARGETS] = out[TARGETS].fillna(0.5)\n    out.to_csv(\"submission.csv\", index=False)\n    return out\n\n\ndef main():\n    dev = require_cuda()\n    bundle = load_bundle()\n    apply_bundle_config(bundle)\n\n    test_df = pd.read_csv(ROOT / \"test.csv\")\n    test_df[\"StudyInstanceUID\"] = test_df[\"StudyInstanceUID\"].astype(str)\n    test_series = pd.read_csv(ROOT / \"test_series.csv\")\n    test_series[\"StudyInstanceUID\"] = test_series[\"StudyInstanceUID\"].astype(str)\n    test_series[\"SeriesInstanceUID\"] = test_series[\"SeriesInstanceUID\"].astype(str)\n    log(f\"test {test_df.shape}; test_series {test_series.shape}\")\n\n    plane_map = dict(zip(test_series[\"SeriesInstanceUID\"], test_series[\"Anatomical_Plane\"]))\n    log(\"header pass: test\")\n    hte = annotate(walk(\"test_series\"))\n    log(f\"  {len(hte)} test series\")\n    lat_te, lat_info_te = laterality_maps(hte)\n    log(f\"laterality info: {lat_info_te}\")\n    slots_te = pick_slots(hte, plane_map)\n    cov = pd.Series([len(v) for v in slots_te.values()]).describe()\n    log(f\"test slots per study: mean {cov['mean']:.2f} min {cov['min']:.0f} max {cov['max']:.0f}\")\n    st_te, Cte, Mte = build_cache(slots_te, plane_map, lat_te, \"test\")\n\n    fold_states = bundle.get(\"fold_states\", [])\n    if not fold_states:\n        raise ValueError(\"saved model bundle has no fold_states\")\n\n    rank_sum = np.zeros((len(st_te), len(TARGETS)), np.float64)\n    for k, fold in enumerate(fold_states, 1):\n        model = build_model().to(dev)\n        model.load_state_dict(fold[\"state_dict\"])\n        P = predict(model, Cte, Mte, np.arange(len(st_te)), dev)\n        rank_sum += pd.DataFrame(P).rank(pct=True).values\n        log(f\"inferred fold {fold.get('fold', k - 1)} ({k}/{len(fold_states)})\")\n        del model\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n    sub = write_submission(rank_sum, len(fold_states), st_te, test_df)\n    log(f\"submission.csv {sub.shape}; models {len(fold_states)}\")\n    print(sub.head().to_string())\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"try:\n    main()\nexcept Exception:\n    traceback.print_exc()\n    t = pd.read_csv(find_root() / \"test.csv\")\n    for c in TARGETS:\n        t[c] = 0.5\n    t.to_csv(\"submission_fallback.csv\", index=False)\n    print(\"wrote submission_fallback.csv; re-raising so this is not submit-ready\")\n    raise\nlog(\"done\")\n","metadata":{},"outputs":[],"execution_count":null}]}