{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11"},"kaggle":{"accelerator":"gpu"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"fa0d0b94-50ce-41e2-b27a-b891ffe26c59","cell_type":"markdown","source":"# RSNA Knee D0 — simple ResNet34 2.5D single-model inference\n\nHidden-test inference for the bundle produced by the paired D0 training notebook.\n\n- exact W01 Stage 0.5 DICOM geometry and intensity preprocessing;\n- deterministic center-bin selection from the same dense candidate layout;\n- one ResNet34 model, raw sigmoid probabilities;\n- no fold/model ensemble, rank averaging, TTA, attention, or report branch.\n\nOnly the training notebook needs ImageNet weights. This notebook reconstructs the model entirely from `simple_resnet34_bundle.pt`.\n","metadata":{}},{"id":"b7605a17-992a-4773-9d5b-acfc78d5554d","cell_type":"code","source":"from __future__ import annotations\n\nimport hashlib, json, os, re, time\nfrom contextlib import nullcontext\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\nfrom torchvision.models import resnet34\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    comp_root: Optional[str] = os.getenv(\"RSNA_COMP_ROOT\")\n    bundle_path: Optional[str] = os.getenv(\"RSNA_MODEL_BUNDLE\")\n    output_dir: str = os.getenv(\"RSNA_INFER_OUTPUT\", \"/kaggle/working\")\n    study_batch_size: int = int(os.getenv(\"RSNA_INFER_STUDY_BATCH\", \"2\"))\n    dicom_workers: int = int(os.getenv(\"RSNA_DICOM_WORKERS\", \"2\"))\n    require_two_gpus: bool = os.getenv(\"RSNA_REQUIRE_2GPU\", \"1\") == \"1\"\n\n\nrt = RuntimeCFG()\nOUT = Path(rt.output_dir)\nOUT.mkdir(parents=True, exist_ok=True)\nDEVICE = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\ntorch.set_float32_matmul_precision(\"high\")\ntorch.manual_seed(2026)\n\n\ndef find_comp_root() -> Path:\n    if rt.comp_root:\n        path = Path(rt.comp_root)\n        if (path / \"test.csv\").is_file() and (path / \"test_series\").is_dir():\n            return path\n        raise FileNotFoundError(f\"Invalid RSNA_COMP_ROOT: {path}\")\n    candidates = [\n        Path(\"/kaggle/input/competitions/rsna-knee-abnormality-detection\"),\n        Path(\"/kaggle/input/rsna-knee-abnormality-detection\"),\n        Path.cwd() / \"data\", Path.cwd(),\n    ]\n    for path in candidates:\n        if (path / \"test.csv\").is_file() and (path / \"test_series\").is_dir():\n            return path\n    base = Path(\"/kaggle/input\")\n    if base.is_dir():\n        for path in base.iterdir():\n            if path.is_dir() and (path / \"test.csv\").is_file() and (path / \"test_series\").is_dir():\n                return path\n    raise FileNotFoundError(\"Attach the competition data or set RSNA_COMP_ROOT\")\n\n\ndef find_bundle() -> Path:\n    filename = \"simple_resnet34_bundle.pt\"\n    if rt.bundle_path:\n        path = Path(rt.bundle_path)\n        path = path if path.is_file() else path / filename\n        if path.is_file():\n            return path\n        raise FileNotFoundError(path)\n    found = []\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[:] = [name for name in dirs if name not in (\"train_series\", \"test_series\")]\n            if filename in files:\n                found.append(Path(root) / filename)\n    found = sorted(set(path.resolve() for path in found), key=str)\n    if len(found) != 1:\n        raise RuntimeError(f\"Expected exactly one {filename}; found {found}. Set RSNA_MODEL_BUNDLE.\")\n    return found[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\")\n\nif bundle.get(\"version\") != \"rsna-knee-w01-d0-resnet34-v1\":\n    raise RuntimeError(f\"Wrong bundle version: {bundle.get('version')}\")\nif bundle.get(\"model_schema\") != \"resnet34-2p5d-sixslot-meanmax-mlp-v1\":\n    raise RuntimeError(f\"Wrong model schema: {bundle.get('model_schema')}\")\nif bundle.get(\"targets\") != TARGETS:\n    raise RuntimeError(\"Target order mismatch\")\n\nSLOT_SPECS = [tuple(value) for value in bundle[\"slot_specs\"]]\nDENSE_TOKENS_PER_SLOT = tuple(int(value) for value in bundle[\"dense_tokens_per_slot\"])\nSELECTED_TOKENS_PER_SLOT = tuple(int(value) for value in bundle[\"selected_tokens_per_slot\"])\nCACHE_IMAGE_SIZE = int(bundle[\"cache_image_size\"])\nMODEL_IMAGE_SIZE = int(bundle[\"model_image_size\"])\nPREPROCESS_PAYLOAD = bundle[\"preprocess_payload\"]\nPREPROCESS_SIGNATURE = hashlib.sha256(\n    json.dumps(PREPROCESS_PAYLOAD, sort_keys=True).encode()\n).hexdigest()\nif PREPROCESS_SIGNATURE != bundle[\"preprocess_signature\"]:\n    raise RuntimeError(\"Bundle preprocessing payload/signature mismatch\")\nif [list(value) for value in SLOT_SPECS] != [\n    list(value) for value in PREPROCESS_PAYLOAD[\"slot_specs\"]\n]:\n    raise RuntimeError(\"Slot specification mismatch\")\nif list(DENSE_TOKENS_PER_SLOT) != PREPROCESS_PAYLOAD[\"tokens_per_slot\"]:\n    raise RuntimeError(\"Dense token layout mismatch\")\n\nraw_test_df = pd.read_csv(ROOT / \"test.csv\", dtype={UID: str})\nseries_df = pd.read_csv(ROOT / \"test_series.csv\", dtype={UID: str, \"SeriesInstanceUID\": str})\nsample = pd.read_csv(ROOT / \"sample_submission.csv\", dtype={UID: str})\nif sample.columns.tolist() != [UID] + TARGETS:\n    raise RuntimeError(\"Unexpected sample_submission.csv columns\")\nif set(sample[UID]) != set(raw_test_df[UID]):\n    raise RuntimeError(\"test.csv and sample_submission.csv UID sets differ\")\ntest_df = sample[[UID]].merge(raw_test_df, on=UID, how=\"left\", validate=\"one_to_one\")\n\nprint(\"competition root:\", ROOT)\nprint(\"bundle:\", BUNDLE_PATH)\nprint(\"training scope:\", bundle.get(\"training\", {}).get(\"scope\"))\nprint(\"best CV epoch/AUC:\", bundle.get(\"training\", {}).get(\"best_epoch\"),\n      bundle.get(\"training\", {}).get(\"best_report_weighted_macro_auc\"))\nprint(\"selected tokens per slot:\", SELECTED_TOKENS_PER_SLOT)\nprint(\"device:\", DEVICE, \"gpus:\", torch.cuda.device_count())\nif rt.require_two_gpus and torch.cuda.is_available() and torch.cuda.device_count() < 2:\n    print(\"warning: only one GPU is available; inference remains correct but slower\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"74c72103-07d2-432d-b0e0-96d9b9b45913","cell_type":"markdown","source":"## 1. Exact W01 Stage 0.5 preprocessing for hidden test\n\nThe code sorts slices by physical position, deduplicates physical locations, canonicalizes DICOM LPS orientation/laterality, selects six plane/contrast slots, applies a centered 140 mm crop, normalizes with per-series 1st–99th percentiles, and constructs adjacent 2.5D triplets.\n","metadata":{}},{"id":"9389f3cb-e541-457d-84d7-c17bb8f0ab0d","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 [path for path in folder.iterdir() if path.is_file()]\n\n\ndef header_key(path: Path):\n    tags = [\n        \"ImagePositionPatient\", \"ImageOrientationPatient\", \"InstanceNumber\",\n        \"PixelSpacing\", \"Laterality\", \"ImageLaterality\",\n    ]\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    records = []\n    for path in dcm_files(folder):\n        try:\n            records.append((path,) + header_key(path))\n        except Exception:\n            pass\n    if not records:\n        return [], \"\"\n    records.sort(key=lambda item: item[1])\n    deduplicated, seen = [], set()\n    for record in records:\n        key = round(float(record[1]), 3)\n        if key not in seen:\n            deduplicated.append(record)\n            seen.add(key)\n    first_ds = deduplicated[0][4]\n    laterality = str(\n        getattr(first_ds, \"ImageLaterality\", \"\") or getattr(first_ds, \"Laterality\", \"\")\n    ).upper()\n    if laterality not in (\"L\", \"R\"):\n        positions = [record[3] for record in deduplicated if record[3] is not None]\n        if positions:\n            patient_x = float(np.median([position[0] for position in positions]))\n            if abs(patient_x) > 20:\n                laterality = \"L\" if patient_x > 0 else \"R\"\n    return [record[0] for record in deduplicated], laterality\n\n\nTARGET_AXES = {\n    \"Sagittal\": (np.array([0, -1, 0.0]), np.array([0, 0, -1.0])),\n    \"Coronal\": (np.array([1, 0, 0.0]), np.array([0, 0, -1.0])),\n    \"Axial\": (np.array([1, 0, 0.0]), np.array([0, 1, 0.0])),\n}\n\n\ndef canonicalize(array: 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        column_axis, row_axis = iop[:3], iop[3:]\n        wanted_column, wanted_row = TARGET_AXES[plane]\n        if abs(np.dot(row_axis, wanted_column)) > abs(np.dot(column_axis, wanted_column)):\n            array = array.T\n            row_axis, column_axis = column_axis, row_axis\n            spacing = [spacing[1], spacing[0]]\n        if np.dot(column_axis, wanted_column) < 0:\n            array = array[:, ::-1]\n        if np.dot(row_axis, wanted_row) < 0:\n            array = array[::-1]\n    except Exception:\n        pass\n    return np.ascontiguousarray(array), (spacing[0], spacing[1])\n\n\ndef crop_resize(array: np.ndarray, spacing: Tuple[float, float]) -> np.ndarray:\n    height, width = array.shape\n    crop_mm = float(PREPROCESS_PAYLOAD[\"crop_mm\"])\n    crop_h = min(height, max(16, int(round(crop_mm / max(spacing[0], 1e-3)))))\n    crop_w = min(width, max(16, int(round(crop_mm / max(spacing[1], 1e-3)))))\n    y0 = max(0, (height - crop_h) // 2)\n    x0 = max(0, (width - crop_w) // 2)\n    tensor = torch.from_numpy(\n        np.ascontiguousarray(array[y0:y0 + crop_h, x0:x0 + crop_w])\n    ).float()[None, None]\n    tensor = F.interpolate(\n        tensor, (CACHE_IMAGE_SIZE, CACHE_IMAGE_SIZE), mode=\"bilinear\", align_corners=False\n    )\n    return tensor[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\")\n_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    profile = {\n        \"n_files\": len(files),\n        \"fluid\": bool(csv_fluid == 1),\n        \"fatsat\": bool(csv_fs == 1),\n        \"header_ok\": False,\n    }\n    if not files:\n        return profile\n    tags = [\n        \"SeriesDescription\", \"ProtocolName\", \"SequenceName\", \"ScanningSequence\",\n        \"SequenceVariant\", \"ScanOptions\", \"RepetitionTime\", \"EchoTime\",\n    ]\n    try:\n        ds = pydicom.dcmread(\n            files[len(files) // 2], stop_before_pixels=True, force=True, specific_tags=tags\n        )\n        text = \" \".join(str(getattr(ds, key, \"\")) for key in tags[:6]).lower().replace(\"_\", \" \")\n        tr = float(getattr(ds, \"RepetitionTime\", np.nan))\n        te = float(getattr(ds, \"EchoTime\", np.nan))\n        gre = bool(_GRE_RX.search(text))\n        fatsat = bool(_FATSAT_RX.search(text)) or profile[\"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 (\n            np.isfinite(tr) and tr > 800 and np.isfinite(te) and te >= 60\n        )\n        pdw = bool(_PD_RX.search(text)) or (\n            np.isfinite(tr) and tr > 800 and np.isfinite(te) and te < 60\n        )\n        profile.update({\n            \"fluid\": bool(t2 or pdw or (profile[\"fluid\"] and not t1)),\n            \"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        })\n    except Exception:\n        pass\n    profile.setdefault(\"struct\", not profile[\"fluid\"])\n    return profile\n\n\ndef choose_series(rows: pd.DataFrame, image_root: Path) -> Dict[int, Path]:\n    chosen, used, records = {}, set(), []\n    for _, row in rows.iterrows():\n        folder = image_root / str(row[UID]) / str(row[\"SeriesInstanceUID\"])\n        record = row.to_dict()\n        record.update(series_profile(folder, row))\n        record[\"folder\"] = folder\n        records.append(record)\n    metadata = pd.DataFrame(records)\n    if metadata.empty:\n        return chosen\n    for slot, (_, plane, desired) in enumerate(SLOT_SPECS):\n        group = metadata[\n            (metadata[\"Anatomical_Plane\"] == plane)\n            & (~metadata[\"SeriesInstanceUID\"].isin(used))\n        ].copy()\n        if group.empty:\n            continue\n        want_fluid = desired == \"fluid\"\n        match = np.where(\n            want_fluid,\n            group[\"fluid\"] & group[\"fatsat\"],\n            group[\"struct\"] & ~group[\"fatsat\"],\n        )\n        fallback = np.where(want_fluid, group[\"fluid\"], group[\"struct\"])\n        n_files = group[\"n_files\"].clip(lower=1).to_numpy(float)\n        stack_quality = -np.abs(np.log(n_files / 32.0)) - 2.0 * (\n            (n_files < 10) | (n_files > 160)\n        )\n        group[\"quality\"] = 6.0 * match.astype(float) + 2.5 * fallback.astype(float) + stack_quality\n        selected = group.sort_values([\"quality\", \"n_files\"], ascending=[False, False]).iloc[0]\n        used.add(selected[\"SeriesInstanceUID\"])\n        chosen[slot] = Path(selected[\"folder\"])\n    return chosen\n\n\ndef sample_fractions(n_tokens: int) -> np.ndarray:\n    if n_tokens <= 1:\n        return np.asarray([0.5], np.float32)\n    coordinate = np.linspace(-1.0, 1.0, n_tokens)\n    low, high = PREPROCESS_PAYLOAD[\"coverage\"]\n    half = (high - low) / 2.0\n    return (0.5 + half * np.sign(coordinate) * np.abs(coordinate) ** 1.6).astype(np.float32)\n\n\ndef center_bin_positions(n_candidates: int, count: int) -> np.ndarray:\n    if count <= 0 or n_candidates <= 0:\n        return np.empty(0, np.int64)\n    if count >= n_candidates:\n        return np.arange(n_candidates, dtype=np.int64)\n    positions = ((np.arange(count, dtype=np.float64) + 0.5) * n_candidates / count - 0.5)\n    return np.clip(np.rint(positions).astype(np.int64), 0, n_candidates - 1)\n\n\ndef read_selected_triplets(folder: Path, plane: str, dense_count: int, selected_count: int):\n    files, laterality = sorted_series_files(folder)\n    if not files:\n        return (\n            np.zeros((selected_count, 3, CACHE_IMAGE_SIZE, CACHE_IMAGE_SIZE), np.uint8),\n            np.zeros(selected_count, bool),\n        )\n    if laterality == \"R\" and plane == \"Sagittal\":\n        files = files[::-1]\n    dense_centers = (sample_fractions(dense_count) * (len(files) - 1)).round().astype(int)\n    offset = int(PREPROCESS_PAYLOAD[\"triplet_offset\"])\n    dense_triplets = [\n        np.clip([center - offset, center, center + offset], 0, len(files) - 1)\n        for center in dense_centers\n    ]\n    unique = sorted(set(int(index) for triplet in dense_triplets for index in triplet))\n    decoded, sample_values = {}, []\n    for index in unique:\n        try:\n            ds = pydicom.dcmread(files[index], force=True)\n            array = ds.pixel_array.astype(np.float32)\n            if array.ndim == 3:\n                array = array[len(array) // 2]\n            array = (\n                array * float(getattr(ds, \"RescaleSlope\", 1) or 1)\n                + float(getattr(ds, \"RescaleIntercept\", 0) or 0)\n            )\n            if str(getattr(ds, \"PhotometricInterpretation\", \"\")).upper() == \"MONOCHROME1\":\n                array = float(array.max() + array.min()) - array\n            array, spacing = canonicalize(array, ds, plane)\n            if laterality == \"R\" and plane in (\"Coronal\", \"Axial\"):\n                array = array[:, ::-1]\n            decoded[index] = (array, spacing)\n            sample_values.append(array[::4, ::4].reshape(-1))\n        except Exception:\n            pass\n    if not decoded:\n        return (\n            np.zeros((selected_count, 3, CACHE_IMAGE_SIZE, CACHE_IMAGE_SIZE), np.uint8),\n            np.zeros(selected_count, bool),\n        )\n    q_low, q_high = np.percentile(np.concatenate(sample_values), [1, 99])\n    available = sorted(decoded)\n    chosen_dense = center_bin_positions(dense_count, selected_count)\n    output = np.zeros((selected_count, 3, CACHE_IMAGE_SIZE, CACHE_IMAGE_SIZE), np.uint8)\n    mask = np.zeros(selected_count, bool)\n    for output_index, dense_index in enumerate(chosen_dense):\n        planes = []\n        for index in dense_triplets[int(dense_index)]:\n            nearest = min(available, key=lambda value: abs(value - int(index)))\n            array, spacing = decoded[nearest]\n            normalized = np.clip((array - q_low) / max(q_high - q_low, 1e-6), 0, 1)\n            planes.append(crop_resize(normalized, spacing))\n        output[output_index] = np.clip(np.stack(planes) * 255, 0, 255).round().astype(np.uint8)\n        mask[output_index] = True\n    return output, mask\n","metadata":{},"outputs":[],"execution_count":null},{"id":"b3174d81-75fa-408d-9b0e-d8d2319f3e51","cell_type":"markdown","source":"## 2. Test study dataset and model\n","metadata":{}},{"id":"2a419dfd-1add-4a57-9774-9c7603b34675","cell_type":"code","source":"class TestStudyDataset(Dataset):\n    def __init__(self):\n        self.studies = test_df[[UID]].reset_index(drop=True)\n        self.series_groups = {key: group.copy() for key, group in 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, index: int):\n        uid = str(self.studies.iloc[index][UID])\n        rows = self.series_groups.get(uid, pd.DataFrame(columns=series_df.columns))\n        selected_series = choose_series(rows, self.image_root)\n        images, masks, slots = [], [], []\n        for slot_id, (dense_count, selected_count) in enumerate(\n            zip(DENSE_TOKENS_PER_SLOT, SELECTED_TOKENS_PER_SLOT)\n        ):\n            if slot_id in selected_series:\n                slot_images, slot_mask = read_selected_triplets(\n                    selected_series[slot_id], SLOT_SPECS[slot_id][1], dense_count, selected_count\n                )\n            else:\n                slot_images = np.zeros(\n                    (selected_count, 3, CACHE_IMAGE_SIZE, CACHE_IMAGE_SIZE), np.uint8\n                )\n                slot_mask = np.zeros(selected_count, bool)\n            images.append(slot_images)\n            masks.append(slot_mask)\n            slots.extend([slot_id] * selected_count)\n        return {\n            \"index\": index,\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        }\n\n\nclass ResNet34Encoder(nn.Module):\n    feature_dim = 512\n\n    def __init__(self):\n        super().__init__()\n        backbone = resnet34(weights=None)\n        backbone.fc = nn.Identity()\n        self.backbone = backbone\n        normalization = bundle[\"normalization\"]\n        self.register_buffer(\n            \"pixel_mean\", torch.tensor(normalization[\"mean\"]).view(1, 3, 1, 1)\n        )\n        self.register_buffer(\n            \"pixel_std\", torch.tensor(normalization[\"std\"]).view(1, 3, 1, 1)\n        )\n\n    def forward(self, images: torch.Tensor) -> torch.Tensor:\n        x = images.float().div_(255.0)\n        if x.shape[-2:] != (MODEL_IMAGE_SIZE, MODEL_IMAGE_SIZE):\n            x = F.interpolate(\n                x,\n                size=(MODEL_IMAGE_SIZE, MODEL_IMAGE_SIZE),\n                mode=\"bilinear\",\n                align_corners=False,\n            )\n        x = (x - self.pixel_mean) / self.pixel_std\n        return self.backbone(x)\n\n\nclass SixSlotMeanMaxHead(nn.Module):\n    def __init__(self):\n        super().__init__()\n        fusion_dim = len(SLOT_SPECS) * (2 * ResNet34Encoder.feature_dim + 1)\n        hidden_dim = int(bundle[\"hidden_dim\"])\n        dropout = float(bundle[\"dropout\"])\n        self.net = nn.Sequential(\n            nn.LayerNorm(fusion_dim),\n            nn.Dropout(dropout),\n            nn.Linear(fusion_dim, hidden_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, len(TARGETS)),\n        )\n\n    def forward(self, features: torch.Tensor, mask: torch.Tensor, slot: torch.Tensor):\n        descriptors = []\n        for slot_id in range(len(SLOT_SPECS)):\n            selected = mask & slot.eq(slot_id)\n            count = selected.sum(1, keepdim=True)\n            mean = (features * selected.unsqueeze(-1)).sum(1) / count.clamp_min(1)\n            maximum = features.masked_fill(~selected.unsqueeze(-1), -1e4).max(1).values\n            present = count.gt(0)\n            mean = torch.where(present, mean, torch.zeros_like(mean))\n            maximum = torch.where(present, maximum, torch.zeros_like(maximum))\n            descriptors.extend([mean, maximum, present.float()])\n        return self.net(torch.cat(descriptors, dim=1))\n\n\nencoder = ResNet34Encoder().to(DEVICE)\nhead = SixSlotMeanMaxHead().to(DEVICE)\nencoder.load_state_dict(bundle[\"encoder_state\"], strict=True)\nhead.load_state_dict(bundle[\"head_state\"], strict=True)\nencoder.eval()\nhead.eval()\nrunner = (\n    nn.DataParallel(encoder, device_ids=[0, 1])\n    if DEVICE.type == \"cuda\" and torch.cuda.device_count() >= 2\n    else encoder\n)\n\n\ndef encode_studies(images: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:\n    batch, tokens = images.shape[:2]\n    flat_images = images.reshape(batch * tokens, *images.shape[2:])\n    flat_mask = mask.reshape(-1)\n    indices = flat_mask.nonzero(as_tuple=False).squeeze(1)\n    if len(indices):\n        valid_features = runner(flat_images.index_select(0, indices))\n        all_features = valid_features.new_zeros((batch * tokens, ResNet34Encoder.feature_dim))\n        all_features = all_features.index_copy(0, indices, valid_features)\n    else:\n        all_features = next(encoder.parameters()).new_zeros(\n            (batch * tokens, ResNet34Encoder.feature_dim)\n        )\n    return all_features.view(batch, tokens, ResNet34Encoder.feature_dim)\n\n\ndef autocast_context():\n    if DEVICE.type == \"cuda\":\n        return torch.autocast(device_type=\"cuda\", dtype=torch.float16)\n    return nullcontext()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"ee231659-fb6b-47e9-860f-7470342c8a7d","cell_type":"markdown","source":"## 3. Single-model prediction and submission\n","metadata":{}},{"id":"34791a2b-ae6a-40a9-b505-a354b8c32409","cell_type":"code","source":"loader_kwargs = dict(\n    dataset=TestStudyDataset(),\n    batch_size=rt.study_batch_size,\n    shuffle=False,\n    num_workers=rt.dicom_workers,\n    pin_memory=DEVICE.type == \"cuda\",\n    persistent_workers=bool(rt.dicom_workers > 0),\n)\nif rt.dicom_workers > 0:\n    loader_kwargs[\"prefetch_factor\"] = 2\nloader = DataLoader(**loader_kwargs)\n\npredictions, row_order, valid_counts = [], [], []\nstarted = time.time()\nwith torch.inference_mode():\n    for step, batch in enumerate(loader):\n        images = batch[\"images\"].to(DEVICE, non_blocking=True)\n        mask = batch[\"mask\"].to(DEVICE, non_blocking=True)\n        slot = batch[\"slot\"].to(DEVICE, non_blocking=True)\n        with autocast_context():\n            features = encode_studies(images, mask)\n            logits = head(features, mask, slot)\n        predictions.append(torch.sigmoid(logits.float()).cpu().numpy())\n        row_order.extend(batch[\"index\"].numpy().astype(int).tolist())\n        valid_counts.append(np.stack([\n            ((batch[\"mask\"] & batch[\"slot\"].eq(slot_id)).sum(1)).numpy()\n            for slot_id in range(len(SLOT_SPECS))\n        ], axis=1))\n        if (step + 1) % 50 == 0:\n            print(f\"batch={step + 1}/{len(loader)} elapsed_min={(time.time() - started) / 60:.1f}\")\n\nprediction = np.concatenate(predictions, axis=0)\nrow_order = np.asarray(row_order, np.int64)\nvalid_counts = np.concatenate(valid_counts, axis=0)\nif sorted(row_order.tolist()) != list(range(len(test_df))):\n    raise RuntimeError(\"Inference row coverage/order is incomplete\")\nordered_prediction = np.empty_like(prediction)\nordered_prediction[row_order] = prediction\nordered_prediction = np.clip(ordered_prediction, 1e-5, 1 - 1e-5)\n\nsubmission = sample[[UID]].copy()\nsubmission[TARGETS] = ordered_prediction\nif submission[TARGETS].isna().any().any():\n    raise RuntimeError(\"NaN detected in submission\")\nif submission[UID].tolist() != sample[UID].tolist():\n    raise RuntimeError(\"Submission UID order changed\")\n\nsubmission_path = OUT / \"submission.csv\"\nsubmission.to_csv(submission_path, index=False)\ndiagnostics = {\n    \"version\": bundle[\"version\"],\n    \"bundle_path\": str(BUNDLE_PATH),\n    \"bundle_sha256\": hashlib.sha256(BUNDLE_PATH.read_bytes()).hexdigest(),\n    \"single_model\": True,\n    \"preprocess_signature\": PREPROCESS_SIGNATURE,\n    \"selected_tokens_per_slot\": list(SELECTED_TOKENS_PER_SLOT),\n    \"mean_valid_tokens_per_slot\": valid_counts.mean(0).round(4).tolist(),\n    \"prediction_mean\": dict(zip(TARGETS, submission[TARGETS].mean().round(6))),\n    \"prediction_std\": dict(zip(TARGETS, submission[TARGETS].std().round(6))),\n    \"elapsed_minutes\": (time.time() - started) / 60,\n}\n(OUT / \"inference_diagnostics.json\").write_text(\n    json.dumps(diagnostics, indent=2), encoding=\"utf-8\"\n)\nprint(submission.head().to_string(index=False))\nprint(json.dumps(diagnostics, indent=2))\nprint(\"saved:\", submission_path)\nprint(\"SIMPLE_RESNET34_INFERENCE_COMPLETE\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"e034203d-21e1-4ddb-a107-8a732c205396","cell_type":"markdown","source":"## Kaggle run contract\n\n1. Attach competition data and exactly one completed D0 training output.\n2. Run all cells with 2×T4. No external pretrained checkpoint is needed here.\n3. Submit `/kaggle/working/simple_resnet34_inference/submission.csv`.\n","metadata":{}}]}