{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11"},"w01":{"stage":"RSNA_Knee_W01_02_DINOv2S_SelectEpoch_Feature_Factory_Resume_2xT4.ipynb","backbone":"DINOv2-S","slots":6,"data_sources":"intentionally_unpinned_for_easy_kaggle_override"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# RSNA Knee W01 — Stage 02: Stage-0.5-cached DINOv2-S feature factory\n\nReads the exact uint8 image tokens produced by Stage 0.5, loads one explicitly selected\nStage 1 checkpoint (epoch 10, 20, or 30), and extracts the dense 64-token feature cache.\nThe checkpoint tensor hash and preprocessing signature are recorded for Stage 3.\n","metadata":{}},{"cell_type":"code","source":"from __future__ import annotations\n\nimport hashlib, json, math, os, random, re, shutil, time, threading\nfrom dataclasses import asdict, dataclass\nfrom pathlib import Path\nfrom typing import Dict, List, Optional, Tuple\nfrom concurrent.futures import ThreadPoolExecutor\n\nimport h5py\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset\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\"\nMODE = \"w01_cache\"\nCACHE_MANIFEST = \"w01_stage05_manifest.json\"\n\n@dataclass\nclass CFG:\n    seed: int = 2026\n    smoke: bool = os.getenv(\"RSNA_SMOKE_TEST\", \"0\") == \"1\"\n    comp_root: Optional[str] = os.getenv(\"RSNA_COMP_ROOT\")\n    cache_dirs: Optional[str] = os.getenv(\"RSNA_CACHE_DIRS\")\n    adapt_dir: Optional[str] = os.getenv(\"RSNA_ADAPT_DIR\")\n    adapted_checkpoint: Optional[str] = os.getenv(\"RSNA_ADAPTED_CHECKPOINT\")\n    feature_resume_dir: Optional[str] = os.getenv(\"RSNA_FEATURE_RESUME_DIR\")\n    output_dir: str = os.getenv(\"RSNA_FEATURE_OUTPUT\", \"/kaggle/working/w01_features\")\n    dinov2_model: str = \"vit_small_patch14_dinov2.lvd142m\"\n    image_size: int = 336\n    input_mode: str = \"adjacent\"\n    spatial_grid: int = 7\n    projection_dim: int = 256\n    crop_mm: float = 140.0\n    triplet_offset: int = 1\n    coverage: Tuple[float, float] = (0.02, 0.98)\n    tokens_per_slot: Tuple[int, ...] = (16, 12, 12, 8, 10, 6)\n    checkpoint_epochs: Tuple[int, ...] = (10, 20, 30)\n    selected_adapt_epoch: int = int(os.getenv(\"RSNA_ADAPT_EPOCH\", \"30\"))\n    encode_batch: int = int(os.getenv(\"RSNA_FEATURE_ENCODE_BATCH\", \"32\"))\n    time_budget_minutes: float = float(os.getenv(\"RSNA_TIME_BUDGET_MIN\", \"650\"))\n    stop_margin_minutes: float = float(os.getenv(\"RSNA_STOP_MARGIN_MIN\", \"10\"))\n    require_two_gpus: bool = os.getenv(\"RSNA_REQUIRE_2GPU\", \"1\") == \"1\"\n    verify_cache_hashes: bool = os.getenv(\"RSNA_VERIFY_CACHE_HASHES\", \"0\") == \"1\"\n\ncfg = CFG()\nif cfg.smoke:\n    cfg.output_dir = os.getenv(\"RSNA_FEATURE_OUTPUT\", str(Path.cwd() / \"smoke_w01_features\"))\n    cfg.image_size = 64\n    cfg.spatial_grid = 7\n    cfg.projection_dim = 32\n    cfg.tokens_per_slot = (2, 2, 2, 2, 2, 2)\n    cfg.checkpoint_epochs = (1,)\n    cfg.selected_adapt_epoch = 1\n    cfg.encode_batch = 12\n    cfg.require_two_gpus = False\n    cfg.time_budget_minutes = 30\n    cfg.stop_margin_minutes = 1\nOUT = Path(cfg.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\")\nrandom.seed(cfg.seed); np.random.seed(cfg.seed); torch.manual_seed(cfg.seed)\ntorch.cuda.manual_seed_all(cfg.seed)\nprint(json.dumps(asdict(cfg), indent=2))\nprint(\"mode:\", MODE, \"gpus:\", torch.cuda.device_count())\nif not cfg.smoke and cfg.require_two_gpus and torch.cuda.device_count() < 2:\n    raise RuntimeError(\"Select Kaggle accelerator GPU T4 x2 before running Stage 02\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SLOT_SPECS = [\n    (\"sag_fluid\", \"Sagittal\", \"fluid\"),\n    (\"sag_struct\", \"Sagittal\", \"struct\"),\n    (\"cor_fluid\", \"Coronal\", \"fluid\"),\n    (\"cor_struct\", \"Coronal\", \"struct\"),\n    (\"ax_fluid\", \"Axial\", \"fluid\"),\n    (\"ax_struct\", \"Axial\", \"struct\"),\n]\nassert len(SLOT_SPECS) == len(cfg.tokens_per_slot) == 6\nif not cfg.smoke:\n    assert sum(cfg.tokens_per_slot) == 64\n\ndef find_comp_root() -> Path:\n    if cfg.comp_root:\n        path = Path(cfg.comp_root)\n        if (path / \"train.csv\").is_file():\n            return path\n        raise FileNotFoundError(f\"Invalid RSNA_COMP_ROOT: {path}\")\n    if cfg.smoke:\n        path = Path.cwd() / \"smoke_rsna_data\"\n        if (path / \"train.csv\").is_file():\n            return path\n        raise FileNotFoundError(\"Run Stage 0.5 smoke before Stage 2 smoke\")\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 / \"train.csv\").is_file():\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 / \"train.csv\").is_file():\n                return path\n    raise FileNotFoundError(\"Attach the competition data or set RSNA_COMP_ROOT\")\n\nROOT = find_comp_root()\ntrain_df = pd.read_csv(ROOT / \"train.csv\", dtype={UID: str})\nif train_df[UID].duplicated().any():\n    raise RuntimeError(\"train.csv contains duplicate StudyInstanceUID values\")\nordered_uids = train_df[UID].astype(str).tolist()\ntrain_uid_sha256 = hashlib.sha256(\"\\n\".join(ordered_uids).encode()).hexdigest()\n\ndef preprocessing_signature() -> Tuple[str, dict]:\n    payload = {\n        \"version\": \"w01-stage05-cache-v1\",\n        \"image_size\": cfg.image_size,\n        \"input_mode\": cfg.input_mode,\n        \"crop_mm\": cfg.crop_mm,\n        \"triplet_offset\": cfg.triplet_offset,\n        \"coverage\": list(cfg.coverage),\n        \"tokens_per_slot\": list(cfg.tokens_per_slot),\n        \"slot_specs\": SLOT_SPECS,\n        \"normalization\": \"series-percentile-1-99-subsample4\",\n        \"series_selection\": \"w01-six-slot-v1\",\n        \"orientation\": \"dicom-lps-canonical-v1\",\n    }\n    digest = hashlib.sha256(json.dumps(payload, sort_keys=True).encode()).hexdigest()\n    return digest, payload\n\nPREPROCESS_SIGNATURE, PREPROCESS_PAYLOAD = preprocessing_signature()\n\ndef configured_roots(value: Optional[str]) -> List[Path]:\n    if not value:\n        return []\n    return [Path(x.strip()) for x in re.split(r\"[;,]\", value) if x.strip()]\n\ndef sha256_file(path: Path) -> str:\n    digest = hashlib.sha256()\n    with path.open(\"rb\") as handle:\n        for block in iter(lambda: handle.read(8 << 20), b\"\"):\n            digest.update(block)\n    return digest.hexdigest()\n\ndef discover_cache_manifests() -> List[Path]:\n    roots = configured_roots(cfg.cache_dirs) or [Path(\"/kaggle/input\"), Path.cwd()]\n    found = []\n    for root in roots:\n        if root.is_file() and root.name == CACHE_MANIFEST:\n            found.append(root)\n        elif root.is_dir():\n            found.extend(root.rglob(CACHE_MANIFEST))\n    return sorted(set(path.resolve() for path in found))\n\ncache_index: Dict[str, Tuple[Path, str]] = {}\ncompatible_manifests, ignored_manifests = [], []\nfor manifest_path in discover_cache_manifests():\n    try:\n        manifest = json.loads(manifest_path.read_text(encoding=\"utf-8\"))\n    except Exception:\n        ignored_manifests.append(str(manifest_path)); continue\n    if (manifest.get(\"preprocess_signature\") != PREPROCESS_SIGNATURE or\n            manifest.get(\"train_uid_sha256\") != train_uid_sha256):\n        ignored_manifests.append(str(manifest_path)); continue\n    compatible_manifests.append(str(manifest_path))\n    for shard in manifest.get(\"shards\", []):\n        shard_path = manifest_path.parent / shard[\"file\"]\n        if not shard_path.is_file():\n            raise FileNotFoundError(f\"cache manifest references missing shard: {shard_path}\")\n        if cfg.verify_cache_hashes and sha256_file(shard_path) != shard.get(\"sha256\"):\n            raise RuntimeError(f\"cache shard checksum mismatch: {shard_path}\")\n        for item in shard.get(\"studies\", []):\n            uid = str(item[\"uid\"])\n            location = (shard_path, str(item[\"key\"]))\n            if uid in cache_index and cache_index[uid] != location:\n                raise RuntimeError(f\"duplicate StudyInstanceUID across Stage 0.5 parts: {uid}\")\n            cache_index[uid] = location\n\nmissing = [uid for uid in ordered_uids if uid not in cache_index]\nextra = sorted(set(cache_index) - set(ordered_uids))\nif missing or extra:\n    raise RuntimeError(\n        f\"Stage 0.5 cache is incomplete/incompatible: cached={len(cache_index)} \"\n        f\"expected={len(ordered_uids)} missing={missing[:10]} extra={extra[:10]}. \"\n        \"Attach every Stage 0.5 output part and set RSNA_CACHE_DIRS if needed.\"\n    )\nprint(f\"Stage 0.5 cache ready: {len(cache_index)} studies from \"\n      f\"{len(compatible_manifests)} manifest(s)\")\nif ignored_manifests:\n    print(\"ignored incompatible cache manifests:\", ignored_manifests)\n\nclass CachedStudyPixels(Dataset):\n    def __init__(self):\n        self._handles = {}\n        self.max_tokens = sum(cfg.tokens_per_slot)\n\n    def __len__(self):\n        return len(ordered_uids)\n\n    def __getstate__(self):\n        state = self.__dict__.copy(); state[\"_handles\"] = {}\n        return state\n\n    def _open(self, path: Path):\n        key = str(path)\n        handle = self._handles.get(key)\n        if handle is None:\n            handle = h5py.File(path, \"r\")\n            if handle.attrs.get(\"preprocess_signature\", \"\") != PREPROCESS_SIGNATURE:\n                handle.close()\n                raise RuntimeError(f\"cache signature mismatch: {path}\")\n            self._handles[key] = handle\n        return handle\n\n    def __getitem__(self, idx: int):\n        uid = ordered_uids[int(idx)]\n        shard_path, group_key = cache_index[uid]\n        group = self._open(shard_path)[group_key]\n        if str(group.attrs.get(\"uid\", \"\")) != uid:\n            raise RuntimeError(f\"cache UID mismatch at {shard_path}:{group_key}\")\n        images = group[\"images\"][:]\n        mask = group[\"mask\"][:].astype(bool)\n        slot = group[\"slot\"][:].astype(np.int64)\n        z = group[\"z\"][:].astype(np.float32)\n        expected = (self.max_tokens, 3, cfg.image_size, cfg.image_size)\n        if images.shape != expected:\n            raise RuntimeError(f\"cached image shape mismatch for {uid}: {images.shape} != {expected}\")\n        return {\n            \"index\": int(idx), \"uid\": uid,\n            \"images\": torch.from_numpy(images),\n            \"mask\": torch.from_numpy(mask),\n            \"slot\": torch.from_numpy(slot),\n            \"z\": torch.from_numpy(z),\n        }\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Frozen encoder\n\nStage 2 does not decode DICOM again. It streams the dense uint8 tokens from the\nStage 0.5 shards, while each GPU owns an independent frozen encoder replica.\n","metadata":{}},{"cell_type":"code","source":"\nclass TinyEncoder(nn.Module):\n    \"\"\"CPU smoke surrogate with the same regional output contract.\"\"\"\n    name = \"tiny\"\n    source_dim = 64\n    region_dim = 32\n\n    def __init__(self):\n        super().__init__()\n        self.n_regions = 1 + cfg.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, (cfg.spatial_grid, cfg.spatial_grid)).flatten(2).transpose(1, 2)\n        return self.project_norm(self.projector(torch.cat([cls, spatial], 1) @ self.base_proj))\n\n\ndef _unwrap_state(obj):\n    if isinstance(obj, dict):\n        for key in (\"state_dict\", \"model\", \"teacher\", \"student\"):\n            if key in obj and isinstance(obj[key], dict):\n                return _unwrap_state(obj[key])\n    return obj\n\n\ndef build_dinov2_backbone(pretrained: bool) -> nn.Module:\n    import timm\n\n    model = timm.create_model(\n        cfg.dinov2_model,\n        pretrained=pretrained,\n        img_size=cfg.image_size,\n        num_classes=0,\n        dynamic_img_size=False,\n    )\n\n    return model\n\n\nclass DinoV2SmallEncoder(nn.Module):\n    name = \"dinov2_small\"\n    source_dim = 384\n\n    def __init__(self, pretrained: bool):\n        super().__init__()\n        self.model = build_dinov2_backbone(pretrained)\n        self.region_dim = int(cfg.projection_dim)\n        self.n_regions = 1 + cfg.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        nn.init.orthogonal_(self.projector.weight)\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        if not cfg.smoke and side != cfg.image_size // 14:\n            raise RuntimeError(f\"expected {cfg.image_size // 14}x{cfg.image_size // 14} patches, got {side}x{side}\")\n        grid = patch.reshape(len(patch), side, side, self.source_dim).permute(0, 3, 1, 2)\n        spatial = F.adaptive_avg_pool2d(grid, (cfg.spatial_grid, cfg.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() -> Dict[str, nn.Module]:\n    if cfg.smoke:\n        return {\"tiny\": TinyEncoder().to(DEVICE).eval()}\n    enc = DinoV2SmallEncoder(pretrained=False).to(DEVICE)\n    return {enc.name: enc}\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The complete DINOv2-S state comes from the selected Stage 1 checkpoint. No online\npretrained download is required in Stage 2.\n","metadata":{}},{"cell_type":"markdown","source":"## Select the Stage 1 checkpoint and extract the dense feature cache\n\nSelection priority is: `RSNA_ADAPTED_CHECKPOINT` (exact file), then\n`RSNA_ADAPT_DIR/w01_epoch{RSNA_ADAPT_EPOCH}.pt`, then an unambiguous search.\n","metadata":{}},{"cell_type":"code","source":"def resolve_selected_checkpoint() -> Path:\n    filename = f\"w01_epoch{cfg.selected_adapt_epoch}.pt\"\n    if cfg.adapted_checkpoint:\n        path = Path(cfg.adapted_checkpoint)\n        path = path / filename if path.is_dir() else path\n        if not path.is_file():\n            raise FileNotFoundError(f\"RSNA_ADAPTED_CHECKPOINT does not exist: {path}\")\n        return path\n    if cfg.selected_adapt_epoch not in set(cfg.checkpoint_epochs):\n        raise ValueError(\n            f\"RSNA_ADAPT_EPOCH must be one of {cfg.checkpoint_epochs}; \"\n            f\"got {cfg.selected_adapt_epoch}\"\n        )\n    if cfg.adapt_dir:\n        path = Path(cfg.adapt_dir) / filename\n        if not path.is_file():\n            raise FileNotFoundError(f\"Selected Stage 1 checkpoint not found: {path}\")\n        return path\n    found = []\n    for base in (Path(\"/kaggle/input\"), Path.cwd()):\n        if base.exists():\n            found.extend(base.rglob(filename))\n    found = sorted(set(found), key=str)\n    if len(found) != 1:\n        raise RuntimeError(\n            f\"Expected exactly one {filename}; got {found}. Set \"\n            \"RSNA_ADAPTED_CHECKPOINT to the exact .pt file or RSNA_ADAPT_DIR \"\n            \"to the Stage 1 dataset directory.\"\n        )\n    return found[0]\n\ndef load_pt(path: Path):\n    try:\n        return torch.load(path, map_location=\"cpu\", weights_only=False)\n    except TypeError:\n        return torch.load(path, map_location=\"cpu\")\n\ndef encoder_state_signature(state: dict) -> str:\n    digest = hashlib.sha256()\n    for key in sorted(state):\n        value = state[key]\n        digest.update(key.encode())\n        digest.update(value.detach().cpu().contiguous().numpy().tobytes())\n    return digest.hexdigest()\n\ncheckpoint_path = resolve_selected_checkpoint()\ncheckpoint = load_pt(checkpoint_path)\nselected_epoch = int(checkpoint.get(\"epoch_completed\", -1))\nif selected_epoch not in set(cfg.checkpoint_epochs):\n    raise RuntimeError(f\"Stage 2 accepts checkpoint epochs {cfg.checkpoint_epochs}; got {selected_epoch}\")\nif selected_epoch != cfg.selected_adapt_epoch:\n    raise RuntimeError(\n        f\"Requested epoch {cfg.selected_adapt_epoch}, but checkpoint contains epoch {selected_epoch}\"\n    )\nif checkpoint.get(\"preprocess_signature\") != PREPROCESS_SIGNATURE:\n    raise RuntimeError(\"Stage 1 checkpoint and Stage 0.5 preprocessing signatures differ\")\nif \"encoder_state\" not in checkpoint:\n    raise RuntimeError(\"Selected Stage 1 checkpoint does not contain encoder_state\")\ncheckpoint_signature = encoder_state_signature(checkpoint[\"encoder_state\"])\n\nencoders = build_encoders()\nname, primary = next(iter(encoders.items()))\nresult = primary.load_state_dict(checkpoint[\"encoder_state\"], strict=True)\nif result.missing_keys or result.unexpected_keys:\n    raise RuntimeError((result.missing_keys[:10], result.unexpected_keys[:10]))\nprimary.eval()\nprint(\"MODEL_SELECTION:\", {\n    \"path\": str(checkpoint_path), \"epoch\": selected_epoch,\n    \"encoder\": name, \"state_sha256\": checkpoint_signature,\n})\n\ndataset = CachedStudyPixels()\nfeature_path = OUT / f\"train_features_{name}.npy\"\nmask_path, slot_path, z_path, progress_path = [OUT / value for value in (\n    \"train_token_mask.npy\", \"train_slot_ids.npy\", \"train_zpos.npy\", \"feature_progress.npy\"\n)]\n\nif cfg.feature_resume_dir:\n    resume_dir = Path(cfg.feature_resume_dir)\n    prior_manifest_path = resume_dir / \"feature_manifest.json\"\n    if not prior_manifest_path.is_file():\n        raise FileNotFoundError(prior_manifest_path)\n    prior = json.loads(prior_manifest_path.read_text(encoding=\"utf-8\"))\n    expected_resume = {\n        \"mode\": MODE,\n        \"checkpoint_signature\": checkpoint_signature,\n        \"preprocess_signature\": PREPROCESS_SIGNATURE,\n        \"train_uid_sha256\": train_uid_sha256,\n    }\n    mismatch = {key: (prior.get(key), value) for key, value in expected_resume.items()\n                if prior.get(key) != value}\n    if mismatch:\n        raise RuntimeError(f\"Incompatible Stage 2 resume cache: {mismatch}\")\n    for path in (feature_path, mask_path, slot_path, z_path, progress_path):\n        candidate = resume_dir / path.name\n        if candidate.is_file() and not path.is_file():\n            shutil.copy2(candidate, path)\n\nshape = (len(dataset), dataset.max_tokens, int(primary.n_regions), int(primary.region_dim))\nspecs = [\n    (feature_path, np.float16, shape),\n    (mask_path, np.bool_, shape[:2]),\n    (slot_path, np.int8, shape[:2]),\n    (z_path, np.float16, shape[:2]),\n    (progress_path, np.bool_, (len(dataset),)),\n]\nfor path, dtype, expected_shape in specs:\n    if not path.is_file():\n        np.lib.format.open_memmap(path, mode=\"w+\", dtype=dtype, shape=expected_shape).flush()\n    actual = np.load(path, mmap_mode=\"r\")\n    if actual.shape != expected_shape or actual.dtype != np.dtype(dtype):\n        raise RuntimeError(\n            f\"Resume file contract mismatch for {path}: \"\n            f\"shape={actual.shape}, dtype={actual.dtype}; expected {expected_shape}, {np.dtype(dtype)}\"\n        )\n\ndef replicate(device: torch.device):\n    if name == \"tiny\":\n        replica = TinyEncoder()\n    elif name == \"dinov2_small\":\n        replica = DinoV2SmallEncoder(pretrained=False)\n    else:\n        raise KeyError(name)\n    replica.load_state_dict(checkpoint[\"encoder_state\"], strict=True)\n    return replica.to(device).eval()\n\ndevices = [torch.device(f\"cuda:{i}\") for i in range(min(2, torch.cuda.device_count()))]\nif not devices:\n    devices = [torch.device(\"cpu\"), torch.device(\"cpu\")] if cfg.smoke else [torch.device(\"cpu\")]\nif not cfg.smoke and cfg.require_two_gpus and len(devices) < 2:\n    raise RuntimeError(\"Stage 2 requires Kaggle T4 x2\")\nreplicas = [primary.to(devices[0]).eval()] + [replicate(device) for device in devices[1:]]\npending_all = np.flatnonzero(~np.asarray(np.load(progress_path, mmap_mode=\"r\")))\npending_shards = [pending_all[i::len(devices)] for i in range(len(devices))]\nstarted = time.time()\nstop = threading.Event()\nprint_lock = threading.Lock()\n\n@torch.inference_mode()\ndef encode_valid(enc, images, valid, device):\n    out = np.zeros((len(valid), int(enc.n_regions), int(enc.region_dim)), np.float32)\n    batch_size, start = cfg.encode_batch, 0\n    while start < len(valid):\n        positions = valid[start:start + batch_size]\n        try:\n            with torch.autocast(\"cuda\", enabled=device.type == \"cuda\"):\n                value = enc(images[positions].to(device, non_blocking=True)).float().cpu().numpy()\n            out[start:start + len(positions)] = value\n            start += len(positions)\n        except torch.cuda.OutOfMemoryError:\n            if batch_size <= 1:\n                raise\n            batch_size = max(1, batch_size // 2)\n            if device.type == \"cuda\":\n                torch.cuda.empty_cache()\n            with print_lock:\n                print(f\"CUDA OOM on {device}: reducing encode batch to {batch_size}\")\n    return out\n\ndef worker(worker_id: int):\n    device, enc = devices[worker_id], replicas[worker_id]\n    if device.type == \"cuda\":\n        torch.cuda.set_device(device)\n    local_ds = CachedStudyPixels()\n    fmap = np.lib.format.open_memmap(feature_path, mode=\"r+\")\n    mmap = np.lib.format.open_memmap(mask_path, mode=\"r+\")\n    smap = np.lib.format.open_memmap(slot_path, mode=\"r+\")\n    zmap = np.lib.format.open_memmap(z_path, mode=\"r+\")\n    pmap = np.lib.format.open_memmap(progress_path, mode=\"r+\")\n    completed = 0\n    pending = pending_shards[worker_id]\n    for step, idx in enumerate(pending):\n        elapsed = (time.time() - started) / 60\n        if stop.is_set() or elapsed >= cfg.time_budget_minutes - cfg.stop_margin_minutes:\n            stop.set(); break\n        record = local_ds[int(idx)]\n        mask = record[\"mask\"].numpy()\n        valid = np.flatnonzero(mask)\n        fmap[idx] = 0\n        if len(valid):\n            fmap[idx, valid] = encode_valid(\n                enc, record[\"images\"], valid, device\n            ).astype(np.float16)\n        mmap[idx] = mask\n        smap[idx] = record[\"slot\"].numpy()\n        zmap[idx] = record[\"z\"].numpy()\n        pmap[idx] = True\n        completed += 1\n        if completed % 16 == 0:\n            for array in (fmap, mmap, smap, zmap, pmap):\n                array.flush()\n        if step % max(1, len(pending) // 10) == 0:\n            with print_lock:\n                print(f\"worker={worker_id} device={device} {step + 1}/{len(pending)} \"\n                      f\"elapsed={elapsed:.1f}m\")\n    for array in (fmap, mmap, smap, zmap, pmap):\n        array.flush()\n    return completed\n\nwith ThreadPoolExecutor(max_workers=len(devices), thread_name_prefix=\"w01-feature\") as pool:\n    worker_counts = list(pool.map(worker, range(len(devices))))\n\nprogress = np.load(progress_path, mmap_mode=\"r\")\ncomplete = bool(progress.all())\nuid_path = OUT / \"feature_uids.csv\"\npd.DataFrame({UID: ordered_uids, \"row\": np.arange(len(ordered_uids))}).to_csv(uid_path, index=False)\nmanifest = {\n    \"version\": \"w01-feature-cache-stage05-2\", \"mode\": MODE, \"complete\": complete,\n    \"completed_studies\": int(progress.sum()), \"n_studies\": len(progress),\n    \"feature_shape\": list(shape), \"feature_dtype\": \"float16\",\n    \"checkpoint\": checkpoint_path.name, \"checkpoint_signature\": checkpoint_signature,\n    \"selected_adapt_epoch\": selected_epoch,\n    \"preprocess_signature\": PREPROCESS_SIGNATURE,\n    \"preprocess_payload\": PREPROCESS_PAYLOAD,\n    \"train_uid_sha256\": train_uid_sha256,\n    \"stage05_manifests\": len(compatible_manifests),\n    \"devices\": [str(device) for device in devices],\n    \"worker_completed_this_run\": worker_counts,\n    \"encode_batch_start\": cfg.encode_batch,\n    \"elapsed_minutes\": (time.time() - started) / 60,\n    \"files\": [feature_path.name, mask_path.name, slot_path.name,\n              z_path.name, progress_path.name, uid_path.name],\n}\nmanifest_path = OUT / \"feature_manifest.json\"\ntemporary = manifest_path.with_suffix(\".json.tmp\")\ntemporary.write_text(json.dumps(manifest, indent=2), encoding=\"utf-8\")\nos.replace(temporary, manifest_path)\nprint(json.dumps(manifest, indent=2))\nif cfg.smoke:\n    assert complete and tuple(np.load(feature_path, mmap_mode=\"r\").shape) == shape\n    print(\"STAGE02_SMOKE_PASS\")\nprint(\"STAGE02_COMPLETE\" if complete else\n      \"STAGE02_GRACEFUL_PARTIAL: attach output and set RSNA_FEATURE_RESUME_DIR next run\")\n","metadata":{},"outputs":[],"execution_count":null}]}