{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.x"},"rsna_submission_version":"v1-2026-08-28"},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n# RSNA Knee Abnormality Detection — Submission Inference v1\n\nThis notebook is **inference-only**.\n\nIt expects:\n- the official RSNA competition input;\n- the **output of the successful training notebook** attached as an input, containing:\n\n```text\nrsna_25d_radimagenet_resnet50_fold0.pt\nrsna_25d_radimagenet_resnet50_fold1.pt\nrsna_25d_radimagenet_resnet50_fold2.pt\nrsna_25d_radimagenet_resnet50_fold3.pt\nrsna_25d_radimagenet_resnet50_fold4.pt\n```\n\nFlow:\n\n```text\nhidden test DICOM\n    ↓\nsame 4-slot + 2.5D preprocessing\n    ↓\nload fold0..fold4\n    ↓\npredict 3 z-windows/fold\n    ↓\nrank-mean ensemble\n    ↓\n/kaggle/working/submission.csv\n```\n\nIt does **not train**, and it does **not need reports, pseudo-labels, or the original RadImageNet source checkpoint**.\nThe five saved fold checkpoints already contain the complete trained model state.\n","metadata":{}},{"cell_type":"code","source":"\nfrom __future__ import annotations\n\nimport gc, math, time, warnings\nfrom pathlib import Path\nfrom concurrent.futures import ThreadPoolExecutor, as_completed\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nwarnings.filterwarnings(\"ignore\")\n\nTARGETS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n    \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\",\n    \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\",\n]\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"device:\", DEVICE)\nif torch.cuda.is_available():\n    print(torch.cuda.get_device_name(0))\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ---------------- INFERENCE CONFIG ----------------\nIMG_SIZE = 224\nCROP_MM = 130.0\nN_SAMPLE_SLICES = 9\nN_WINDOWS = 3\n\nBATCH_STUDIES = 8\nCACHE_THREADS = 12\n\nSLOTS = [\n    (\"SAG_FS\",     \"Sagittal\", \"fluid_fs\"),\n    (\"COR_FS\",     \"Coronal\",  \"fluid_fs\"),\n    (\"AX_FS\",      \"Axial\",    \"fluid_fs\"),\n    (\"SAG_STRUCT\", \"Sagittal\", \"struct\"),\n]\nN_SLOTS = len(SLOTS)\n\n# Optional exact mounted training-output directory.\n# Leave None for automatic discovery.\nWEIGHTS_DIR = None\n\n# Temporary cache is intentionally NOT in /kaggle/working.\nCACHE_DIR = Path(\"/kaggle/temp/rsna_knee_submission_cache\")\nCACHE_DIR.mkdir(parents=True, exist_ok=True)\n\nprint(\"slots:\", [x[0] for x in SLOTS])\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n## 1. Locate hidden-test data and saved fold weights\n\nThe checkpoint search skips the official competition tree, so it will not recursively scan the huge DICOM input.\n","metadata":{}},{"cell_type":"code","source":"\ndef find_root():\n    candidates = [\n        Path(\"/kaggle/input/competitions/rsna-knee-abnormality-detection\"),\n        Path(\"/kaggle/input/rsna-knee-abnormality-detection\"),\n    ]\n    for p in candidates:\n        if (p / \"test.csv\").is_file() and (p / \"test_series\").is_dir():\n            return p\n    for p in Path(\"/kaggle/input\").glob(\"*\"):\n        if p.is_dir() and (p / \"test.csv\").is_file() and (p / \"test_series\").is_dir():\n            return p\n    raise FileNotFoundError(\"Cannot locate official RSNA competition input.\")\n\nROOT = find_root()\ntest_df = pd.read_csv(ROOT / \"test.csv\")\ntest_series = pd.read_csv(ROOT / \"test_series.csv\")\n\nprint(\"ROOT:\", ROOT)\nprint(\"test:\", test_df.shape)\nprint(\"test series:\", test_series.shape)\n\n\ndef discover_fold_weights(explicit_dir=None):\n    pattern = \"rsna_25d_radimagenet_resnet50_fold*.pt\"\n\n    if explicit_dir is not None:\n        d = Path(explicit_dir)\n        hits = sorted(d.glob(pattern))\n        if len(hits) != 5:\n            raise FileNotFoundError(\n                f\"Expected 5 fold weights in {d}, found {len(hits)}\"\n            )\n        return hits\n\n    hits = []\n    base = Path(\"/kaggle/input\")\n\n    # Search only non-competition mounted inputs.\n    for d in base.iterdir():\n        if not d.is_dir() or d.name == \"competitions\":\n            continue\n        for pat in (\n            pattern,\n            f\"*/{pattern}\",\n            f\"*/*/{pattern}\",\n            f\"*/*/*/{pattern}\",\n        ):\n            hits.extend(d.glob(pat))\n\n    hits = sorted(set(hits), key=lambda p: str(p))\n\n    by_fold = {}\n    for p in hits:\n        try:\n            fold = int(p.stem.rsplit(\"fold\", 1)[1])\n        except Exception:\n            continue\n        if fold not in by_fold or len(str(p)) < len(str(by_fold[fold])):\n            by_fold[fold] = p\n\n    missing = [f for f in range(5) if f not in by_fold]\n    if missing:\n        print(\"candidate checkpoints found:\")\n        for p in hits:\n            print(\" \", p)\n        raise FileNotFoundError(\n            \"Could not find all 5 fold checkpoints. \"\n            f\"Missing folds: {missing}. \"\n            \"Attach the successful training notebook Output Files as an input, \"\n            \"or set WEIGHTS_DIR explicitly.\"\n        )\n\n    return [by_fold[f] for f in range(5)]\n\n\nFOLD_WEIGHTS = discover_fold_weights(WEIGHTS_DIR)\n\nprint(\"\\nFold weights:\")\nfor i, p in enumerate(FOLD_WEIGHTS):\n    print(f\" fold {i}: {p} ({p.stat().st_size / 1024**2:.1f} MB)\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Select the same four MRI slots as training","metadata":{}},{"cell_type":"code","source":"\ndef series_priority(row, kind):\n    fluid = int(row[\"Fluid_Sensitive\"])\n    fs = int(row[\"Fat_Suppression\"])\n    if kind == \"fluid_fs\":\n        return 100 * fluid + 40 * fs + 20 * (fluid and fs)\n    if kind == \"struct\":\n        return 80 * (not fs) + 20 * (not fluid) + 5 * fluid\n    raise ValueError(kind)\n\ndef choose_slots(meta_df):\n    chosen = {}\n    for uid, g in meta_df.groupby(\"StudyInstanceUID\"):\n        one = {}\n        for slot_name, plane, kind in SLOTS:\n            cand = g[g[\"Anatomical_Plane\"] == plane].copy()\n            if len(cand) == 0:\n                continue\n            cand[\"priority\"] = cand.apply(lambda r: series_priority(r, kind), axis=1)\n            cand = cand.sort_values(\"priority\", ascending=False)\n            one[slot_name] = cand.iloc[0][\"SeriesInstanceUID\"]\n        chosen[uid] = one\n    return chosen\n\ntest_slot_map = choose_slots(test_series)\n\ncoverage = {\n    slot_name: np.mean([\n        slot_name in test_slot_map.get(uid, {})\n        for uid in test_df[\"StudyInstanceUID\"]\n    ])\n    for slot_name, _, _ in SLOTS\n}\nprint(\"test slot coverage:\", {k: round(v, 3) for k, v in coverage.items()})\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n## 3. DICOM → 2.5D temporary test cache\n\nSame preprocessing as training:\n- physical DICOM slice ordering;\n- central 15–85% coverage;\n- 9 sampled slices;\n- ~130 mm center crop;\n- percentile normalization;\n- resize 224×224;\n- three 3-slice 2.5D windows.\n","metadata":{}},{"cell_type":"code","source":"\ndef dicom_sort_key(path):\n    try:\n        ds = pydicom.dcmread(path, stop_before_pixels=True, force=True)\n        iop = getattr(ds, \"ImageOrientationPatient\", None)\n        ipp = getattr(ds, \"ImagePositionPatient\", None)\n        if iop is not None and ipp is not None and len(iop) >= 6 and len(ipp) >= 3:\n            row = np.array(iop[:3], dtype=float)\n            col = np.array(iop[3:6], dtype=float)\n            normal = np.cross(row, col)\n            return (0, float(np.dot(np.array(ipp[:3], dtype=float), normal)))\n        if hasattr(ds, \"InstanceNumber\"):\n            return (1, float(ds.InstanceNumber))\n    except Exception:\n        pass\n    return (2, str(path))\n\ndef read_one_dicom(path):\n    ds = pydicom.dcmread(path, force=True)\n    arr = ds.pixel_array.astype(np.float32)\n    arr = arr * float(getattr(ds, \"RescaleSlope\", 1.0) or 1.0)\n    arr = arr + float(getattr(ds, \"RescaleIntercept\", 0.0) or 0.0)\n\n    if str(getattr(ds, \"PhotometricInterpretation\", \"\")).upper() == \"MONOCHROME1\":\n        arr = arr.max() - arr\n\n    px = getattr(ds, \"PixelSpacing\", None)\n    spacing = None\n    if px is not None and len(px) >= 2:\n        try:\n            spacing = float(np.mean([float(px[0]), float(px[1])]))\n        except Exception:\n            pass\n    return arr, spacing\n\ndef center_crop_physical(arr, spacing, crop_mm=CROP_MM):\n    if spacing is None or not np.isfinite(spacing) or spacing <= 0:\n        side = min(arr.shape[-2:])\n    else:\n        side = int(round(crop_mm / spacing))\n        side = max(64, min(side, min(arr.shape[-2:])))\n\n    h, w = arr.shape[-2:]\n    cy, cx = h // 2, w // 2\n    y0 = max(0, cy - side // 2)\n    x0 = max(0, cx - side // 2)\n    return arr[y0:y0 + side, x0:x0 + side]\n\ndef read_series_windows(series_dir, img_size=IMG_SIZE):\n    files = list(Path(series_dir).glob(\"*.dcm\"))\n    if len(files) < 3:\n        return None\n\n    files = sorted(files, key=dicom_sort_key)\n\n    lo = int(round(0.15 * (len(files) - 1)))\n    hi = int(round(0.85 * (len(files) - 1)))\n    idx = np.linspace(lo, hi, N_SAMPLE_SLICES).round().astype(int)\n\n    slices = []\n    for i in idx:\n        try:\n            arr, spacing = read_one_dicom(files[int(i)])\n            arr = center_crop_physical(arr, spacing)\n            t = torch.from_numpy(np.ascontiguousarray(arr)).float()[None, None]\n            t = F.interpolate(\n                t, size=(img_size, img_size),\n                mode=\"bilinear\", align_corners=False\n            )\n            slices.append(t[0, 0].numpy())\n        except Exception:\n            slices.append(np.zeros((img_size, img_size), np.float32))\n\n    vol = np.stack(slices, axis=0)\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    vol = (vol * 255.0).round().astype(np.uint8)\n\n    return np.stack([vol[0:3], vol[3:6], vol[6:9]], axis=0)\n\ndef build_test_cache():\n    series_root = ROOT / \"test_series\"\n    shape = (\n        len(test_df), N_SLOTS, N_WINDOWS,\n        3, IMG_SIZE, IMG_SIZE\n    )\n\n    cache_path = CACHE_DIR / \"test_2p5d_uint8.dat\"\n    mask_path = CACHE_DIR / \"test_2p5d_mask.npy\"\n    expected_bytes = int(np.prod(shape)) * np.dtype(np.uint8).itemsize\n\n    if (\n        cache_path.exists()\n        and mask_path.exists()\n        and cache_path.stat().st_size == expected_bytes\n    ):\n        arr = np.memmap(cache_path, dtype=np.uint8, mode=\"r+\", shape=shape)\n        mask = np.load(mask_path)\n        if mask.shape == (len(test_df), N_SLOTS):\n            print(\"reusing temporary test cache\")\n            return arr, mask\n\n    arr = np.memmap(cache_path, dtype=np.uint8, mode=\"w+\", shape=shape)\n    arr[:] = 0\n    mask = np.zeros((len(test_df), N_SLOTS), np.float32)\n\n    slot_to_i = {name: i for i, (name, _, _) in enumerate(SLOTS)}\n    jobs = []\n    for row_i, uid in enumerate(test_df[\"StudyInstanceUID\"]):\n        for slot_name, sid in test_slot_map.get(uid, {}).items():\n            sdir = series_root / uid / sid\n            if sdir.is_dir():\n                jobs.append((row_i, slot_to_i[slot_name], sdir))\n\n    print(f\"decoding {len(jobs)} selected test series\")\n    t0 = time.time()\n\n    def task(job):\n        i, s, d = job\n        return i, s, read_series_windows(d)\n\n    done = 0\n    with ThreadPoolExecutor(max_workers=CACHE_THREADS) as ex:\n        futures = [ex.submit(task, j) for j in jobs]\n        for fut in as_completed(futures):\n            i, s, win = fut.result()\n            if win is not None:\n                arr[i, s] = win\n                mask[i, s] = 1.0\n            done += 1\n            if done % 500 == 0 or done == len(jobs):\n                print(f\"  {done}/{len(jobs)} | {(time.time()-t0)/60:.1f} min\")\n\n    arr.flush()\n    np.save(mask_path, mask)\n    print(\"filled-slot fraction:\", float(mask.mean()))\n    return arr, mask\n\ntest_cache, test_mask = build_test_cache()\nprint(\"test cache shape:\", test_cache.shape)\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n## 4. Model definition\n\nThe fold `.pt` files saved by the training notebook contain the full `state_dict`, including the frozen lower ResNet layers.\n","metadata":{}},{"cell_type":"code","source":"\nfrom torchvision.models import resnet50\n\nclass SlotAttentionHead(nn.Module):\n    def __init__(self, dim, hidden=384, dropout=0.20):\n        super().__init__()\n        self.proj = nn.Sequential(\n            nn.LayerNorm(dim),\n            nn.Linear(dim, hidden),\n            nn.GELU(),\n        )\n        self.slot_emb = nn.Parameter(torch.randn(N_SLOTS, hidden) * 0.02)\n        self.query = nn.Parameter(torch.randn(len(TARGETS), hidden) * 0.02)\n        self.dropout = nn.Dropout(dropout)\n        self.classifier = nn.Parameter(torch.randn(len(TARGETS), hidden) * 0.02)\n        self.bias = nn.Parameter(torch.zeros(len(TARGETS)))\n        self.hidden = hidden\n\n    def forward(self, feat, mask):\n        h = self.proj(feat) + self.slot_emb[None]\n        att = torch.einsum(\"bsh,th->bts\", h, self.query) / math.sqrt(self.hidden)\n        att = att.masked_fill(mask[:, None, :] < 0.5, -1e4)\n        att = att.softmax(dim=-1)\n        ctx = torch.einsum(\"bts,bsh->bth\", att, h)\n        ctx = self.dropout(ctx)\n        logits = (ctx * self.classifier[None]).sum(-1) + self.bias\n        return logits\n\nclass KneeModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = resnet50(weights=None)\n        self.backbone.fc = nn.Identity()\n        self.head = SlotAttentionHead(2048)\n        self.register_buffer(\n            \"mean\", torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)\n        )\n        self.register_buffer(\n            \"std\", torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)\n        )\n\n    def forward(self, imgs, mask):\n        B, S, C, H, W = imgs.shape\n        x = imgs.reshape(B * S, C, H, W).float() / 255.0\n        x = (x - self.mean) / self.std\n        feat = self.backbone(x).reshape(B, S, -1)\n        return self.head(feat, mask)\n\n@torch.no_grad()\ndef predict_all_windows(model, cache, mask_arr, batch_size=BATCH_STUDIES):\n    model.eval()\n    n = len(test_df)\n    acc = np.zeros((n, len(TARGETS)), np.float32)\n\n    for win in range(N_WINDOWS):\n        parts = []\n        for st in range(0, n, batch_size):\n            sel = np.arange(st, min(st + batch_size, n))\n\n            imgs = np.stack([\n                np.stack([cache[i, s, win] for s in range(N_SLOTS)], axis=0)\n                for i in sel\n            ], axis=0)\n\n            imgs = torch.from_numpy(np.ascontiguousarray(imgs)).to(\n                DEVICE, non_blocking=True\n            )\n            masks = torch.from_numpy(mask_arr[sel].astype(np.float32)).to(\n                DEVICE, non_blocking=True\n            )\n\n            with torch.autocast(\n                device_type=\"cuda\",\n                dtype=torch.float16,\n                enabled=DEVICE.type == \"cuda\",\n            ):\n                logits = model(imgs, masks)\n\n            parts.append(torch.sigmoid(logits).float().cpu().numpy())\n\n        acc += np.concatenate(parts, axis=0) / N_WINDOWS\n\n    return acc\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Load five folds, predict, rank-ensemble, write `submission.csv`","metadata":{}},{"cell_type":"code","source":"\nfold_preds = []\n\nfor fold, ckpt_path in enumerate(FOLD_WEIGHTS):\n    print(f\"\\n===== fold {fold} =====\")\n    print(\"loading:\", ckpt_path)\n\n    model = KneeModel().to(DEVICE)\n    ckpt = torch.load(ckpt_path, map_location=\"cpu\", weights_only=False)\n    state = ckpt[\"model\"] if isinstance(ckpt, dict) and \"model\" in ckpt else ckpt\n\n    missing, unexpected = model.load_state_dict(state, strict=False)\n    if missing or unexpected:\n        print(\"missing:\", missing)\n        print(\"unexpected:\", unexpected)\n        raise RuntimeError(\n            \"Checkpoint/model mismatch. \"\n            \"Use the inference notebook matching the training version.\"\n        )\n\n    pred = predict_all_windows(\n        model, test_cache, test_mask,\n        batch_size=BATCH_STUDIES\n    )\n    fold_preds.append(pred)\n\n    print(\"prediction range:\", float(pred.min()), \"to\", float(pred.max()))\n\n    del model, ckpt, state, pred\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\nrank_preds = [\n    pd.DataFrame(p, columns=TARGETS).rank(pct=True).values\n    for p in fold_preds\n]\nfinal_pred = np.mean(rank_preds, axis=0)\n\nsubmission = pd.DataFrame(final_pred, columns=TARGETS)\nsubmission.insert(0, \"StudyInstanceUID\", test_df[\"StudyInstanceUID\"].values)\n\nassert submission.shape == (len(test_df), 1 + len(TARGETS))\nassert submission[\"StudyInstanceUID\"].is_unique\nassert np.isfinite(submission[TARGETS].values).all()\n\nout_path = Path(\"/kaggle/working/submission.csv\")\nsubmission.to_csv(out_path, index=False)\n\nprint(\"\\nSaved:\", out_path)\nprint(\"shape:\", submission.shape)\ndisplay(submission.head())\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Final format sanity check\nsample_path = ROOT / \"sample_submission.csv\"\nif sample_path.is_file():\n    sample = pd.read_csv(sample_path)\n    assert list(submission.columns) == list(sample.columns)\n\nprint(\"submission.csv exists:\", Path(\"/kaggle/working/submission.csv\").is_file())\nprint(\"rows:\", len(submission))\nprint(\"ready for scoring\")\n","metadata":{},"outputs":[],"execution_count":null}]}