{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"title":"RSNA Knee Research Tony-Sakhawat V1"}},"nbformat_minor":5,"nbformat":4,"cells":[{"cell_type":"markdown","id":"apache-provenance-00","metadata":{},"source":"# RSNA Knee Research Candidate — Tony/Sakhawat V1 (public LB 0.836)\n\nSPDX-License-Identifier: Apache-2.0\n\nThis inference notebook is derived from two public Kaggle notebooks released under the Apache License 2.0:\n\n1. Tony Li, [RSNA Knee infer, Version 1](https://www.kaggle.com/code/tonylica/rsna-knee-infer?scriptVersionId=340809629), the checkpoint-compatible inference and model architecture. Pinned code SHA-256: `9c460f1f1387f2c7e9bfcf26d24fb10663df80f8cdab4e738c50cb398b46fad7`.\n2. Sakhawat Hossen, [Knee RSNA, Version 1](https://www.kaggle.com/code/sakhawathossen/knee-rsna?scriptVersionId=340853917), the overlapping-window TTA and hybrid rank ensemble. Pinned notebook SHA-256: `fc5465a8b7bc1e36f693e29acf2216f1225c2a8f4278aed0f089252138346e28`; pinned code SHA-256: `9ebc90c30036e11934da4e35d209033d6b927409c6cb717d0bf567daf36f3625`.\n\nChanges in this research candidate are limited to the Kaggle owner/slug, title, file name, this explicit provenance notice, an explicit dataset-version pin, and `nbformat_minor` 4→5 so the upstream cell IDs are valid under the notebook schema. No upstream cell ID, source, executable code, output, or execution count was changed. All executable code is byte-for-byte identical to Sakhawat Version 1. Retain this notice and the Apache-2.0 license when redistributing the notebook or a derivative.\n\nThe checkpoint is referenced, not redistributed: `tonylica/rsna2026-models`, dataset ID `11546462`, dataset version `2`, Kaggle source/version ID `18706996`, file `rsna_20260807_v1.pt`. Kaggle currently reports that dataset's license as `Unknown`; obtain an explicit license before copying, republishing, or distributing the checkpoint outside its Kaggle dataset. The DINOv2-small dependency is pinned to `metaresearch/dinov2/PyTorch/small/1` (model ID `986`, instance ID `3325`, version source ID `4533`) and is Apache-2.0 licensed.\n"},{"id":"096cd11a","cell_type":"markdown","source":"\n# RSNA Knee Abnormality Detection — Multi-View DINOv2 Inference\n\nThis notebook performs study-level inference from multi-sequence knee MRI series. Each study is reduced to a small set of clinically useful view/sequence slots, representative slices are converted into three-channel inputs, and a DINOv2 encoder produces features that are aggregated by a diagnosis-specific attention head.\n\n### What this version changes\n\n- Keeps the checkpoint-compatible model architecture and preprocessing contract intact.\n- Uses DICOM geometry rather than filename order when arranging slices.\n- Normalizes right/left knees into a consistent anatomical convention when the metadata supports it.\n- Adds **overlapping slice-window test-time averaging**: instead of evaluating only three disjoint 3-slice windows, the model can evaluate every consecutive 3-slice window in the cached 9-slice stack.\n- Uses a conservative **hybrid fold ensemble** that is dominated by fold-wise percentile ranks while adding a small rank of the mean-probability signal.\n- Keeps the entire inference path offline and Kaggle-compatible.\n\n> **Reproducibility / attribution:** this notebook is an independent cleanup and extension of the supplied baseline implementation. If the original baseline, checkpoint, or model dataset belongs to another Kaggle author, preserve the license and add the original source attribution before publishing. Reformatting or rewriting code does not remove attribution obligations.\n\nThe public leaderboard score can move up or down because the hidden evaluation set is different from any local sample. The two inference changes above are intentionally modest and are designed to reduce prediction variance without changing the learned model itself.\n","metadata":{}},{"id":"88780063","cell_type":"markdown","source":"\n## 1. Configuration\n\nThe defaults below match the attached model bundle. `OVERLAP_TTA=True` is the main experimental change. For a 9-slice cache and 3-channel inputs it evaluates 7 consecutive windows instead of only 3 disjoint windows.\n\n`FOLD_RANK_WEIGHT=0.90` keeps the final ensemble close to the original rank-averaging behavior. The remaining 10% comes from the rank of the probability-mean ensemble; this can break a few disagreements between folds without making calibration dominate an AUC-oriented submission.\n","metadata":{}},{"id":"2b82502b","cell_type":"code","source":"\nfrom __future__ import annotations\n\nimport gc\nimport os\nimport re\nimport time\nimport traceback\nfrom concurrent.futures import ThreadPoolExecutor\nfrom pathlib import Path\n\nfor _var in (\"OMP_NUM_THREADS\", \"OPENBLAS_NUM_THREADS\", \"MKL_NUM_THREADS\"):\n    os.environ.setdefault(_var, \"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\nSTART_TIME = time.time()\nSEED = 20260808\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\nTARGETS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\", \"Medial OA\",\n    \"Lateral OA\", \"PF OA\", \"Effusion\", \"Synovitis\", \"Baker's\",\n    \"Contusion\", \"Fracture\",\n]\n\nIMG = 224\nCROP_MM = 160.0\nGROUP = 3\nN_GROUP = 3\nCACHE_SLICES = GROUP * N_GROUP\n\nHDR_THREADS = 16\nPIX_THREADS = 12\nEVAL_BATCH = 12\nTIME_BUDGET = 8.0 * 3600\n\nBACKBONE_VARIANT = \"small\"\nUNFREEZE_LAST = 6\nMODEL_FILE = \"rsna_20260807_v1.pt\"\n\n# Inference improvements.\nOVERLAP_TTA = True\nLOGIT_POOL_WEIGHT = 0.90      # 1.0 reproduces sigmoid(mean(logits))\nFOLD_RANK_WEIGHT = 0.90       # dominant component of the final AUC-oriented blend\nFOLD_SCORE_POWER = 4.0        # used only when fold validation AUC is saved in the bundle\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(\n    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)\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\n\ndef log(message: str) -> None:\n    print(f\"[{time.time() - START_TIME:7.1f}s] {message}\", flush=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T19:36:07.285213Z","iopub.execute_input":"2026-08-07T19:36:07.285478Z","iopub.status.idle":"2026-08-07T19:36:07.296635Z","shell.execute_reply.started":"2026-08-07T19:36:07.285455Z","shell.execute_reply":"2026-08-07T19:36:07.295949Z"}},"outputs":[],"execution_count":null},{"id":"da360f55","cell_type":"markdown","source":"## 2. Locate competition data and attached weights","metadata":{}},{"id":"180c527d","cell_type":"code","source":"\ndef find_root() -> Path:\n    candidates = [\n        Path(\"/kaggle/input/competitions/rsna-knee-abnormality-detection\"),\n        Path(\"/kaggle/input/rsna-knee-abnormality-detection\"),\n        Path(\"data\"),\n        Path(\".\"),\n    ]\n    for candidate in candidates:\n        if (candidate / \"test.csv\").is_file() and (candidate / \"test_series\").is_dir():\n            return candidate\n\n    base = Path(\"/kaggle/input\")\n    if base.is_dir():\n        for level1 in sorted(p for p in base.iterdir() if p.is_dir()):\n            nested = [level1] + sorted(p for p in level1.iterdir() if p.is_dir())\n            for candidate in nested:\n                if (candidate / \"test.csv\").is_file() and (candidate / \"test_series\").is_dir():\n                    return candidate\n    raise FileNotFoundError(\"RSNA knee competition data was not found in the Kaggle input mount.\")\n\n\ndef find_dinov2(variant: str = \"small\") -> Path | None:\n    base = Path(\"/kaggle/input\")\n    if not base.is_dir():\n        return None\n\n    hits: list[Path] = []\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\n    for hit in hits:\n        if variant in str(hit).lower():\n            return hit\n    return hits[0] if hits else None\n\n\ndef find_model_path() -> Path:\n    direct = [\n        Path(\"/kaggle/input/datasets/tonylica/rsna2026-models\") / MODEL_FILE,\n        Path(MODEL_FILE),\n    ]\n    for candidate in direct:\n        if candidate.is_file():\n            return candidate\n\n    base = Path(\"/kaggle/input\")\n    if base.is_dir():\n        for candidate in base.rglob(MODEL_FILE):\n            if candidate.is_file():\n                return candidate\n    raise FileNotFoundError(f\"Required model bundle is missing: {MODEL_FILE}\")\n\n\nROOT = find_root()\nlog(f\"input root: {ROOT}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T19:36:07.300217Z","iopub.execute_input":"2026-08-07T19:36:07.300534Z","iopub.status.idle":"2026-08-07T19:36:07.31432Z","shell.execute_reply.started":"2026-08-07T19:36:07.300492Z","shell.execute_reply":"2026-08-07T19:36:07.313649Z"}},"outputs":[],"execution_count":null},{"id":"4083ebe6","cell_type":"markdown","source":"\n## 3. Read DICOM headers and recover sequence semantics\n\nMRI series names are not perfectly standardized across scanners. The header pass therefore combines sequence description, pulse timing, scan options, and the provided anatomical plane table to identify useful fat-suppressed fluid-sensitive and structural series.\n","metadata":{}},{"id":"1ee67c96","cell_type":"code","source":"\nHDR_TAGS = [\n    \"SeriesDescription\", \"SequenceName\", \"ScanOptions\", \"ScanningSequence\",\n    \"RepetitionTime\", \"EchoTime\", \"Laterality\", \"ImageLaterality\",\n    \"ImagePositionPatient\", \"PixelSpacing\", \"Rows\", \"Columns\",\n    \"RescaleSlope\", \"RescaleIntercept\",\n]\n\n\ndef probe_series(item):\n    split, study_uid, series_uid, directory = item\n    row = {\n        \"split\": split,\n        \"StudyInstanceUID\": study_uid,\n        \"SeriesInstanceUID\": series_uid,\n        \"dir\": directory,\n    }\n    try:\n        files = sorted(e.name for e in os.scandir(directory) if e.name.endswith(\".dcm\"))\n        row[\"files\"] = files\n        row[\"n_slices\"] = len(files)\n        if not files:\n            return row\n\n        middle = files[len(files) // 2]\n        ds = pydicom.dcmread(\n            os.path.join(directory, middle),\n            stop_before_pixels=True,\n            force=True,\n        )\n        for tag in HDR_TAGS:\n            value = getattr(ds, tag, None)\n            if value is None:\n                row[tag] = None\n            elif isinstance(value, (list, tuple)) or type(value).__name__ == \"MultiValue\":\n                row[tag] = \"|\".join(str(x) for x in value)\n            else:\n                row[tag] = str(value)\n    except Exception as exc:\n        row[\"err\"] = str(exc)[:160]\n    return row\n\n\ndef scan_series(split: str) -> pd.DataFrame:\n    base = ROOT / split\n    if not base.is_dir():\n        return pd.DataFrame()\n\n    jobs = []\n    for study in os.scandir(base):\n        if not study.is_dir():\n            continue\n        for series in os.scandir(study.path):\n            if series.is_dir():\n                jobs.append((split, study.name, series.name, series.path))\n\n    with ThreadPoolExecutor(max_workers=HDR_THREADS) as pool:\n        rows = list(pool.map(probe_series, jobs))\n    return pd.DataFrame(rows)\n\n\ndef annotate_sequences(df: pd.DataFrame) -> pd.DataFrame:\n    df = df.copy()\n    desc = (df[\"SeriesDescription\"].fillna(\"\") + \" \" + df[\"SequenceName\"].fillna(\"\"))\n    desc = desc.str.lower().str.replace(_SEP, \" \", regex=True)\n\n    scan_options = df[\"ScanOptions\"].fillna(\"\").str.upper().str.split(\"|\")\n    option_fatsat = scan_options.apply(\n        lambda tokens: any(token.strip() in FATSAT_OPTS for token in tokens)\n    )\n    df[\"fatsat\"] = desc.str.contains(_FATSAT_RX) | option_fatsat\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\n    named_t1 = desc.str.contains(_T1_RX)\n    named_t2 = desc.str.contains(_T2_RX)\n    named_pd = desc.str.contains(_PD_RX)\n\n    df[\"weight\"] = np.where(\n        named_t1 & ~named_t2 & ~named_pd, \"T1\",\n        np.where(\n            named_t2 & ~named_pd, \"T2\",\n            np.where(\n                named_pd, \"PD\",\n                np.where(\n                    gre, \"GRE\",\n                    np.where(tr < 800, \"T1\", np.where(te > 60, \"T2\", np.where(tr >= 800, \"PD\", \"UNK\"))),\n                ),\n            ),\n        ),\n    )\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    )\n    return df\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T19:36:09.044635Z","iopub.execute_input":"2026-08-07T19:36:09.045126Z","iopub.status.idle":"2026-08-07T19:36:09.058628Z","shell.execute_reply.started":"2026-08-07T19:36:09.0451Z","shell.execute_reply":"2026-08-07T19:36:09.057797Z"}},"outputs":[],"execution_count":null},{"id":"3ff3b83d","cell_type":"markdown","source":"## 4. Laterality normalization and slot selection","metadata":{}},{"id":"445642ad","cell_type":"code","source":"\ndef _tag_side(group: pd.DataFrame) -> str | None:\n    values = [str(x).strip().upper() for x in group[\"Laterality\"].dropna()]\n    if \"ImageLaterality\" in group.columns:\n        values += [str(x).strip().upper() for x in group[\"ImageLaterality\"].dropna()]\n    values = [x[0] for x in values if x and x[0] in (\"L\", \"R\")]\n    return values[0] if values else None\n\n\ndef _position_side(group: pd.DataFrame) -> str | None:\n    xs = []\n    for raw in group.get(\"ImagePositionPatient\", pd.Series(dtype=object)).dropna():\n        try:\n            xs.append(float(str(raw).split(\"|\")[0]))\n        except Exception:\n            pass\n    if not xs:\n        return None\n\n    median_x = float(np.median(xs))\n    if abs(median_x) < LAT_MIN_OFFSET_MM:\n        return None\n    return \"R\" if median_x < 0 else \"L\"  # DICOM patient coordinates use LPS.\n\n\ndef laterality_maps(headers: pd.DataFrame):\n    tagged, positioned = {}, {}\n    for study_uid, group in headers.groupby(\"StudyInstanceUID\"):\n        tagged[study_uid] = _tag_side(group)\n        positioned[study_uid] = _position_side(group)\n\n    comparable = [s for s in tagged if tagged[s] and positioned[s]]\n    agreement = (\n        float(np.mean([tagged[s] == positioned[s] for s in comparable]))\n        if comparable else np.nan\n    )\n    tag_coverage = float(np.mean([v is not None for v in tagged.values()]))\n\n    if LAT_FALLBACK == \"on\":\n        use_position = True\n    elif LAT_FALLBACK == \"off\":\n        use_position = False\n    else:\n        use_position = bool(comparable) and np.isfinite(agreement) and agreement >= LAT_MIN_AGREEMENT\n\n    resolved = {\n        study_uid: (tagged[study_uid] or (positioned[study_uid] if use_position else None))\n        for study_uid in tagged\n    }\n    final_coverage = float(np.mean([v is not None for v in resolved.values()]))\n\n    info = {\n        \"tag_coverage\": tag_coverage,\n        \"agreement\": agreement,\n        \"n_compared\": len(comparable),\n        \"fallback_used\": use_position,\n        \"final_coverage\": final_coverage,\n    }\n    log(\n        f\"laterality: tag={tag_coverage:.1%}, agreement={agreement:.1%} \"\n        f\"on {len(comparable)} comparable studies, final={final_coverage:.1%}\"\n    )\n    return resolved, info\n\n\ndef pick_slots(series_df: pd.DataFrame, plane_map: dict) -> dict:\n    table = series_df.copy()\n    table[\"plane\"] = table[\"SeriesInstanceUID\"].map(plane_map)\n\n    selected = {}\n    for study_uid, group in table.groupby(\"StudyInstanceUID\"):\n        study_slots = {}\n        for slot_name, plane, fluid, fatsat in SLOTS:\n            keep = (group[\"plane\"] == plane) & (group[\"fatsat\"] == fatsat)\n            if fluid is not None:\n                keep &= group[\"fluid\"] == fluid\n\n            candidates = group[keep]\n            if len(candidates) == 0 and fluid is False:\n                candidates = group[(group[\"plane\"] == plane) & (~group[\"fatsat\"])]\n\n            if len(candidates):\n                study_slots[slot_name] = candidates.sort_values(\n                    \"n_slices\", ascending=False\n                ).iloc[0]\n        selected[study_uid] = study_slots\n    return selected\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T19:36:14.742937Z","iopub.execute_input":"2026-08-07T19:36:14.743326Z","iopub.status.idle":"2026-08-07T19:36:14.759006Z","shell.execute_reply.started":"2026-08-07T19:36:14.743301Z","shell.execute_reply":"2026-08-07T19:36:14.758052Z"}},"outputs":[],"execution_count":null},{"id":"5b49094a","cell_type":"markdown","source":"\n## 5. Spatial sorting, physical crop, and study cache\n\nThe image pipeline samples the central 60% of each selected MRI stack. A constant physical field of view is used when pixel spacing is available, then every slice is resized to the network input resolution and normalized using per-series robust percentiles.\n","metadata":{}},{"id":"2f6a4c96","cell_type":"code","source":"\ndef _natural_key(name):\n    return tuple(int(x) if x.isdigit() else x.lower() for x in re.split(r\"(\\d+)\", str(name)))\n\n\ndef spatially_sorted_files(record) -> list[str]:\n    files = list(record[\"files\"])\n    directory = record[\"dir\"]\n    rows = []\n\n    for original_pos, name in enumerate(files):\n        ipp, instance = None, None\n        try:\n            ds = pydicom.dcmread(\n                os.path.join(directory, name),\n                stop_before_pixels=True,\n                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, original_pos))\n\n    positioned = [row for row in rows if row[1] is not None]\n    threshold = max(2, int(0.8 * len(rows)))\n\n    if len(positioned) >= threshold:\n        xyz = np.stack([row[1] for row in positioned])\n        varying_axis = int(np.argmax(np.ptp(xyz, axis=0)))\n        fallback = float(np.nanmedian(xyz[:, varying_axis]))\n        rows.sort(\n            key=lambda row: (\n                float(row[1][varying_axis]) if row[1] is not None else fallback,\n                row[2] if row[2] is not None else float(\"inf\"),\n                row[3],\n            )\n        )\n    elif sum(row[2] is not None for row in rows) >= threshold:\n        rows.sort(key=lambda row: (row[2] if row[2] is not None else float(\"inf\"), row[3]))\n    else:\n        rows.sort(key=lambda row: _natural_key(row[0]))\n\n    return [row[0] for row in rows]\n\n\ndef read_slot(record, n_slice: int | None = None, out_size: int | None = None):\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\n    files = spatially_sorted_files(record)\n    directory = record[\"dir\"]\n    px = record[\"px\"]\n    n_files = len(files)\n    if n_files == 0:\n        return None\n\n    low_idx = int(0.20 * (n_files - 1))\n    high_idx = int(0.80 * (n_files - 1))\n    if high_idx > low_idx:\n        sample_idx = np.unique(np.linspace(low_idx, high_idx, n_slice).astype(int))\n    else:\n        sample_idx = np.array([n_files // 2])\n    while len(sample_idx) < n_slice:\n        sample_idx = np.append(sample_idx, sample_idx[-1])\n\n    planes = []\n    for index in sample_idx[:n_slice]:\n        try:\n            ds = pydicom.dcmread(os.path.join(directory, files[int(index)]), force=True)\n            image = ds.pixel_array.astype(np.float32)\n            slope = float(getattr(ds, \"RescaleSlope\", 1) or 1)\n            intercept = float(getattr(ds, \"RescaleIntercept\", 0) or 0)\n            image = image * slope + intercept\n        except Exception:\n            image = np.zeros((out_size, out_size), dtype=np.float32)\n        planes.append(image)\n\n    shape = planes[0].shape\n    planes = [p if p.shape == shape else np.zeros(shape, np.float32) for p in planes]\n    volume = np.stack(planes)\n\n    if px and np.isfinite(px) and px > 0:\n        desired = int(round(CROP_MM / px))\n        h, w = shape\n        if 16 < desired < min(h, w):\n            cy, cx = h // 2, w // 2\n            half = desired // 2\n            volume = volume[:, max(0, cy - half):cy + half, max(0, cx - half):cx + half]\n\n    p01, p99 = np.percentile(volume, [1, 99])\n    volume = np.clip((volume - p01) / max(p99 - p01, 1e-6), 0, 1)\n\n    tensor = torch.from_numpy(np.ascontiguousarray(volume)).unsqueeze(0)\n    tensor = F.interpolate(\n        tensor, size=(out_size, out_size), mode=\"bilinear\", align_corners=False\n    )\n    return (tensor.squeeze(0) * 255).round().clamp(0, 255).to(torch.uint8)\n\n\ndef normalise_laterality(image: torch.Tensor, plane: str, laterality: str | None):\n    if laterality != \"R\":\n        return image\n    if plane in (\"Coronal\", \"Axial\"):\n        return torch.flip(image, dims=[-1])\n    return torch.flip(image, dims=[0])\n\n\ndef build_cache(slot_map: dict, laterality_map: dict, tag: str):\n    studies = sorted(slot_map)\n    study_index = {study_uid: i for i, study_uid in enumerate(studies)}\n\n    cache = np.zeros((len(studies), N_SLOT, CACHE_SLICES, IMG, IMG), dtype=np.uint8)\n    mask = np.zeros((len(studies), N_SLOT), dtype=np.float32)\n    log(f\"{tag}: cache {cache.shape} = {cache.nbytes / 1024**3:.2f} GB\")\n\n    jobs = [\n        (study_uid, slot_idx, plane, slot_map[study_uid][slot_name])\n        for study_uid in studies\n        for slot_idx, (slot_name, plane, _, _) in enumerate(SLOTS)\n        if slot_name in slot_map[study_uid]\n    ]\n    log(f\"{tag}: decoding {len(jobs)} selected series\")\n\n    chunk_size = 512\n    completed = 0\n    with ThreadPoolExecutor(max_workers=PIX_THREADS) as pool:\n        for start in range(0, len(jobs), chunk_size):\n            block = jobs[start:start + chunk_size]\n            decoded = pool.map(lambda job: read_slot(job[3], CACHE_SLICES, IMG), block)\n\n            for (study_uid, slot_idx, plane, _), image in zip(block, decoded):\n                completed += 1\n                if image is None:\n                    continue\n                image = normalise_laterality(image, plane, laterality_map.get(study_uid))\n                cache[study_index[study_uid], slot_idx] = image.numpy()\n                mask[study_index[study_uid], slot_idx] = 1.0\n\n            if completed % 4096 < chunk_size:\n                log(f\"  {tag}: {completed}/{len(jobs)} series decoded\")\n            if time.time() - START_TIME > TIME_BUDGET:\n                log(f\"  {tag}: time budget reached during DICOM decoding\")\n                break\n\n    gc.collect()\n    return studies, cache, mask\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T19:36:16.96876Z","iopub.execute_input":"2026-08-07T19:36:16.969173Z","iopub.status.idle":"2026-08-07T19:36:16.992443Z","shell.execute_reply.started":"2026-08-07T19:36:16.969144Z","shell.execute_reply":"2026-08-07T19:36:16.991671Z"}},"outputs":[],"execution_count":null},{"id":"de4d38e4","cell_type":"markdown","source":"\n## 6. Checkpoint-compatible model\n\nThe module names and tensor shapes in this section intentionally match the saved checkpoint. The head learns a separate attention distribution over MRI slots for every diagnosis, while the DINOv2 representation combines the CLS token, average patch representation, and a focal top-k patch summary.\n","metadata":{}},{"id":"847034de","cell_type":"code","source":"\nclass SlotHead(nn.Module):\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        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),\n                \"MCL\": (1, 4),\n                \"Medial Meniscus\": (0, 1, 3, 4),\n                \"Lateral Meniscus\": (0, 1, 3, 4),\n                \"Medial OA\": (1, 4, 5),\n                \"Lateral OA\": (1, 4, 5),\n                \"PF OA\": (0, 2, 5),\n                \"Effusion\": (0, 2),\n                \"Synovitis\": (0, 2),\n                \"Baker's\": (0,),\n                \"Contusion\": (0, 1, 2),\n                \"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        attention = (\n            torch.einsum(\"bsh,oh->bos\", h, self.query) / self.hidden**0.5\n            + self.slot_prior.unsqueeze(0)\n        )\n        attention = attention.masked_fill(mask.unsqueeze(1) < 0.5, -1e4).softmax(-1)\n        context = self.drop(torch.einsum(\"bos,bsh->boh\", attention, h))\n        return (context * self.out.weight.unsqueeze(0)).sum(-1) + self.out.bias\n\n\nclass Model(nn.Module):\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        batch, slots = imgs.shape[:2]\n        x = imgs.reshape(batch * slots, *imgs.shape[2:]).float().div_(255.0)\n        x = (x - self.mean) / self.std\n\n        encoded = self.backbone(pixel_values=x).last_hidden_state\n        patches = encoded[:, 1:]\n        k = max(1, patches.shape[1] // 8)\n        focal = patches.topk(k, dim=1).values.mean(1)\n        features = torch.cat([encoded[:, 0], patches.mean(1), focal], dim=1)\n        features = features.reshape(batch, slots, -1)\n        return self.head(features, mask)\n\n\ndef build_model():\n    from transformers import AutoModel\n\n    backbone_path = find_dinov2(BACKBONE_VARIANT)\n    if backbone_path is None:\n        raise FileNotFoundError(\"Offline DINOv2 weights are not attached to this Kaggle notebook.\")\n\n    backbone = AutoModel.from_pretrained(str(backbone_path))\n    n_layers = len(backbone.encoder.layer)\n\n    for parameter in backbone.parameters():\n        parameter.requires_grad = False\n    for block in backbone.encoder.layer[max(0, n_layers - UNFREEZE_LAST):]:\n        for parameter in block.parameters():\n            parameter.requires_grad = True\n    for parameter in backbone.layernorm.parameters():\n        parameter.requires_grad = True\n\n    feature_dim = backbone.config.hidden_size * 3\n    trainable = sum(p.numel() for p in backbone.parameters() if p.requires_grad)\n    log(\n        f\"backbone: {n_layers} blocks, last {UNFREEZE_LAST} trainable \"\n        f\"({trainable / 1e6:.1f}M params), feature dim={feature_dim}\"\n    )\n    return Model(backbone, feature_dim)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T19:36:20.313992Z","iopub.execute_input":"2026-08-07T19:36:20.314451Z","iopub.status.idle":"2026-08-07T19:36:20.329167Z","shell.execute_reply.started":"2026-08-07T19:36:20.314423Z","shell.execute_reply":"2026-08-07T19:36:20.328331Z"}},"outputs":[],"execution_count":null},{"id":"756d88d8","cell_type":"markdown","source":"\n## 7. Enhanced test-time aggregation\n\nThe baseline cache contains 9 slices per slot and feeds them to the model as three non-overlapping triplets. Here, the default inference evaluates all 7 consecutive triplets: `[0:3]`, `[1:4]`, …, `[6:9]`. This reuses the same decoded cache, so there is no additional DICOM I/O.\n\nFor each fold, most of the prediction comes from `sigmoid(mean(logits))`, matching the original behavior. A small mean-probability component is included to reduce sensitivity to one extreme window.\n","metadata":{}},{"id":"ca0446a5","cell_type":"code","source":"\ndef require_cuda() -> torch.device:\n    if not torch.cuda.is_available():\n        raise RuntimeError(\"A CUDA GPU is required for this inference notebook.\")\n\n    device_name = torch.cuda.get_device_name(0)\n    capability = torch.cuda.get_device_capability(0)\n    arch = f\"sm_{capability[0]}{capability[1]}\"\n    supported = set(torch.cuda.get_arch_list())\n    log(f\"cuda: {device_name}, capability={arch}\")\n    if arch not in supported:\n        raise RuntimeError(f\"The installed PyTorch build does not support the assigned GPU ({arch}).\")\n\n    torch.backends.cuda.matmul.allow_tf32 = True\n    try:\n        torch.set_float32_matmul_precision(\"high\")\n    except Exception:\n        pass\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\n    TARGETS = list(bundle.get(\"targets\", TARGETS))\n    SLOTS = [tuple(slot) for slot 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\n    variant = str(bundle.get(\"model_variant\", \"dinov2-small\")).split(\"-\")[-1]\n    BACKBONE_VARIANT = \"base\" if variant == \"base\" else \"small\"\n    log(\n        f\"bundle config: backbone={BACKBONE_VARIANT}, img={IMG}, \"\n        f\"group={GROUP}, cached_slices={CACHE_SLICES}, slots={N_SLOT}\"\n    )\n\n\ndef window_starts() -> list[int]:\n    if OVERLAP_TTA and CACHE_SLICES >= GROUP:\n        return list(range(CACHE_SLICES - GROUP + 1))\n    return [group_idx * GROUP for group_idx in range(N_GROUP)]\n\n\ndef take_window(cache_rows: torch.Tensor, start: int) -> torch.Tensor:\n    return cache_rows[:, :, start:start + GROUP]\n\n\n@torch.no_grad()\ndef predict_fold(model, cache, mask, indices, device):\n    model.eval()\n    starts = window_starts()\n    outputs = []\n\n    for batch_start in range(0, len(indices), EVAL_BATCH):\n        batch_idx = indices[batch_start:batch_start + EVAL_BATCH]\n        rows = torch.from_numpy(cache[batch_idx]).to(device, non_blocking=True)\n        slot_mask = torch.from_numpy(mask[batch_idx]).to(device, non_blocking=True)\n\n        logit_sum = None\n        probability_sum = None\n        for start in starts:\n            with torch.autocast(\"cuda\", enabled=device.type == \"cuda\"):\n                logits = model(take_window(rows, start), slot_mask).float()\n            probabilities = torch.sigmoid(logits)\n            logit_sum = logits if logit_sum is None else logit_sum + logits\n            probability_sum = probabilities if probability_sum is None else probability_sum + probabilities\n\n        mean_logits = logit_sum / len(starts)\n        mean_probabilities = probability_sum / len(starts)\n        pooled = (\n            LOGIT_POOL_WEIGHT * torch.sigmoid(mean_logits)\n            + (1.0 - LOGIT_POOL_WEIGHT) * mean_probabilities\n        )\n        outputs.append(pooled.cpu().numpy())\n\n    if not outputs:\n        return np.zeros((0, len(TARGETS)), dtype=np.float32)\n    return np.concatenate(outputs, axis=0)\n\n\ndef percentile_rank(matrix: np.ndarray) -> np.ndarray:\n    return pd.DataFrame(matrix).rank(axis=0, pct=True, method=\"average\").to_numpy(dtype=np.float64)\n\n\ndef extract_fold_quality(fold: dict) -> float | None:\n    # Only accept explicit AUC-like metadata; generic loss values are intentionally ignored.\n    for key in (\"val_auc\", \"macro_auc\", \"auc\", \"best_auc\", \"valid_auc\"):\n        value = fold.get(key)\n        if isinstance(value, (int, float, np.number)) and np.isfinite(value) and 0.5 <= float(value) <= 1.0:\n            return float(value)\n    return None\n\n\ndef fold_weights(fold_states: list[dict]) -> np.ndarray:\n    quality = [extract_fold_quality(fold) for fold in fold_states]\n    if any(value is None for value in quality):\n        log(\"fold validation AUC metadata not available -> equal fold weights\")\n        return np.ones(len(fold_states), dtype=np.float64)\n\n    raw = np.asarray(quality, dtype=np.float64) ** FOLD_SCORE_POWER\n    raw /= raw.mean()\n    log(\"fold weights from saved validation AUC: \" + \", \".join(f\"{w:.3f}\" for w in raw))\n    return raw\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T19:36:23.064725Z","iopub.execute_input":"2026-08-07T19:36:23.065166Z","iopub.status.idle":"2026-08-07T19:36:23.080571Z","shell.execute_reply.started":"2026-08-07T19:36:23.065137Z","shell.execute_reply":"2026-08-07T19:36:23.079807Z"}},"outputs":[],"execution_count":null},{"id":"7ef36b6b","cell_type":"markdown","source":"## 8. Run inference and write `submission.csv`","metadata":{}},{"id":"5907ce35","cell_type":"code","source":"\ndef write_submission(final_scores: np.ndarray, study_uids: list[str], test_df: pd.DataFrame):\n    prediction_table = pd.DataFrame(final_scores, columns=TARGETS)\n    prediction_table.insert(0, \"StudyInstanceUID\", study_uids)\n    prediction_table[\"StudyInstanceUID\"] = prediction_table[\"StudyInstanceUID\"].astype(str)\n\n    submission = test_df[[\"StudyInstanceUID\"]].merge(\n        prediction_table, on=\"StudyInstanceUID\", how=\"left\"\n    )\n    submission[TARGETS] = submission[TARGETS].fillna(0.5)\n    submission.to_csv(\"submission.csv\", index=False)\n    return submission\n\n\ndef main():\n    device = 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\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\n    log(\"reading test DICOM headers\")\n    headers = annotate_sequences(scan_series(\"test_series\"))\n    log(f\"header rows: {len(headers)}\")\n\n    laterality, lat_info = laterality_maps(headers)\n    log(f\"laterality diagnostics: {lat_info}\")\n\n    slot_map = pick_slots(headers, plane_map)\n    slot_counts = pd.Series([len(slots) for slots in slot_map.values()])\n    if len(slot_counts):\n        log(\n            f\"slots/study: mean={slot_counts.mean():.2f}, \"\n            f\"min={slot_counts.min():.0f}, max={slot_counts.max():.0f}\"\n        )\n\n    study_uids, cache, mask = build_cache(slot_map, laterality, \"test\")\n\n    fold_states = bundle.get(\"fold_states\", [])\n    if not fold_states:\n        raise ValueError(\"The saved model bundle contains no fold_states.\")\n\n    weights = fold_weights(fold_states)\n    total_weight = float(weights.sum())\n    rank_accumulator = np.zeros((len(study_uids), len(TARGETS)), dtype=np.float64)\n    probability_accumulator = np.zeros_like(rank_accumulator)\n\n    log(f\"TTA windows per fold: {len(window_starts())} -> {window_starts()}\")\n\n    all_indices = np.arange(len(study_uids))\n    for fold_number, (fold, weight) in enumerate(zip(fold_states, weights), start=1):\n        model = build_model().to(device)\n        model.load_state_dict(fold[\"state_dict\"], strict=True)\n\n        probabilities = predict_fold(model, cache, mask, all_indices, device)\n        rank_accumulator += weight * percentile_rank(probabilities)\n        probability_accumulator += weight * probabilities\n\n        fold_id = fold.get(\"fold\", fold_number - 1)\n        log(f\"inferred fold {fold_id} ({fold_number}/{len(fold_states)}), weight={weight:.3f}\")\n\n        del model\n        gc.collect()\n        torch.cuda.empty_cache()\n\n    mean_fold_rank = rank_accumulator / total_weight\n    mean_probability = probability_accumulator / total_weight\n    probability_rank = percentile_rank(mean_probability)\n\n    final_scores = (\n        FOLD_RANK_WEIGHT * mean_fold_rank\n        + (1.0 - FOLD_RANK_WEIGHT) * probability_rank\n    )\n\n    submission = write_submission(final_scores, study_uids, test_df)\n    log(f\"submission.csv written: {submission.shape}\")\n    print(submission.head().to_string(index=False))\n    return submission\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T19:36:26.48876Z","iopub.execute_input":"2026-08-07T19:36:26.489491Z","iopub.status.idle":"2026-08-07T19:36:26.501198Z","shell.execute_reply.started":"2026-08-07T19:36:26.489456Z","shell.execute_reply":"2026-08-07T19:36:26.500598Z"}},"outputs":[],"execution_count":null},{"id":"0813000b","cell_type":"code","source":"\ntry:\n    submission = main()\nexcept Exception:\n    traceback.print_exc()\n    fallback = pd.read_csv(find_root() / \"test.csv\")\n    for target in TARGETS:\n        fallback[target] = 0.5\n    fallback.to_csv(\"submission_fallback.csv\", index=False)\n    print(\"A fallback file was written for debugging; the notebook is re-raising the error.\")\n    raise\n\nlog(\"done\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-07T19:36:29.085376Z","iopub.execute_input":"2026-08-07T19:36:29.085763Z"}},"outputs":[],"execution_count":null},{"id":"519f19e0-5440-4b94-a5b5-d2e326969612","cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}