{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11"},"v7b":{"architecture":"CoAtNet Raptor","single_method":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"e6ed740b-d428-439c-ab43-5fb278e9ef60","cell_type":"markdown","source":"# RSNA Knee — V7B single-CoAtNet inference\n\nRestores exactly one V7B CoAtNet checkpoint, rebuilds the public 44-slice/140-mm/15–85% stack from hidden-test DICOM, applies deterministic 24-window 2.5D inference, and writes the 12-target submission. No OrthoFoundation and no multi-architecture fusion.","metadata":{}},{"id":"c0ee256d-8b49-47ab-8ecf-4674ce1c7782","cell_type":"code","source":"from __future__ import annotations\n\nimport gc, hashlib, json, math, os, re, threading, time\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 Dataset\n\ntry:\n    import timm\nexcept Exception:\n    timm = None\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    output_dir: str = os.getenv(\"RSNA_INFER_OUTPUT\", \"/kaggle/working\")\n    max_gpus: int = 2\n\n\nrt = RuntimeCFG()\nif rt.smoke:\n    rt.output_dir = os.getenv(\"RSNA_INFER_OUTPUT\", str(Path.cwd() / \"smoke_v7b_infer\"))\nOUT = Path(rt.output_dir); OUT.mkdir(parents=True, exist_ok=True)\ntorch.manual_seed(42); torch.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(): return Path(rt.comp_root)\n    for path in (Path(\"/kaggle/input/competitions/rsna-knee-abnormality-detection\"),\n                 Path(\"/kaggle/input/rsna-knee-abnormality-detection\"), Path.cwd() / \"data\", Path.cwd()):\n        if (path / \"test.csv\").is_file() and (path / \"test_series\").is_dir(): 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(\"competition test data not found\")\n\n\ndef find_bundle() -> Path:\n    if rt.bundle_path and Path(rt.bundle_path).is_file(): return Path(rt.bundle_path)\n    matches = []\n    for base in (Path(\"/kaggle/input\"), Path.cwd()):\n        if not base.is_dir(): 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_coatnet_v7b_bundle.pt\" in files:\n                matches.append(Path(root) / \"rsna_knee_coatnet_v7b_bundle.pt\")\n    matches = sorted(set(matches))\n    if len(matches) != 1:\n        raise FileNotFoundError(f\"expected exactly one V7B bundle, found: {matches}\")\n    return matches[0]\n\n\nROOT = find_comp_root(); BUNDLE_PATH = find_bundle()\ntry: bundle = torch.load(BUNDLE_PATH, map_location=\"cpu\", weights_only=False)\nexcept TypeError: bundle = torch.load(BUNDLE_PATH, map_location=\"cpu\")\nassert bundle[\"version\"] == \"rsna-knee-v7b-coatnet-raptor-single-method\"\nassert bundle[\"model_schema\"] == \"raptor-coatnet-target-attention-v1\"\nassert bundle[\"targets\"] == TARGETS and bundle.get(\"model_state\")\n\nSLOT_SPECS = [tuple(x) for x in bundle[\"slot_specs\"]]\nTOKENS_PER_SLOT = tuple(bundle[\"tokens_per_slot\"])\nIMAGE_SIZE = int(bundle[\"corpus_size\"])\nMODEL_RESOLUTION = int(bundle[\"resolution\"])\nEVAL_WINDOWS = int(bundle[\"eval_windows\"])\nCROP_MM = float(bundle[\"crop_mm\"])\nCOVERAGE = tuple(bundle[\"coverage\"])\nINPUT_MODE = \"repeat_center\"\nTRIPLET_OFFSET = 1\nassert TOKENS_PER_SLOT == (12, 10, 8, 6, 8) and sum(TOKENS_PER_SLOT) == 44\n\ncontract = {\"slot_specs\": bundle[\"slot_specs\"], \"tokens_per_slot\": bundle[\"tokens_per_slot\"],\n            \"corpus_size\": bundle[\"corpus_size\"], \"crop_mm\": bundle[\"crop_mm\"],\n            \"coverage\": bundle[\"coverage\"], \"input_mode\": bundle[\"input_mode\"],\n            \"resolution\": bundle[\"resolution\"], \"eval_windows\": bundle[\"eval_windows\"]}\nsignature = hashlib.sha256(json.dumps(contract, sort_keys=True).encode()).hexdigest()\nassert signature == bundle[\"preprocess_signature\"]\n\nraw_test = 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 and set(sample[UID]) == set(raw_test[UID])\ntest_df = sample[[UID]].merge(raw_test, on=UID, how=\"left\", validate=\"one_to_one\")\nprint(\"root:\", ROOT, \"bundle:\", BUNDLE_PATH)\nprint(\"model:\", bundle[\"arch\"], \"selection:\", bundle[\"selection\"])\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 [])","metadata":{},"outputs":[],"execution_count":null},{"id":"9c6a3226-a59b-47fd-9b5e-9c97cfbf3e52","cell_type":"markdown","source":"## 1. DICOM → fixed 44×336 stack, identical slot/crop contract","metadata":{}},{"id":"543ade06-c6ff-4aa3-8734-633d46352798","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    lo, hi = COVERAGE\n    return np.linspace(lo, hi, n_tokens, dtype=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        }","metadata":{},"outputs":[],"execution_count":null},{"id":"75750743-2a47-459d-8f10-9e051175a879","cell_type":"markdown","source":"## 2. CoAtNet Raptor target-attention model","metadata":{}},{"id":"ed79b20e-2822-49db-bda5-64e0b06eebae","cell_type":"code","source":"class RaptorClassifier(nn.Module):\n    \"\"\"One CoAtNet encoder plus pathology-specific attention pooling.\"\"\"\n    def __init__(self, backbone: nn.Module, feature_dim: int, n_targets: int = 12, dropout: float = 0.2):\n        super().__init__()\n        self.backbone = backbone\n        self.norm = nn.LayerNorm(feature_dim)\n        self.att = nn.Sequential(\n            nn.Linear(feature_dim, 256), nn.Tanh(), nn.Dropout(dropout), nn.Linear(256, n_targets)\n        )\n        self.cls_weight = nn.Parameter(torch.zeros(n_targets, feature_dim))\n        self.cls_bias = nn.Parameter(torch.zeros(n_targets))\n        nn.init.trunc_normal_(self.cls_weight, std=0.02)\n\n    def encode(self, x):\n        batch, windows = x.shape[:2]\n        feature = self.backbone(x.flatten(0, 1))\n        return feature.reshape(batch, windows, -1)\n\n    def head(self, feature):\n        hidden = self.norm(feature)\n        attention = torch.softmax(self.att(hidden), dim=1)\n        pooled = torch.einsum(\"bkn,bkf->bnf\", attention, hidden)\n        return (pooled * self.cls_weight).sum(-1) + self.cls_bias\n\n    def forward(self, x):\n        return self.head(self.encode(x))","metadata":{},"outputs":[],"execution_count":null},{"id":"c40e1f1c-c5da-4a96-bf82-548af8f03e00","cell_type":"markdown","source":"## 3. One-model inference; optional two-GPU study sharding","metadata":{}},{"id":"bdcc64f4-2662-4797-a5dc-95692569c5de","cell_type":"code","source":"class TinyBackbone(nn.Module):\n    num_features = 32\n    def __init__(self):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(3, 16, 3, 2, 1), nn.GELU(),\n            nn.Conv2d(16, 32, 3, 2, 1), nn.GELU(), nn.AdaptiveAvgPool2d(1), nn.Flatten(),\n        )\n    def forward(self, x): return self.net(x)\n\n\ndef build_model(device: torch.device):\n    if bundle[\"arch\"] == \"tiny\":\n        backbone = TinyBackbone()\n    else:\n        if timm is None: raise ImportError(\"timm is required for CoAtNet inference\")\n        backbone = timm.create_model(bundle[\"arch\"], pretrained=False, num_classes=0,\n                                     in_chans=3, global_pool=\"avg\")\n    model = RaptorClassifier(backbone, int(bundle[\"feature_dim\"]), len(TARGETS),\n                             float(bundle[\"dropout\"]))\n    model.load_state_dict(bundle[\"model_state\"], strict=True)\n    for parameter in model.parameters(): parameter.requires_grad = False\n    return model.to(device).eval()\n\n\ndef inference_devices():\n    if torch.cuda.is_available():\n        return [torch.device(f\"cuda:{i}\") for i in range(min(rt.max_gpus, torch.cuda.device_count()))]\n    return [torch.device(\"cpu\")] * (2 if rt.smoke else 1)\n\n\ndef make_eval_windows(record: dict) -> torch.Tensor:\n    # DICOM preprocessing emitted the fixed 44-slot stack as repeated RGB centers.\n    gray = record[\"images\"][:, 1].float() / 255.0\n    valid = np.flatnonzero(record[\"mask\"].numpy())\n    if len(valid) < 3: valid = np.arange(min(3, len(gray)))\n    lo, hi = int(valid.min()), int(valid.max())\n    candidates = [c for c in range(lo + 1, hi) if c - 1 >= lo and c + 1 <= hi]\n    if not candidates: candidates = [max(1, min((lo + hi) // 2, len(gray) - 2))]\n    positions = np.linspace(0, len(candidates) - 1, EVAL_WINDOWS).round().astype(int)\n    windows = torch.stack([\n        torch.stack([gray[candidates[j] - 1], gray[candidates[j]], gray[candidates[j] + 1]])\n        for j in positions\n    ])\n    if windows.shape[-2:] != (MODEL_RESOLUTION, MODEL_RESOLUTION):\n        windows = F.interpolate(windows, (MODEL_RESOLUTION, MODEL_RESOLUTION),\n                                mode=\"bilinear\", align_corners=False)\n    return windows\n\n\ndef column_rank(values: np.ndarray):\n    return pd.DataFrame(values).rank(method=\"average\", pct=True).to_numpy(np.float64)\n\n\n@torch.inference_mode()\ndef run_inference():\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))\n    assert np.array_equal(merged, np.arange(len(test_df)))\n    models = [build_model(device) for device in devices]\n    raw = np.zeros((len(test_df), len(TARGETS)), np.float32)\n    worker_metrics, lock = [], threading.Lock()\n    print(\"inference devices:\", [str(x) for x in devices], \"studies/shard:\", [len(x) for x in shards])\n\n    def worker(worker_id: int):\n        torch.set_grad_enabled(False)\n        device, model = devices[worker_id], models[worker_id]\n        if device.type == \"cuda\": torch.cuda.set_device(device)\n        dataset = StudyPixels(); started = time.time()\n        for local_step, index in enumerate(shards[worker_id]):\n            record = dataset[int(index)]\n            windows = make_eval_windows(record).unsqueeze(0).to(device, non_blocking=True)\n            with torch.autocast(\"cuda\", dtype=torch.float16, enabled=device.type == \"cuda\"):\n                prediction = torch.sigmoid(model(windows)).float().cpu().numpy()[0]\n            raw[int(index)] = prediction\n            if local_step % max(1, len(shards[worker_id]) // 10) == 0 or local_step + 1 == len(shards[worker_id]):\n                with lock:\n                    print(f\"worker={worker_id} device={device} {local_step+1}/{len(shards[worker_id])} \"\n                          f\"elapsed={(time.time()-started)/60:.1f}m\")\n        return {\"worker\": worker_id, \"device\": str(device), \"studies\": int(len(shards[worker_id])),\n                \"elapsed_minutes\": (time.time() - started) / 60}\n\n    with ThreadPoolExecutor(max_workers=len(devices), thread_name_prefix=\"v7b-infer\") as pool:\n        worker_metrics.extend(pool.map(worker, range(len(devices))))\n    prediction = column_rank(raw)\n    submission = test_df[[UID]].copy(); submission[TARGETS] = prediction\n    assert submission.columns.tolist() == [UID] + TARGETS\n    assert submission[TARGETS].notna().all().all() and len(submission) == len(test_df)\n    submission.to_csv(OUT / \"submission.csv\", index=False)\n    np.save(OUT / \"predictions_raw.npy\", raw)\n    with open(OUT / \"inference_diagnostics_v7b.json\", \"w\", encoding=\"utf-8\") as f:\n        json.dump({\"devices\": [str(x) for x in devices], \"workers\": worker_metrics,\n                   \"n_studies\": len(test_df), \"preprocess_signature\": signature,\n                   \"prediction_std\": dict(zip(TARGETS, submission[TARGETS].std().round(6)))}, f, indent=2)\n    print(submission.head().to_string(index=False)); print(\"saved:\", OUT / \"submission.csv\")\n    return submission\n\n\nsubmission = run_inference()","metadata":{},"outputs":[],"execution_count":null}]}