{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11"},"w01":{"stage":"RSNA_Knee_W01_04_Inference_2xT4.ipynb","backbone":"DINOv2-S","slots":6,"data_sources":"intentionally_unpinned_for_easy_kaggle_override"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"9a941d70","cell_type":"markdown","source":"# RSNA Knee W01 — Stage 04: six-slot DINOv2-S five-seed inference\n\nHidden-test inference from one explicitly resolved `rsna_knee_w01_bundle.pt`.\nThe bundle embeds the exact tuned DINOv2-S encoder used by Stage 2 and all five\nStage 3 heads; no separate backbone checkpoint is searched at inference time.\n","metadata":{}},{"id":"bb16b32b","cell_type":"code","source":"from __future__ import annotations\n\nimport gc\nimport hashlib\nimport json\nimport math\nimport os\nimport re\nimport sys\nimport time\nimport traceback\nimport threading\nfrom concurrent.futures import ThreadPoolExecutor\nfrom dataclasses import dataclass\nfrom pathlib import Path\nfrom typing import Dict, List, Optional, Tuple\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\n\n\nTARGETS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n    \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\",\n    \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\",\n]\nUID = \"StudyInstanceUID\"\n\n\n@dataclass\nclass RuntimeCFG:\n    smoke: bool = os.getenv(\"RSNA_SMOKE_TEST\", \"0\") == \"1\"\n    comp_root: Optional[str] = os.getenv(\"RSNA_COMP_ROOT\")\n    bundle_path: Optional[str] = os.getenv(\"RSNA_MODEL_BUNDLE\")\n    model_dir: Optional[str] = os.getenv(\"RSNA_MODEL_DIR\")\n    expected_adapt_epoch: Optional[int] = (int(os.environ[\"RSNA_ADAPT_EPOCH\"]) if os.getenv(\"RSNA_ADAPT_EPOCH\") else None)\n    output_dir: str = os.getenv(\"RSNA_INFER_OUTPUT\", \"/kaggle/working\")\n    encode_batch: int = int(os.getenv(\"RSNA_INFER_ENCODE_BATCH\", \"16\"))\n    inference_gpu_count: int = 2\n    dicom_workers: int = 2\n\n\nrt = RuntimeCFG()\nif rt.smoke:\n    rt.output_dir = os.getenv(\"RSNA_INFER_OUTPUT\", str(Path.cwd() / \"smoke_infer_output\"))\n    rt.encode_batch = 8\n    rt.dicom_workers = 0\nOUT = Path(rt.output_dir)\nOUT.mkdir(parents=True, exist_ok=True)\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ntorch.manual_seed(2026)\ntorch.set_float32_matmul_precision(\"high\")\n\n\ndef find_comp_root() -> Path:\n    if rt.comp_root and (Path(rt.comp_root) / \"test.csv\").is_file():\n        return Path(rt.comp_root)\n    candidates = [\n        Path(\"/kaggle/input/competitions/rsna-knee-abnormality-detection\"),\n        Path(\"/kaggle/input/rsna-knee-abnormality-detection\"), Path.cwd() / \"data\", Path.cwd(),\n    ]\n    for p in candidates:\n        if (p / \"test.csv\").is_file() and (p / \"test_series\").is_dir():\n            return p\n    base = Path(\"/kaggle/input\")\n    if base.is_dir():\n        for p in base.iterdir():\n            if p.is_dir() and (p / \"test.csv\").is_file() and (p / \"test_series\").is_dir():\n                return p\n    raise FileNotFoundError(\"competition data not found\")\n\n\ndef find_bundle() -> Path:\n    if rt.bundle_path:\n        path = Path(rt.bundle_path)\n        path = path / \"rsna_knee_w01_bundle.pt\" if path.is_dir() else path\n        if not path.is_file():\n            raise FileNotFoundError(f\"RSNA_MODEL_BUNDLE does not exist: {path}\")\n        return path\n    if rt.model_dir:\n        path = Path(rt.model_dir) / \"rsna_knee_w01_bundle.pt\"\n        if not path.is_file():\n            raise FileNotFoundError(f\"Bundle not found in RSNA_MODEL_DIR: {path}\")\n        return path\n    candidates = []\n    for base in (Path(\"/kaggle/input\"), Path.cwd()):\n        if not base.is_dir():\n            continue\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 \"rsna_knee_w01_bundle.pt\" in files:\n                candidates.append(Path(root) / \"rsna_knee_w01_bundle.pt\")\n    candidates = sorted(set(candidates), key=lambda p: str(p))\n    if not candidates:\n        raise FileNotFoundError(\"W01 bundle not found; attach the completed W01 Stage03 output\")\n    if len(candidates) > 1:\n        listed = \"\\n\".join(f\"  - {p}\" for p in candidates)\n        raise RuntimeError(\"Multiple bundles found; set RSNA_MODEL_BUNDLE explicitly:\\n\" + listed)\n    return candidates[0]\n\n\nROOT = find_comp_root()\nBUNDLE_PATH = find_bundle()\ntry:\n    bundle = torch.load(BUNDLE_PATH, map_location=\"cpu\", weights_only=False)\nexcept TypeError:\n    bundle = torch.load(BUNDLE_PATH, map_location=\"cpu\")\nassert bundle[\"targets\"] == TARGETS\nassert bundle.get(\"version\")==\"rsna-knee-w01-fullgold-5seed-fixed28-swa5-2\", bundle.get(\"version\")\nassert bundle.get(\"experiment\")==\"w01_fullgold_5seed\" and bundle.get(\"base_mode\")==\"w01_cache\"\ntraining=bundle.get(\"full_gold_training\",{})\nassert int(training.get(\"gold_oversample\",-1))==8\nassert int(training.get(\"head_epochs\",-1))==(3 if rt.smoke else 28)\nassert int(training.get(\"swa_last_epochs\",-1))==(2 if rt.smoke else 5)\nassert training.get(\"validation_used\") is False and training.get(\"early_stopping_used\") is False\nassert bundle.get(\"model_schema\") == \"labelmil-series-transformer-w01-dinov2s-6slot\"\nassert bundle.get(\"backbone_tuned_state\"), \"missing shared tuned-backbone state\"\nexpected_heads = int(bundle[\"n_models\"])\nassert expected_heads == len(bundle[\"head_seeds\"])\nactual_heads = sum(len(branch[\"states\"]) for branch in bundle[\"branches\"].values())\nassert actual_heads == expected_heads, (expected_heads, actual_heads)\nif not rt.smoke:\n    assert expected_heads == 5, expected_heads\nSLOT_SPECS = [tuple(x) for x in bundle[\"slot_specs\"]]\nTOKENS_PER_SLOT = tuple(bundle[\"tokens_per_slot\"])\nIMAGE_SIZE = int(bundle[\"image_size\"])\nINPUT_MODE = str(bundle[\"input_mode\"])\nSPATIAL_GRID = int(bundle[\"spatial_grid\"])\nCROP_MM = float(bundle[\"crop_mm\"])\nTRIPLET_OFFSET = int(bundle[\"triplet_offset\"])\nCOVERAGE = tuple(bundle[\"coverage\"])\nMAX_TOKENS = sum(TOKENS_PER_SLOT)\nassert len(SLOT_SPECS) == len(TOKENS_PER_SLOT) == 6\nif not rt.smoke:\n    assert MAX_TOKENS == 64\nprint(\"root:\", ROOT)\nprint(\"bundle:\", BUNDLE_PATH, bundle[\"version\"])\npreprocess_payload = bundle.get(\"preprocess_payload\")\nif not isinstance(preprocess_payload, dict):\n    raise RuntimeError(\"Bundle lacks the Stage 0.5 preprocessing payload\")\nactual_signature = hashlib.sha256(json.dumps(preprocess_payload, sort_keys=True).encode()).hexdigest()\nassert actual_signature == bundle[\"preprocess_signature\"], (actual_signature, bundle[\"preprocess_signature\"])\nfor key, value in {\n    \"slot_specs\": bundle[\"slot_specs\"], \"tokens_per_slot\": bundle[\"tokens_per_slot\"],\n    \"image_size\": bundle[\"image_size\"], \"input_mode\": bundle[\"input_mode\"],\n    \"crop_mm\": bundle[\"crop_mm\"], \"triplet_offset\": bundle[\"triplet_offset\"],\n    \"coverage\": bundle[\"coverage\"],\n}.items():\n    left = list(value) if isinstance(value, tuple) else value\n    right = preprocess_payload[key]\n    if key == \"slot_specs\":\n        left = [list(x) for x in left]\n    if left != right:\n        raise RuntimeError(f\"Bundle field {key} differs from preprocessing payload\")\nselected_epoch = int(bundle.get(\"adaptation\", {}).get(\"selected_epoch\", -1))\nif rt.expected_adapt_epoch is not None and selected_epoch != rt.expected_adapt_epoch:\n    raise RuntimeError(f\"Requested adaptation epoch {rt.expected_adapt_epoch}, bundle uses {selected_epoch}\")\nbundle_digest = hashlib.sha256(BUNDLE_PATH.read_bytes()).hexdigest()\nprint(\"MODEL_SELECTION:\", {\"path\": str(BUNDLE_PATH), \"bundle_sha256\": bundle_digest,\n                           \"adapt_epoch\": selected_epoch, \"head_seeds\": list(bundle[\"head_seeds\"])})\nprint(\"branches:\", list(bundle[\"branches\"]))\nif not rt.smoke:\n    assert (IMAGE_SIZE, SPATIAL_GRID, int(bundle[\"projection_dim\"])) == (336, 7, 256)\nprint(\"available gpus:\", torch.cuda.device_count(),\n      [torch.cuda.get_device_name(i) for i in range(torch.cuda.device_count())] if torch.cuda.is_available() else [])\nprint(\"preprocess signature:\", actual_signature)\n\nraw_test_df = pd.read_csv(ROOT / \"test.csv\", dtype={UID: str})\ntest_series_df = pd.read_csv(ROOT / \"test_series.csv\", dtype={UID: str, \"SeriesInstanceUID\": str})\nsample = pd.read_csv(ROOT / \"sample_submission.csv\", dtype={UID: str})\nassert sample.columns.tolist() == [UID] + TARGETS\nassert set(sample[UID]) == set(raw_test_df[UID])\ntest_df = sample[[UID]].merge(raw_test_df, on=UID, how=\"left\", validate=\"one_to_one\")\n","metadata":{"lines_to_next_cell":1},"outputs":[],"execution_count":null},{"id":"320bb964","cell_type":"markdown","source":"## 1. DICOM preprocessing — identical to train\n","metadata":{}},{"id":"a930e5a8","cell_type":"code","source":"def dcm_files(folder: Path) -> List[Path]:\n    if not folder.is_dir():\n        return []\n    files = list(folder.glob(\"*.dcm\"))\n    return files if files else [p for p in folder.iterdir() if p.is_file()]\n\n\ndef header_key(path: Path):\n    tags = [\"ImagePositionPatient\", \"ImageOrientationPatient\", \"InstanceNumber\",\n            \"PixelSpacing\", \"Laterality\", \"ImageLaterality\"]\n    ds = pydicom.dcmread(path, stop_before_pixels=True, force=True, specific_tags=tags)\n    try:\n        iop = np.asarray(ds.ImageOrientationPatient, np.float64)\n        ipp = np.asarray(ds.ImagePositionPatient, np.float64)\n        key = float(np.dot(ipp, np.cross(iop[:3], iop[3:])))\n    except Exception:\n        iop = ipp = None\n        key = float(getattr(ds, \"InstanceNumber\", 0))\n    return key, iop, ipp, ds\n\n\ndef sorted_series_files(folder: Path):\n    recs = []\n    for p in dcm_files(folder):\n        try:\n            recs.append((p,) + header_key(p))\n        except Exception:\n            pass\n    if not recs:\n        return [], None, \"\"\n    recs.sort(key=lambda x: x[1])\n    dedup, seen = [], set()\n    for rec in recs:\n        k = round(float(rec[1]), 3)\n        if k not in seen:\n            dedup.append(rec); seen.add(k)\n    ds = dedup[0][4]\n    lat = str(getattr(ds, \"ImageLaterality\", \"\") or getattr(ds, \"Laterality\", \"\")).upper()\n    if lat not in (\"L\", \"R\"):\n        ipps = [r[3] for r in dedup if r[3] is not None]\n        if ipps:\n            x = float(np.median([p[0] for p in ipps]))\n            if abs(x) > 20:\n                lat = \"L\" if x > 0 else \"R\"\n    return [r[0] for r in dedup], dedup[0][2], lat\n\n\nTARGET_AXES = {\n    \"Sagittal\": (np.array([0, -1, 0.]), np.array([0, 0, -1.])),\n    \"Coronal\": (np.array([1, 0, 0.]), np.array([0, 0, -1.])),\n    \"Axial\": (np.array([1, 0, 0.]), np.array([0, 1, 0.])),\n}\n\n\ndef canonicalize(arr: np.ndarray, ds, plane: str):\n    spacing = list(map(float, getattr(ds, \"PixelSpacing\", [1.0, 1.0])))\n    try:\n        iop = np.asarray(ds.ImageOrientationPatient, np.float64)\n        col_axis, row_axis = iop[:3], iop[3:]\n        want_col, want_row = TARGET_AXES[plane]\n        if abs(np.dot(row_axis, want_col)) > abs(np.dot(col_axis, want_col)):\n            arr = arr.T\n            row_axis, col_axis = col_axis, row_axis\n            spacing = [spacing[1], spacing[0]]\n        if np.dot(col_axis, want_col) < 0:\n            arr = arr[:, ::-1]\n        if np.dot(row_axis, want_row) < 0:\n            arr = arr[::-1]\n    except Exception:\n        pass\n    return np.ascontiguousarray(arr), (spacing[0], spacing[1])\n\n\ndef crop_resize(arr: np.ndarray, spacing: Tuple[float, float], out: int):\n    h, w = arr.shape\n    ch = min(h, max(16, int(round(CROP_MM / max(spacing[0], 1e-3)))))\n    cw = min(w, max(16, int(round(CROP_MM / max(spacing[1], 1e-3)))))\n    y0, x0 = max(0, (h - ch) // 2), max(0, (w - cw) // 2)\n    x = torch.from_numpy(np.ascontiguousarray(arr[y0:y0 + ch, x0:x0 + cw])).float()[None, None]\n    return F.interpolate(x, (out, out), mode=\"bilinear\", align_corners=False)[0, 0].numpy()\n\n\n_FATSAT_RX = re.compile(r\"\\bfs\\b|fatsat|fat sat|\\bstir\\b|\\bspair\\b|\\bspir\\b|water excit|fatsup\")\n_T1_RX = re.compile(r\"\\bt1\\b|\\bt1w\\b\"); _T2_RX = re.compile(r\"\\bt2\\b|\\bt2w\\b\")\n_PD_RX = re.compile(r\"\\bpd\\b|\\bpdw\\b|proton|dens\")\n_GRE_RX = re.compile(r\"gradient|\\bgre\\b|\\bffe\\b|\\bflash\\b\")\n\n\ndef series_profile(folder: Path, row: pd.Series) -> dict:\n    files = dcm_files(folder)\n    csv_fluid = pd.to_numeric(pd.Series([row.get(\"Fluid_Sensitive\", np.nan)]), errors=\"coerce\").iloc[0]\n    csv_fs = pd.to_numeric(pd.Series([row.get(\"Fat_Suppression\", np.nan)]), errors=\"coerce\").iloc[0]\n    prof = {\"n_files\": len(files), \"fluid\": bool(csv_fluid == 1),\n            \"fatsat\": bool(csv_fs == 1), \"header_ok\": False}\n    if not files: return prof\n    tags = [\"SeriesDescription\", \"ProtocolName\", \"SequenceName\", \"ScanningSequence\",\n            \"SequenceVariant\", \"ScanOptions\", \"RepetitionTime\", \"EchoTime\"]\n    try:\n        ds = pydicom.dcmread(files[len(files)//2], stop_before_pixels=True, force=True, specific_tags=tags)\n        text = \" \".join(str(getattr(ds, k, \"\")) for k in tags[:6]).lower().replace(\"_\", \" \")\n        tr = float(getattr(ds, \"RepetitionTime\", np.nan)); te = float(getattr(ds, \"EchoTime\", np.nan))\n        gre = bool(_GRE_RX.search(text)); fatsat = bool(_FATSAT_RX.search(text)) or prof[\"fatsat\"]\n        t1 = bool(_T1_RX.search(text)) or (np.isfinite(tr) and tr <= 800 and not gre)\n        t2 = bool(_T2_RX.search(text)) or (np.isfinite(tr) and tr > 800 and np.isfinite(te) and te >= 60)\n        pdw = bool(_PD_RX.search(text)) or (np.isfinite(tr) and tr > 800 and np.isfinite(te) and te < 60)\n        prof.update({\"fluid\": bool(t2 or pdw or (prof[\"fluid\"] and not t1)), \"fatsat\": fatsat,\n                     \"struct\": bool(t1 or pdw or gre),\n                     \"header_ok\": bool(text.strip() or np.isfinite(tr) or np.isfinite(te))})\n    except Exception: pass\n    prof.setdefault(\"struct\", not prof[\"fluid\"]); return prof\n\n\ndef choose_series(rows: pd.DataFrame, image_root: Path) -> Dict[int, Path]:\n    chosen, used, records = {}, set(), []\n    for _, r in rows.iterrows():\n        folder = image_root / str(r[UID]) / str(r[\"SeriesInstanceUID\"])\n        rec = r.to_dict(); rec.update(series_profile(folder, r)); rec[\"folder\"] = folder; records.append(rec)\n    meta = pd.DataFrame(records)\n    if meta.empty: return chosen\n    for slot, (_, plane, desired) in enumerate(SLOT_SPECS):\n        g = meta[(meta[\"Anatomical_Plane\"] == plane) & (~meta[\"SeriesInstanceUID\"].isin(used))].copy()\n        if g.empty: continue\n        want_fluid = desired == \"fluid\"\n        match = np.where(want_fluid, g[\"fluid\"] & g[\"fatsat\"], g[\"struct\"] & ~g[\"fatsat\"])\n        fallback = np.where(want_fluid, g[\"fluid\"], g[\"struct\"])\n        n = g[\"n_files\"].clip(lower=1).to_numpy(float)\n        stack_quality = -np.abs(np.log(n / 32.0)) - 2.0 * ((n < 10) | (n > 160))\n        g[\"quality\"] = 6.0*match.astype(float) + 2.5*fallback.astype(float) + stack_quality\n        r = g.sort_values([\"quality\", \"n_files\"], ascending=[False, False]).iloc[0]\n        used.add(r[\"SeriesInstanceUID\"]); chosen[slot] = Path(r[\"folder\"])\n    return chosen\n\n\ndef sample_fractions(n_tokens: int) -> np.ndarray:\n    if n_tokens <= 1: return np.asarray([0.5], np.float32)\n    u = np.linspace(-1.0, 1.0, n_tokens); lo, hi = COVERAGE\n    return (0.5 + (hi-lo)/2.0 * np.sign(u) * np.abs(u)**1.6).astype(np.float32)\n\n\ndef read_series_triplets(folder: Path, plane: str, n_tokens: int):\n    files, _, lat = sorted_series_files(folder)\n    if not files:\n        return np.zeros((n_tokens, 3, IMAGE_SIZE, IMAGE_SIZE), np.uint8), np.zeros(n_tokens, bool)\n    if lat == \"R\" and plane == \"Sagittal\":\n        files = files[::-1]\n    centers = (sample_fractions(n_tokens) * (len(files) - 1)).round().astype(int)\n    if INPUT_MODE == \"repeat_center\":\n        triplets = [np.asarray([c, c, c], dtype=int) for c in centers]\n    else:\n        triplets = [np.clip([c - TRIPLET_OFFSET, c, c + TRIPLET_OFFSET], 0, len(files) - 1) for c in centers]\n    unique = sorted(set(int(i) for tri in triplets for i in tri))\n    decoded, sample_values = {}, []\n    for i in unique:\n        try:\n            ds = pydicom.dcmread(files[i], force=True)\n            arr = ds.pixel_array.astype(np.float32)\n            if arr.ndim == 3: arr = arr[len(arr)//2]\n            arr = arr * float(getattr(ds, \"RescaleSlope\", 1) or 1) + float(getattr(ds, \"RescaleIntercept\", 0) or 0)\n            if str(getattr(ds, \"PhotometricInterpretation\", \"\")).upper() == \"MONOCHROME1\":\n                arr = float(arr.max() + arr.min()) - arr\n            arr, spacing = canonicalize(arr, ds, plane)\n            if lat == \"R\" and plane in (\"Coronal\", \"Axial\"):\n                arr = arr[:, ::-1]\n            decoded[i] = (arr, spacing); sample_values.append(arr[::4, ::4].reshape(-1))\n        except Exception:\n            pass\n    if not decoded:\n        return np.zeros((n_tokens, 3, IMAGE_SIZE, IMAGE_SIZE), np.uint8), np.zeros(n_tokens, bool)\n    qlo, qhi = np.percentile(np.concatenate(sample_values), [1, 99])\n    available = sorted(decoded)\n    out = np.zeros((n_tokens, 3, IMAGE_SIZE, IMAGE_SIZE), np.uint8)\n    mask = np.zeros(n_tokens, bool)\n    for k, tri in enumerate(triplets):\n        planes = []\n        for i in tri:\n            j = min(available, key=lambda a: abs(a - int(i)))\n            arr, spacing = decoded[j]\n            arr = np.clip((arr - qlo) / max(qhi - qlo, 1e-6), 0, 1)\n            planes.append(crop_resize(arr, spacing, IMAGE_SIZE))\n        out[k] = np.clip(np.stack(planes) * 255, 0, 255).round().astype(np.uint8)\n        mask[k] = True\n    return out, mask\n\n\nclass StudyPixels(Dataset):\n    def __init__(self):\n        self.studies = test_df[[UID]].reset_index(drop=True)\n        self.series_groups = {k: g.copy() for k, g in test_series_df.groupby(UID)}\n        self.image_root = ROOT / \"test_series\"\n\n    def __len__(self):\n        return len(self.studies)\n\n    def __getitem__(self, idx):\n        study = str(self.studies.iloc[idx][UID])\n        rows = self.series_groups.get(study, pd.DataFrame(columns=test_series_df.columns))\n        selected = choose_series(rows, self.image_root)\n        images, masks, slots, zpos = [], [], [], []\n        for slot, n_tok in enumerate(TOKENS_PER_SLOT):\n            if slot in selected:\n                x, m = read_series_triplets(selected[slot], SLOT_SPECS[slot][1], n_tok)\n            else:\n                x = np.zeros((n_tok, 3, IMAGE_SIZE, IMAGE_SIZE), np.uint8)\n                m = np.zeros(n_tok, bool)\n            images.append(x); masks.append(m); slots.extend([slot] * n_tok)\n            zpos.extend(sample_fractions(n_tok).tolist())\n        return {\n            \"index\": idx, \"uid\": study,\n            \"images\": torch.from_numpy(np.concatenate(images)),\n            \"mask\": torch.from_numpy(np.concatenate(masks)),\n            \"slot\": torch.tensor(slots, dtype=torch.long),\n            \"z\": torch.tensor(zpos, dtype=torch.float32),\n        }\n","metadata":{"lines_to_next_cell":1},"outputs":[],"execution_count":null},{"id":"ded19273","cell_type":"markdown","source":"## 2. Bundled DINOv2-S encoder and unchanged LabelMIL fusion head\n","metadata":{}},{"id":"f3d8994e","cell_type":"code","source":"class TinyEncoder(nn.Module):\n    name = \"tiny\"\n    source_dim = 64\n    region_dim = 32\n\n    def __init__(self):\n        super().__init__()\n        self.n_regions = 1 + SPATIAL_GRID ** 2\n        g = torch.Generator().manual_seed(12345)\n        self.register_buffer(\"base_proj\", torch.randn(3, self.source_dim, generator=g) / math.sqrt(3))\n        self.projector = nn.Linear(self.source_dim, self.region_dim, bias=False)\n        self.project_norm = nn.LayerNorm(self.region_dim)\n        nn.init.orthogonal_(self.projector.weight)\n\n    def forward(self, x):\n        x = x.float() / 255.0\n        cls = x.mean((2, 3))[:, None]\n        spatial = F.adaptive_avg_pool2d(x, (SPATIAL_GRID, SPATIAL_GRID)).flatten(2).transpose(1, 2)\n        return self.project_norm(self.projector(torch.cat([cls, spatial], 1) @ self.base_proj))\n\n\nclass DinoV2SmallEncoder(nn.Module):\n    name = \"dinov2_small\"\n    source_dim = 384\n\n    def __init__(self):\n        super().__init__()\n        import timm\n        self.model = timm.create_model(\n            bundle[\"assets\"][\"timm_model\"], pretrained=False, img_size=IMAGE_SIZE,\n            num_classes=0, dynamic_img_size=False,\n        )\n        self.region_dim = int(bundle[\"projection_dim\"])\n        self.n_regions = 1 + SPATIAL_GRID ** 2\n        self.projector = nn.Linear(self.source_dim, self.region_dim, bias=False)\n        self.project_norm = nn.LayerNorm(self.region_dim)\n        self.register_buffer(\"mean\", torch.tensor([0.485, 0.456, 0.406])[None, :, None, None])\n        self.register_buffer(\"std\", torch.tensor([0.229, 0.224, 0.225])[None, :, None, None])\n\n    def forward(self, x):\n        x = (x.float() / 255.0 - self.mean) / self.std\n        tokens = self.model.forward_features(x)\n        if isinstance(tokens, dict):\n            tokens = tokens.get(\"x_norm_patchtokens\", tokens.get(\"x_prenorm\"))\n        if tokens.ndim != 3:\n            raise RuntimeError(f\"unexpected DINOv2-S feature shape: {tokens.shape}\")\n        n_prefix = int(getattr(self.model, \"num_prefix_tokens\", 1))\n        cls, patch = tokens[:, 0], tokens[:, n_prefix:]\n        side = int(round(math.sqrt(patch.shape[1])))\n        if side * side != patch.shape[1]:\n            raise RuntimeError(f\"non-square DINOv2-S patch layout: {patch.shape}\")\n        grid = patch.reshape(len(patch), side, side, self.source_dim).permute(0, 3, 1, 2)\n        spatial = F.adaptive_avg_pool2d(grid, (SPATIAL_GRID, SPATIAL_GRID)).flatten(2).transpose(1, 2)\n        return self.project_norm(self.projector(torch.cat([cls[:, None], spatial], 1)))\n\n\ndef build_encoders(device: torch.device) -> Dict[str, nn.Module]:\n    encoders = {}\n    for name in bundle[\"branches\"]:\n        if name == \"tiny\":\n            enc = TinyEncoder()\n        elif name == \"dinov2_small\":\n            enc = DinoV2SmallEncoder()\n        else:\n            raise KeyError(f\"unknown branch {name}\")\n        tuned = bundle[\"backbone_tuned_state\"].get(name)\n        if not tuned:\n            raise RuntimeError(f\"missing full tuned-backbone state for {name}\")\n        result = enc.load_state_dict(tuned, strict=True)\n        if result.missing_keys or result.unexpected_keys:\n            raise RuntimeError((result.missing_keys[:10], result.unexpected_keys[:10]))\n        enc = enc.to(device).eval()\n        expected = bundle[\"branches\"][name]\n        assert enc.region_dim == int(expected[\"region_dim\"])\n        assert enc.n_regions == int(expected[\"n_regions\"])\n        encoders[name] = enc\n        print(f\"loaded bundled {name} on {device}: {len(tuned)} tensors\")\n    return encoders\n\n\nclass LabelMIL(nn.Module):\n    def __init__(self, region_dim: int, hidden: int, n_slots: int, dropout: float,\n                 series_layers: int, series_heads: int, series_dropout: float):\n        super().__init__()\n        if hidden % series_heads:\n            raise ValueError(f\"hidden={hidden} must be divisible by series_heads={series_heads}\")\n        self.region_dim = region_dim\n        self.hidden = hidden\n        self.n_slots = n_slots\n        self.proj = nn.Sequential(nn.LayerNorm(region_dim), nn.Linear(region_dim, hidden), nn.GELU())\n        self.slot_emb = nn.Embedding(n_slots, hidden)\n        self.z_mlp = nn.Sequential(nn.Linear(1, hidden), nn.Tanh(), nn.Linear(hidden, hidden))\n        self.spatial_a = nn.Linear(hidden, hidden); self.spatial_b = nn.Linear(hidden, hidden)\n        self.spatial_q = nn.Parameter(torch.randn(len(TARGETS), hidden) * 0.02)\n        layer = nn.TransformerEncoderLayer(\n            d_model=hidden, nhead=series_heads, dim_feedforward=hidden * 3,\n            dropout=series_dropout, activation=\"gelu\", batch_first=True, norm_first=True,\n        )\n        self.series_encoder = nn.TransformerEncoder(layer, num_layers=series_layers)\n        self.series_norm = nn.LayerNorm(hidden)\n        self.depth_a = nn.Linear(hidden, hidden); self.depth_b = nn.Linear(hidden, hidden)\n        self.depth_q = nn.Parameter(torch.randn(len(TARGETS), hidden) * 0.02)\n        # slot order: sag fluid/struct, cor fluid/struct, axial fluid/struct\n        prior = torch.tensor([\n            [1.0, .8, .5, .4, .1, .2], [.4, .4, 1.0, .8, .1, .2],\n            [.8, 1.0, .7, .8, .1, .2], [.8, 1.0, .7, .8, .1, .2],\n            [.4, .5, .8, 1.0, .2, .3], [.4, .5, .8, 1.0, .2, .3],\n            [.2, .2, .2, .2, 1.0, .9], [.6, .2, .4, .2, 1.0, .8],\n            [.6, .2, .4, .2, 1.0, .7], [1.0, .3, .4, .2, .7, .5],\n            [.8, .3, .8, .3, .8, .7], [.6, .6, .6, .6, .6, .7],\n        ], dtype=torch.float32)\n        if n_slots != prior.shape[1]:\n            raise ValueError(f\"W01 requires six slots, got {n_slots}\")\n        prior = prior - prior.mean(1, keepdim=True)\n        self.slot_bias = nn.Parameter(0.35 * prior)\n        self.cls_weight = nn.Parameter(torch.randn(len(TARGETS), hidden) * 0.02)\n        self.cls_bias = nn.Parameter(torch.zeros(len(TARGETS)))\n        self.mean_residual = nn.Linear(hidden, len(TARGETS))\n        self.dropout = nn.Dropout(dropout)\n\n    def forward(self, feat, mask, slot, z):\n        # First perform target-specific spatial pooling inside every 2D slice.\n        base = self.slot_emb(slot) + self.z_mlp(z.unsqueeze(-1))\n        h = self.proj(feat) + base[:, :, None]\n        spatial_gate = torch.tanh(self.spatial_a(h)) * torch.sigmoid(self.spatial_b(h))\n        spatial_score = torch.einsum(\"btrh,lh->bltr\", spatial_gate, self.spatial_q) / math.sqrt(self.hidden)\n        spatial_att = spatial_score.softmax(-1)\n        token_ctx = torch.einsum(\"bltr,btrh->blth\", spatial_att, h)\n        token_ctx = token_ctx * mask[:, None, :, None]\n\n        # V4A: true slice-to-slice interaction. Each diagnostic series is encoded\n        # independently, preserving its ordered z positions and padding mask.\n        batch, labels, _, hidden = token_ctx.shape\n        sequence_ctx = torch.zeros_like(token_ctx)\n        for slot_id in range(self.n_slots):\n            positions = torch.nonzero(slot[0] == slot_id, as_tuple=False).flatten()\n            if not len(positions):\n                continue\n            value = token_ctx.index_select(2, positions)\n            n_depth = value.shape[2]\n            value = value.reshape(batch * labels, n_depth, hidden)\n            padding = (~mask.index_select(1, positions)).unsqueeze(1)\n            padding = padding.expand(batch, labels, n_depth).reshape(batch * labels, n_depth)\n            valid_rows = (~padding).any(1)\n            encoded = torch.zeros_like(value)\n            if valid_rows.any():\n                valid_index = torch.nonzero(valid_rows, as_tuple=False).flatten()\n                valid_encoded = self.series_encoder(\n                    value.index_select(0, valid_index),\n                    src_key_padding_mask=padding.index_select(0, valid_index),\n                )\n                encoded = encoded.index_copy(0, valid_index, valid_encoded)\n            # CUDA autocast may return FP32 from LayerNorm while token_ctx is FP16.\n            # index_copy requires an exact dtype match on GPU.\n            encoded = self.series_norm(encoded).to(sequence_ctx.dtype)\n            encoded = encoded.reshape(batch, labels, n_depth, hidden)\n            sequence_ctx = sequence_ctx.index_copy(2, positions, encoded)\n        sequence_ctx = sequence_ctx * mask[:, None, :, None]\n\n        # Target-specific pooling now consumes contextualized slice features.\n        depth_gate = torch.tanh(self.depth_a(sequence_ctx)) * torch.sigmoid(self.depth_b(sequence_ctx))\n        depth_score = torch.einsum(\"blth,lh->blt\", depth_gate, self.depth_q) / math.sqrt(self.hidden)\n        depth_score = depth_score + self.slot_bias[:, slot].permute(1, 0, 2)\n        depth_score = depth_score.masked_fill(~mask[:, None], -1e4)\n        depth_att = depth_score.softmax(-1)\n        ctx = torch.einsum(\"blt,blth->blh\", depth_att, sequence_ctx)\n        logits = (self.dropout(ctx) * self.cls_weight[None]).sum(-1) + self.cls_bias\n        denom = (mask.sum(1, keepdim=True) * feat.shape[2]).clamp_min(1)\n        mean = (h * mask[:, :, None, None]).sum((1, 2)) / denom\n        return logits + 0.25 * self.mean_residual(self.dropout(mean))\n","metadata":{"lines_to_next_cell":1},"outputs":[],"execution_count":null},{"id":"87d6c96c","cell_type":"markdown","source":"## 3. Hidden-test inference — five-seed rank ensemble with optional T4×2 sharding\n","metadata":{}},{"id":"2b74d889","cell_type":"code","source":"def column_rank(x: np.ndarray) -> np.ndarray:\n    return pd.DataFrame(x).rank(method=\"average\", pct=True).to_numpy(np.float64)\n\n\ndef _inference_devices() -> List[torch.device]:\n    if torch.cuda.is_available():\n        count = min(int(rt.inference_gpu_count), torch.cuda.device_count())\n        return [torch.device(f\"cuda:{i}\") for i in range(max(1, count))]\n    return [torch.device(\"cpu\")] * (2 if rt.smoke else 1)\n\n\ndef build_heads(device: torch.device) -> Dict[str, List[nn.Module]]:\n    heads = {}\n    for name, rec in bundle[\"branches\"].items():\n        models = []\n        for state in rec[\"states\"]:\n            model = LabelMIL(\n                int(rec[\"region_dim\"]), int(bundle[\"hidden_dim\"]), len(SLOT_SPECS),\n                float(bundle[\"head_dropout\"]), int(bundle[\"series_layers\"]),\n                int(bundle[\"series_heads\"]), float(bundle[\"series_dropout\"]),\n            ).to(device)\n            model.load_state_dict(state); model.eval(); models.append(model)\n        heads[name] = models\n    return heads\n\n\n@torch.inference_mode()\ndef encode_valid_slices(enc: nn.Module, images: torch.Tensor, valid: np.ndarray,\n                        device: torch.device) -> torch.Tensor:\n    feat = torch.zeros((1, MAX_TOKENS, int(enc.n_regions), int(enc.region_dim)), device=device)\n    start, batch_size = 0, rt.encode_batch\n    while start < len(valid):\n        pos = valid[start:start + batch_size]\n        try:\n            x = images[pos].to(device, non_blocking=True)\n            with torch.autocast(\"cuda\", enabled=device.type == \"cuda\"):\n                value = enc(x).float()\n            feat[0, torch.as_tensor(pos, device=device, dtype=torch.long)] = value\n            start += len(pos)\n        except torch.cuda.OutOfMemoryError:\n            if batch_size <= 1:\n                raise\n            batch_size = max(1, batch_size // 2)\n            print(f\"CUDA OOM on {device}: reducing encode_batch to {batch_size}\")\n            if device.type == \"cuda\":\n                torch.cuda.empty_cache()\n    return feat\n\n\n@torch.inference_mode()\ndef run_inference() -> pd.DataFrame:\n    devices = _inference_devices()\n    shards = [np.arange(i, len(test_df), len(devices), dtype=np.int64) for i in range(len(devices))]\n    merged = np.sort(np.concatenate(shards)) if shards else np.empty(0, np.int64)\n    if not np.array_equal(merged, np.arange(len(test_df))):\n        raise RuntimeError(\"inference shard planner failed\")\n    print(\"inference devices:\", [str(x) for x in devices], \"studies/shard:\", [len(x) for x in shards])\n    # Build sequentially: torch.hub/module import and checkpoint deserialization are\n    # safer outside worker threads; actual encoding/prediction then runs concurrently.\n    runtimes = [(build_encoders(device), build_heads(device)) for device in devices]\n    raw = {\n        name: np.zeros((len(rec[\"states\"]), len(test_df), len(TARGETS)), np.float32)\n        for name, rec in bundle[\"branches\"].items()\n    }\n    missing = np.zeros(len(test_df), np.float32)\n    valid_by_slot = np.zeros((len(test_df), len(SLOT_SPECS)), np.int16)\n    print_lock = threading.Lock()\n    started = time.time()\n\n    def worker(worker_id: int) -> dict:\n        # Grad/inference mode is thread-local; worker threads do not inherit the\n        # decorator on run_inference(). Disable autograd explicitly per worker.\n        torch.set_grad_enabled(False)\n        device = devices[worker_id]\n        if device.type == \"cuda\":\n            torch.cuda.set_device(device)\n        encoders, heads = runtimes[worker_id]\n        ds = StudyPixels()\n        local_start = time.time()\n        for local_step, idx in enumerate(shards[worker_id]):\n            rec = ds[int(idx)]\n            images = rec[\"images\"]\n            mask_cpu, slot_cpu, z_cpu = rec[\"mask\"], rec[\"slot\"], rec[\"z\"]\n            mask = mask_cpu[None].to(device)\n            slot = slot_cpu[None].to(device)\n            z = z_cpu[None].to(device)\n            mask_np, slot_np = mask_cpu.numpy(), slot_cpu.numpy()\n            missing[int(idx)] = float(1.0 - mask_cpu.float().mean())\n            valid_by_slot[int(idx)] = [(mask_np & (slot_np == s)).sum() for s in range(len(SLOT_SPECS))]\n            valid = np.flatnonzero(mask_np)\n            for name, enc in encoders.items():\n                feat = encode_valid_slices(enc, images, valid, device)\n                for model_id, head in enumerate(heads[name]):\n                    with torch.autocast(\"cuda\", enabled=device.type == \"cuda\"):\n                        pred = torch.sigmoid(head(feat, mask, slot, z)).float().cpu().numpy()[0]\n                    raw[name][model_id, int(idx)] = pred\n            if local_step % max(1, len(shards[worker_id]) // 10) == 0 or local_step + 1 == len(shards[worker_id]):\n                with print_lock:\n                    print(f\"infer worker={worker_id} device={device} {local_step+1}/{len(shards[worker_id])} \"\n                          f\"elapsed={(time.time()-local_start)/60:.1f}m\")\n        return {\"worker\": worker_id, \"device\": str(device), \"studies\": int(len(shards[worker_id])),\n                \"elapsed_minutes\": (time.time() - local_start) / 60}\n\n    with ThreadPoolExecutor(max_workers=len(devices), thread_name_prefix=\"w01-infer\") as pool:\n        worker_metrics = list(pool.map(worker, range(len(devices))))\n\n    branch_rank = {name: np.mean([column_rank(p) for p in pred], axis=0) for name, pred in raw.items()}\n    final = column_rank(branch_rank.get(\"dinov2_small\", next(iter(branch_rank.values()))))\n    primary_raw = raw.get(\"dinov2_small\", next(iter(raw.values())))\n    np.save(OUT / \"seed_predictions_raw.npy\", primary_raw)\n    np.save(OUT / \"predictions_rank_ensemble.npy\", final)\n    sub = test_df[[UID]].copy()\n    sub[TARGETS] = np.clip(final, 1e-6, 1 - 1e-6)\n    assert sub.columns.tolist() == [UID] + TARGETS\n    assert len(sub) == len(test_df) and sub[TARGETS].notna().all().all()\n    sub.to_csv(OUT / \"submission.csv\", index=False)\n    diagnostics = {\n        \"n_studies\": len(test_df), \"bundle_version\": bundle[\"version\"],\n        \"bundle_path\": str(BUNDLE_PATH), \"bundle_sha256\": bundle_digest,\n        \"selected_adapt_epoch\": selected_epoch, \"preprocess_signature\": actual_signature,\n        \"full_gold_training\": bundle.get(\"full_gold_training\", {}), \"head_seeds\": list(bundle[\"head_seeds\"]),\n        \"branches\": list(bundle[\"branches\"]),\n        \"devices\": [str(x) for x in devices],\n        \"dual_gpu_active\": bool(len(devices) == 2 and all(x.type == \"cuda\" for x in devices)),\n        \"workers\": worker_metrics, \"elapsed_minutes\": (time.time() - started) / 60,\n        \"mean_missing_token_fraction\": float(missing.mean()),\n        \"mean_valid_tokens_per_slot\": valid_by_slot.mean(0).round(3).tolist(),\n        \"prediction_std\": dict(zip(TARGETS, sub[TARGETS].std().round(6))),\n    }\n    with open(OUT / \"inference_diagnostics.json\", \"w\", encoding=\"utf-8\") as f:\n        json.dump(diagnostics, f, indent=2)\n    print(sub.head().to_string(index=False))\n    print(\"saved:\", OUT / \"submission.csv\")\n    return sub\n\n\nsubmission = run_inference()\n\nif rt.smoke:\n    assert submission[TARGETS].notna().all().all()\n    print(\"STAGE04_SMOKE_PASS\")\n","metadata":{},"outputs":[],"execution_count":null}]}