{"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":"de574507-26a3-49dd-a783-b53cb2a8156e","cell_type":"markdown","source":"# RSNA Knee D0 — simple ResNet34 2.5D approximate resume\n\nThis is the deliberately small diagnostic pipeline suggested by the Top-13 discussion. It is a data/supervision probe, not an attempt to reproduce the previous large Transformer system:\n\n- consumes the unchanged `w01_supervision.csv` from W01 Stage 0;\n- consumes every HDF5 part from W01 Stage 0.5;\n- resizes cached 336 px triplets to 224 px on GPU;\n- one fully trainable ImageNet ResNet34 shared by every 2.5D triplet;\n- no attention, Transformer, report embedding, multi-seed, or model ensemble;\n- mean + max pooling inside each of six MRI series slots; the six slot descriptors are concatenated and classified by one small MLP;\n- all 4,349 pseudo-labeled studies for training and all 58 official gold studies as an untouched clean validation set;\n- the best pseudo-trained checkpoint is exported directly; there is no gold oversampling or full-data refit.\n\nDefault D0 reads two cached candidates per slot (12 triplets/study). A study batch of 64 therefore means 768 triplet images before the two-GPU split; it is not an image batch of 64. Keep the two-token coverage unchanged for the first run. A later coverage ablation can set `RSNA_SIMPLE_TOKENS_PER_SLOT=4` without rebuilding Stage 0.5.\n\nThis resume edition auto-discovers exactly one attached `simple_resnet34_cv_best.pt` (or accepts `RSNA_RESUME_CHECKPOINT`), restores the encoder/head, recomputes the clean-gold baseline, and continues from `saved_epoch + 1` through epoch 10. AdamW, GradScaler, and RNG state are intentionally restarted because the interrupted checkpoint contains model weights only; the cosine learning-rate schedule still uses the original epoch number. The default study batch is the stable value 64, so no risky batch-128 probe is repeated.\n\nFor an offline Kaggle run, attach the official torchvision file `resnet34-b627a593.pth` and set `RSNA_RESNET34_WEIGHTS` if auto-discovery is ambiguous. Internet is only needed when that file is not attached.\n","metadata":{}},{"id":"85f9f053-fd07-415f-b136-9d78c00c47e1","cell_type":"code","source":"from __future__ import annotations\n\nimport gc, hashlib, json, math, os, random, re, time\nfrom contextlib import nullcontext\nfrom dataclasses import asdict, dataclass\nfrom pathlib import Path\nfrom typing import Dict, List, Optional, Tuple\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 sklearn.metrics import roc_auc_score\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision.models import ResNet34_Weights, 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\"\nSLOT_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]\n\n\ndef repeated_int_tuple(value: str, n: int = 6) -> Tuple[int, ...]:\n    parts = [int(x.strip()) for x in value.split(\",\") if x.strip()]\n    if len(parts) == 1:\n        parts *= n\n    if len(parts) != n or min(parts) < 1:\n        raise ValueError(f\"Expected one positive integer or {n} comma-separated integers, got {value!r}\")\n    return tuple(parts)\n\n\ndef descending_positive_int_tuple(value: str) -> Tuple[int, ...]:\n    parts = tuple(int(x.strip()) for x in value.split(\",\") if x.strip())\n    if not parts or min(parts) < 1:\n        raise ValueError(f\"Expected positive comma-separated batch candidates, got {value!r}\")\n    if any(a <= b for a, b in zip(parts, parts[1:])):\n        raise ValueError(f\"Batch candidates must be strictly descending, got {parts}\")\n    return parts\n\n\ndef optional_float_env(name: str) -> Optional[float]:\n    value = os.getenv(name)\n    return None if value is None or not value.strip() else float(value)\n\n\n@dataclass\nclass CFG:\n    seed: int = int(os.getenv(\"RSNA_SEED\", \"2026\"))\n    smoke: bool = os.getenv(\"RSNA_SMOKE_TEST\", \"0\") == \"1\"\n    comp_root: Optional[str] = os.getenv(\"RSNA_COMP_ROOT\")\n    supervision_dir: Optional[str] = os.getenv(\"RSNA_SUPERVISION_DIR\")\n    cache_dirs: Optional[str] = os.getenv(\"RSNA_CACHE_DIRS\")\n    resnet34_weights: Optional[str] = os.getenv(\"RSNA_RESNET34_WEIGHTS\")\n    resume_checkpoint: Optional[str] = os.getenv(\"RSNA_RESUME_CHECKPOINT\")\n    require_resume: bool = os.getenv(\"RSNA_REQUIRE_RESUME\", \"1\") == \"1\"\n    output_dir: str = os.getenv(\"RSNA_SIMPLE_OUTPUT\", \"/kaggle/working/rsna_simple_resnet34_resume\")\n    cache_image_size: int = 336\n    model_image_size: int = 224\n    input_mode: str = \"adjacent\"\n    crop_mm: float = 140.0\n    triplet_offset: int = 1\n    coverage: Tuple[float, float] = (0.02, 0.98)\n    dense_tokens_per_slot: Tuple[int, ...] = (16, 12, 12, 8, 10, 6)\n    selected_tokens_per_slot: Tuple[int, ...] = repeated_int_tuple(\n        os.getenv(\"RSNA_SIMPLE_TOKENS_PER_SLOT\", \"2\")\n    )\n    # Resume defaults to the stable batch 64 after batch 128 produced a late CUDA launch failure.\n    study_batch_size: int = int(os.getenv(\"RSNA_STUDY_BATCH_SIZE\", \"64\"))\n    batch_candidates: Tuple[int, ...] = descending_positive_int_tuple(\n        os.getenv(\"RSNA_BATCH_CANDIDATES\", \"128,64,32,16,8,4,2,1\")\n    )\n    grad_accum: int = int(os.getenv(\"RSNA_GRAD_ACCUM\", \"1\"))\n    cache_workers: int = int(os.getenv(\"RSNA_CACHE_WORKERS\", \"4\"))\n    cv_epochs: int = int(os.getenv(\"RSNA_CV_EPOCHS\", \"10\"))\n    reference_gold_auc: Optional[float] = optional_float_env(\"RSNA_REFERENCE_GOLD_AUC\")\n    hidden_dim: int = 512\n    dropout: float = 0.20\n    encoder_lr: float = float(os.getenv(\"RSNA_ENCODER_LR\", \"1e-4\"))\n    head_lr: float = float(os.getenv(\"RSNA_HEAD_LR\", \"5e-4\"))\n    weight_decay: float = 1e-4\n    max_pos_weight: float = 6.0\n    decision_min_delta: float = 0.003\n    verify_cache_hashes: bool = os.getenv(\"RSNA_VERIFY_CACHE_HASHES\", \"0\") == \"1\"\n    require_two_gpus: bool = os.getenv(\"RSNA_REQUIRE_2GPU\", \"1\") == \"1\"\n\n\ncfg = CFG()\nif cfg.smoke:\n    cfg.output_dir = os.getenv(\"RSNA_SIMPLE_OUTPUT\", str(Path.cwd() / \"smoke_simple_resnet34\"))\n    cfg.cache_image_size = 64\n    cfg.model_image_size = 64\n    cfg.dense_tokens_per_slot = (2, 2, 2, 2, 2, 2)\n    cfg.selected_tokens_per_slot = (1, 1, 1, 1, 1, 1)\n    cfg.study_batch_size = 2\n    cfg.grad_accum = 1\n    cfg.cache_workers = 0\n    cfg.cv_epochs = 1\n    cfg.hidden_dim = 64\n    cfg.require_two_gpus = False\n    cfg.require_resume = False\n\nif any(a > b for a, b in zip(cfg.selected_tokens_per_slot, cfg.dense_tokens_per_slot)):\n    raise ValueError(\"selected_tokens_per_slot cannot exceed the Stage 0.5 dense cache quotas\")\nif cfg.study_batch_size < 0 or cfg.grad_accum < 1:\n    raise ValueError(\"RSNA_STUDY_BATCH_SIZE must be >= 0 and RSNA_GRAD_ACCUM must be >= 1\")\n\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\")\n\n\ndef seed_everything(seed: int):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    if torch.cuda.is_available():\n        torch.backends.cudnn.benchmark = True\n\n\nseed_everything(cfg.seed)\nprint(json.dumps(asdict(cfg), indent=2))\nprint(\"device:\", DEVICE, \"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, or set RSNA_REQUIRE_2GPU=0\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"b7ba734c-3a40-4d10-b58c-ad4b926359ab","cell_type":"markdown","source":"## 1. Resolve competition data, immutable supervision, and every Stage 0.5 cache part\n\nThis cell refuses an incomplete cache, a UID-order mismatch, or a preprocessing-signature mismatch. Those checks are intentionally stricter than ordinary notebook auto-discovery.\n","metadata":{}},{"id":"8bc19c66-7daf-4022-94b6-3f6e590af524","cell_type":"code","source":"def 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 the W01 Stage 0.5 smoke notebook first\")\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\n\ndef find_unique(filename: str, configured: Optional[str]) -> Path:\n    if configured:\n        path = Path(configured)\n        path = path if path.is_file() else path / filename\n        if path.is_file():\n            return path\n        raise FileNotFoundError(path)\n    roots = [Path.cwd(), Path(\"/kaggle/input\")]\n    found = []\n    for base in roots:\n        if not base.exists():\n            continue\n        found.extend(\n            path for path in base.rglob(filename)\n            if \"train_series\" not in path.parts and \"test_series\" not in path.parts\n        )\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 an explicit path.\")\n    return found[0]\n\n\nROOT = find_comp_root()\ntrain_df = pd.read_csv(ROOT / \"train.csv\", dtype={UID: str})\nfor target in TARGETS:\n    train_df[target] = pd.to_numeric(train_df.get(target), errors=\"coerce\")\nif train_df[UID].isna().any() or train_df[UID].duplicated().any():\n    raise RuntimeError(\"train.csv must contain unique, non-null StudyInstanceUID values\")\n\nsupervision_path = find_unique(\"w01_supervision.csv\", cfg.supervision_dir)\nsupervision = pd.read_csv(supervision_path, dtype={UID: str})\nrequired = (\n    [UID, \"is_gold\", \"report_hash\"]\n    + [f\"y__{target}\" for target in TARGETS]\n    + [f\"w__{target}\" for target in TARGETS]\n)\nmissing_columns = [column for column in required if column not in supervision.columns]\nif missing_columns:\n    raise RuntimeError(f\"Stage 0 supervision is missing columns: {missing_columns}\")\nif supervision[UID].tolist() != train_df[UID].astype(str).tolist():\n    raise RuntimeError(\"Stage 0 supervision UID order differs from competition train.csv\")\n\nY = supervision[[f\"y__{target}\" for target in TARGETS]].to_numpy(np.float32)\nW = supervision[[f\"w__{target}\" for target in TARGETS]].to_numpy(np.float32)\nGOLD_Y = train_df[TARGETS].to_numpy(np.float32)\ngold_mask = supervision[\"is_gold\"].astype(bool).to_numpy()\nordered_uids = train_df[UID].astype(str).tolist()\n\nif not cfg.smoke and (int(gold_mask.sum()) != 58 or int((~gold_mask).sum()) != 4349):\n    raise RuntimeError(f\"Expected 58 gold + 4349 pseudo studies, got {gold_mask.sum()} + {(~gold_mask).sum()}\")\ngold_values = GOLD_Y[gold_mask]\nif not np.isfinite(gold_values).all() or not np.isin(gold_values, [0.0, 1.0]).all():\n    raise RuntimeError(\"The 58 clean validation studies must have complete binary labels in train.csv\")\n\n\ndef preprocessing_signature() -> Tuple[str, dict]:\n    payload = {\n        \"version\": \"w01-stage05-cache-v1\",\n        \"image_size\": cfg.cache_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.dense_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\n\nMANIFEST_NAME = \"w01_stage05_manifest.json\"\nPREPROCESS_SIGNATURE, PREPROCESS_PAYLOAD = preprocessing_signature()\ntrain_uid_sha256 = hashlib.sha256(\"\\n\".join(ordered_uids).encode()).hexdigest()\n\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\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 == MANIFEST_NAME:\n            found.append(root)\n        elif root.is_dir():\n            found.extend(root.rglob(MANIFEST_NAME))\n    return sorted(set(path.resolve() for path in found), key=str)\n\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\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))\n        continue\n    if (\n        manifest.get(\"preprocess_signature\") != PREPROCESS_SIGNATURE\n        or manifest.get(\"train_uid_sha256\") != train_uid_sha256\n    ):\n        ignored_manifests.append(str(manifest_path))\n        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 cache parts: {uid}\")\n            cache_index[uid] = location\n\nmissing_uids = [uid for uid in ordered_uids if uid not in cache_index]\nextra_uids = sorted(set(cache_index) - set(ordered_uids))\nif missing_uids or extra_uids:\n    raise RuntimeError(\n        f\"Stage 0.5 cache incomplete/incompatible: cached={len(cache_index)} expected={len(ordered_uids)} \"\n        f\"missing={missing_uids[:10]} extra={extra_uids[:10]}. Attach every cache part.\"\n    )\n\nprint(\"competition root:\", ROOT)\nprint(\"supervision:\", supervision_path)\nprint(f\"cache ready: {len(cache_index)} studies from {len(compatible_manifests)} manifests\")\nif ignored_manifests:\n    print(\"ignored incompatible manifests:\", ignored_manifests)\nprint(\"preprocess signature:\", PREPROCESS_SIGNATURE)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"b2e78beb-59f9-4026-9f08-451a63905b6f","cell_type":"markdown","source":"## 2. Contract audit and clean 58-study gold validation split\n\nThis is a light cache-contract audit, not the later geometry-debug ablation. Training uses all 4,349 pseudo studies and never exposes the 58 gold studies to the optimizer. Gold labels are read directly from `train.csv`, used without pseudo confidence weights, and serve only for checkpoint selection and the ≥0.003 submit decision.\n","metadata":{}},{"id":"0de77662-d265-4c77-a41d-3b68aa351291","cell_type":"code","source":"def 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\naudit_rows = []\nfor uid in ordered_uids[: min(16, len(ordered_uids))]:\n    shard_path, group_key = cache_index[uid]\n    with h5py.File(shard_path, \"r\") as handle:\n        if handle.attrs.get(\"preprocess_signature\", \"\") != PREPROCESS_SIGNATURE:\n            raise RuntimeError(f\"HDF5 signature mismatch: {shard_path}\")\n        group = handle[group_key]\n        images = group[\"images\"]\n        mask = group[\"mask\"][:].astype(bool)\n        slot = group[\"slot\"][:].astype(np.int64)\n        if images.shape != (sum(cfg.dense_tokens_per_slot), 3, cfg.cache_image_size, cfg.cache_image_size):\n            raise RuntimeError(f\"Unexpected cached image shape for {uid}: {images.shape}\")\n        if images.dtype != np.uint8:\n            raise RuntimeError(f\"Expected uint8 cache, got {images.dtype}\")\n        if not np.array_equal(np.bincount(slot, minlength=6), np.asarray(cfg.dense_tokens_per_slot)):\n            raise RuntimeError(f\"Slot layout mismatch for {uid}\")\n        audit_rows.append({\n            \"uid\": uid,\n            \"valid_total\": int(mask.sum()),\n            **{f\"valid_slot_{s}\": int((mask & (slot == s)).sum()) for s in range(6)},\n            \"sample_min\": int(images[0].min()),\n            \"sample_max\": int(images[0].max()),\n        })\n\ncache_audit = pd.DataFrame(audit_rows)\nprint(cache_audit.to_string(index=False))\ncache_audit.to_csv(OUT / \"cache_contract_audit.csv\", index=False)\n\n\nsplit_role = np.where(gold_mask, \"gold_validation_only\", \"pseudo_diagnostic_train\")\nsplit_table = pd.DataFrame({UID: ordered_uids, \"split_role\": split_role, \"is_gold\": gold_mask})\nsplit_table.to_csv(OUT / \"simple_diagnostic_split.csv\", index=False)\n\ntrain_pseudo_indices = np.flatnonzero(~gold_mask)\ngold_indices = np.flatnonzero(gold_mask)\ncv_train_indices = train_pseudo_indices.copy()\nval_indices = gold_indices.copy()\nprint(\"CV train rows / unique studies:\", len(cv_train_indices), len(np.unique(cv_train_indices)))\nprint(\"Clean gold validation studies:\", len(val_indices))\nprint(\"Gold present in diagnostic training:\", bool(np.intersect1d(cv_train_indices, val_indices).size))\n","metadata":{},"outputs":[],"execution_count":null},{"id":"599a8b1d-c8de-4844-b68a-1dcdffecc72c","cell_type":"markdown","source":"## 3. Cached 2.5D dataset\n\nTraining samples candidate locations randomly inside each slot. Validation uses deterministic bin centers—not the first/last edge slices. Pixel resize and ImageNet normalization happen inside the encoder on GPU.\n","metadata":{}},{"id":"db7a8636-75ed-401e-b71f-98d22fb1965a","cell_type":"code","source":"class CachedStudyDataset(Dataset):\n    def __init__(self, indices: np.ndarray, training: bool):\n        self.indices = np.asarray(indices, dtype=np.int64)\n        self.training = bool(training)\n        self._handles = {}\n        self.n_out = sum(cfg.selected_tokens_per_slot)\n\n    def __len__(self):\n        return len(self.indices)\n\n    def __getstate__(self):\n        state = self.__dict__.copy()\n        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, item: int):\n        study_index = int(self.indices[item])\n        uid = ordered_uids[study_index]\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        dense_mask = group[\"mask\"][:].astype(bool)\n        dense_slot = group[\"slot\"][:].astype(np.int64)\n        selected, out_mask, out_slot = [], [], []\n        for slot_id, quota in enumerate(cfg.selected_tokens_per_slot):\n            candidates = np.flatnonzero(dense_mask & (dense_slot == slot_id))\n            count = min(int(quota), len(candidates))\n            if count:\n                if self.training:\n                    chosen = np.sort(np.random.choice(candidates, count, replace=False))\n                else:\n                    chosen = candidates[center_bin_positions(len(candidates), count)]\n                selected.extend(chosen.tolist())\n                out_mask.extend([True] * count)\n                out_slot.extend([slot_id] * count)\n            for _ in range(int(quota) - count):\n                selected.append(-1)\n                out_mask.append(False)\n                out_slot.append(slot_id)\n\n        images = np.zeros((self.n_out, 3, cfg.cache_image_size, cfg.cache_image_size), np.uint8)\n        positions = [(out_pos, dense_pos) for out_pos, dense_pos in enumerate(selected) if dense_pos >= 0]\n        if positions:\n            dense_positions = [dense_pos for _, dense_pos in positions]\n            values = group[\"images\"][dense_positions]\n            for (out_pos, _), value in zip(positions, values):\n                images[out_pos] = value\n        return {\n            \"index\": study_index,\n            \"images\": torch.from_numpy(images),\n            \"mask\": torch.tensor(out_mask, dtype=torch.bool),\n            \"slot\": torch.tensor(out_slot, dtype=torch.long),\n        }\n\n\ndef seed_worker(worker_id: int):\n    worker_seed = torch.initial_seed() % (2 ** 32)\n    np.random.seed(worker_seed)\n    random.seed(worker_seed)\n\n\ndef make_loader(indices: np.ndarray, training: bool, epoch: int = 0) -> DataLoader:\n    generator = torch.Generator().manual_seed(cfg.seed + 1009 * epoch + (1 if training else 0))\n    kwargs = dict(\n        dataset=CachedStudyDataset(indices, training=training),\n        batch_size=cfg.study_batch_size,\n        shuffle=training,\n        generator=generator,\n        num_workers=cfg.cache_workers,\n        pin_memory=DEVICE.type == \"cuda\",\n        persistent_workers=bool(cfg.cache_workers > 0),\n        worker_init_fn=seed_worker,\n    )\n    if cfg.cache_workers > 0:\n        kwargs[\"prefetch_factor\"] = 2\n    return DataLoader(**kwargs)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"9a9af2ba-7d37-4a30-bd68-9c2310a1aec5","cell_type":"markdown","source":"## 4. ResNet34 encoder and six-slot mean/max head\n\nThe only pooling operations are masked mean and masked max. Concatenating six slot descriptors preserves plane/contrast identity without learned attention.\n","metadata":{}},{"id":"26e5fcbc-a7b1-4ecb-815e-d4bb1656d889","cell_type":"code","source":"def unwrap_state_dict(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_dict(obj[key])\n    return obj\n\n\ndef resolve_resnet34_weights() -> Optional[Path]:\n    if cfg.smoke:\n        return None\n    if cfg.resnet34_weights:\n        path = Path(cfg.resnet34_weights)\n        if not path.is_file():\n            raise FileNotFoundError(path)\n        return path\n    matches = []\n    base = Path(\"/kaggle/input\")\n    if base.is_dir():\n        for path in base.rglob(\"resnet34-b627a593.pth\"):\n            if \"train_series\" not in path.parts and \"test_series\" not in path.parts:\n                matches.append(path)\n    matches = sorted(set(matches), key=str)\n    if len(matches) > 1:\n        raise RuntimeError(f\"Multiple ResNet34 weight files found; set RSNA_RESNET34_WEIGHTS: {matches}\")\n    return matches[0] if matches else None\n\n\nclass ResNet34Encoder(nn.Module):\n    feature_dim = 512\n\n    def __init__(self, load_pretrained: bool):\n        super().__init__()\n        backbone = resnet34(weights=None)\n        self.pretrained_source = \"random\"\n        if load_pretrained and not cfg.smoke:\n            local_path = resolve_resnet34_weights()\n            if local_path is not None:\n                try:\n                    state = torch.load(local_path, map_location=\"cpu\", weights_only=False)\n                except TypeError:\n                    state = torch.load(local_path, map_location=\"cpu\")\n                state = unwrap_state_dict(state)\n                state = {str(k).removeprefix(\"module.\"): v for k, v in state.items()}\n                backbone.load_state_dict(state, strict=True)\n                self.pretrained_source = str(local_path)\n            else:\n                try:\n                    state = ResNet34_Weights.IMAGENET1K_V1.get_state_dict(progress=True)\n                    backbone.load_state_dict(state, strict=True)\n                    self.pretrained_source = \"torchvision:IMAGENET1K_V1\"\n                except Exception as exc:\n                    raise RuntimeError(\n                        \"ImageNet ResNet34 weights are unavailable. Attach resnet34-b627a593.pth \"\n                        \"and set RSNA_RESNET34_WEIGHTS, or enable Internet for training.\"\n                    ) from exc\n        backbone.fc = nn.Identity()\n        self.backbone = backbone\n        mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)\n        std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)\n        self.register_buffer(\"pixel_mean\", mean)\n        self.register_buffer(\"pixel_std\", std)\n\n    def forward(self, images: torch.Tensor) -> torch.Tensor:\n        x = images.float().div_(255.0)\n        if x.shape[-2:] != (cfg.model_image_size, cfg.model_image_size):\n            x = F.interpolate(\n                x,\n                size=(cfg.model_image_size, cfg.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        self.net = nn.Sequential(\n            nn.LayerNorm(fusion_dim),\n            nn.Dropout(cfg.dropout),\n            nn.Linear(fusion_dim, cfg.hidden_dim),\n            nn.GELU(),\n            nn.Dropout(cfg.dropout),\n            nn.Linear(cfg.hidden_dim, len(TARGETS)),\n        )\n\n    def forward(self, features: torch.Tensor, mask: torch.Tensor, slot: torch.Tensor) -> 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\ndef build_modules(load_pretrained: bool = True):\n    encoder = ResNet34Encoder(load_pretrained=load_pretrained).to(DEVICE)\n    head = SixSlotMeanMaxHead().to(DEVICE)\n    runner = (\n        nn.DataParallel(encoder, device_ids=[0, 1])\n        if DEVICE.type == \"cuda\" and torch.cuda.device_count() >= 2\n        else encoder\n    )\n    return encoder, head, runner\n\n\ndef encode_studies(runner, 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        parameter = next(runner.parameters())\n        all_features = parameter.new_zeros((batch * tokens, ResNet34Encoder.feature_dim))\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\n\ndef make_scaler():\n    try:\n        return torch.amp.GradScaler(\"cuda\", enabled=DEVICE.type == \"cuda\")\n    except Exception:\n        return torch.cuda.amp.GradScaler(enabled=DEVICE.type == \"cuda\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"aebcedac-b13a-45c2-9da0-a8d81a2385b5","cell_type":"markdown","source":"## 5. Training and report-label validation utilities\n","metadata":{}},{"id":"19eccee9-5cd0-4a5b-b6e2-60541f2185f9","cell_type":"code","source":"def weighted_macro_auc(\n    truth: np.ndarray,\n    prediction: np.ndarray,\n    weights: Optional[np.ndarray],\n) -> Tuple[float, Dict[str, float]]:\n    scores = {}\n    for j, target in enumerate(TARGETS):\n        y = truth[:, j]\n        p = prediction[:, j]\n        w = None if weights is None else weights[:, j]\n        valid = np.isfinite(y) & np.isfinite(p)\n        if w is not None:\n            valid &= np.isfinite(w) & (w > 0)\n        if valid.sum() < 3 or len(np.unique(y[valid])) < 2:\n            scores[target] = float(\"nan\")\n            continue\n        scores[target] = float(\n            roc_auc_score(y[valid].astype(int), p[valid], sample_weight=None if w is None else w[valid])\n        )\n    finite = [score for score in scores.values() if np.isfinite(score)]\n    return (float(np.mean(finite)) if finite else float(\"nan\")), scores\n\n\ndef make_optimizer(encoder: nn.Module, head: nn.Module):\n    return torch.optim.AdamW(\n        [\n            {\"params\": encoder.parameters(), \"lr\": cfg.encoder_lr, \"name\": \"encoder\"},\n            {\"params\": head.parameters(), \"lr\": cfg.head_lr, \"name\": \"head\"},\n        ],\n        weight_decay=cfg.weight_decay,\n    )\n\n\ndef reset_cuda_peak_stats():\n    if DEVICE.type != \"cuda\":\n        return\n    torch.cuda.synchronize()\n    for device_id in range(torch.cuda.device_count()):\n        torch.cuda.reset_peak_memory_stats(device_id)\n\n\ndef cuda_peak_stats() -> dict:\n    if DEVICE.type != \"cuda\":\n        return {\"devices\": [], \"max_allocated_fraction\": 0.0, \"max_reserved_fraction\": 0.0}\n    torch.cuda.synchronize()\n    devices = []\n    for device_id in range(torch.cuda.device_count()):\n        total = int(torch.cuda.get_device_properties(device_id).total_memory)\n        allocated = int(torch.cuda.max_memory_allocated(device_id))\n        reserved = int(torch.cuda.max_memory_reserved(device_id))\n        devices.append({\n            \"device\": device_id,\n            \"name\": torch.cuda.get_device_name(device_id),\n            \"total_gib\": round(total / 2**30, 3),\n            \"peak_allocated_gib\": round(allocated / 2**30, 3),\n            \"peak_reserved_gib\": round(reserved / 2**30, 3),\n            \"peak_allocated_fraction\": allocated / total,\n            \"peak_reserved_fraction\": reserved / total,\n        })\n    return {\n        \"devices\": devices,\n        \"max_allocated_fraction\": max(x[\"peak_allocated_fraction\"] for x in devices),\n        \"max_reserved_fraction\": max(x[\"peak_reserved_fraction\"] for x in devices),\n    }\n\n\ndef run_worst_case_batch_probe(study_batch: int) -> dict:\n    # Fresh temporary modules make every candidate independent, including candidates that OOM\n    # while AdamW is allocating its moment buffers. All slots are marked valid: this is the\n    # largest activation graph the real loader can produce for this study batch.\n    encoder, head, runner = build_modules(load_pretrained=True)\n    optimizer = make_optimizer(encoder, head)\n    scaler = make_scaler()\n    tokens = sum(cfg.selected_tokens_per_slot)\n    slot_pattern = torch.tensor(\n        [slot_id for slot_id, count in enumerate(cfg.selected_tokens_per_slot) for _ in range(count)],\n        dtype=torch.long, device=DEVICE,\n    )\n    reset_cuda_peak_stats()\n    images = torch.zeros(\n        (study_batch, tokens, 3, cfg.cache_image_size, cfg.cache_image_size),\n        dtype=torch.uint8, device=DEVICE,\n    )\n    mask = torch.ones((study_batch, tokens), dtype=torch.bool, device=DEVICE)\n    slot = slot_pattern.unsqueeze(0).expand(study_batch, -1)\n    labels = torch.zeros((study_batch, len(TARGETS)), dtype=torch.float32, device=DEVICE)\n    optimizer.zero_grad(set_to_none=True)\n    with autocast_context():\n        features = encode_studies(runner, images, mask)\n        logits = head(features, mask, slot)\n        loss = F.binary_cross_entropy_with_logits(logits, labels)\n    scaler.scale(loss).backward()\n    scaler.unscale_(optimizer)\n    nn.utils.clip_grad_norm_(encoder.parameters(), 5.0)\n    nn.utils.clip_grad_norm_(head.parameters(), 5.0)\n    scaler.step(optimizer)  # forces allocation of AdamW first/second moments on GPU0\n    scaler.update()\n    stats = cuda_peak_stats()\n    stats[\"study_batch\"] = int(study_batch)\n    stats[\"triplet_images_total\"] = int(study_batch * tokens)\n    optimizer.zero_grad(set_to_none=True)\n    del loss, logits, features, labels, slot, mask, images\n    del scaler, optimizer, runner, head, encoder\n    gc.collect()\n    torch.cuda.empty_cache()\n    return stats\n\n\ndef configure_study_batch() -> dict:\n    if cfg.study_batch_size > 0:\n        result = {\n            \"mode\": \"manual\",\n            \"selected_study_batch\": int(cfg.study_batch_size),\n            \"grad_accum\": int(cfg.grad_accum),\n            \"effective_study_batch\": int(cfg.study_batch_size * cfg.grad_accum),\n            \"attempts\": [],\n        }\n        print(\"BATCH_TUNING\", json.dumps(result, indent=2))\n        return result\n    if DEVICE.type != \"cuda\":\n        cfg.study_batch_size = 2\n        result = {\n            \"mode\": \"cpu_fallback\", \"selected_study_batch\": 2,\n            \"grad_accum\": int(cfg.grad_accum), \"effective_study_batch\": 2 * int(cfg.grad_accum),\n            \"attempts\": [],\n        }\n        print(\"BATCH_TUNING\", json.dumps(result, indent=2))\n        return result\n\n    attempts = []\n    selected = None\n    for candidate in cfg.batch_candidates:\n        started = time.time()\n        try:\n            stats = run_worst_case_batch_probe(int(candidate))\n            stats[\"seconds\"] = round(time.time() - started, 3)\n            stats[\"status\"] = \"selected\"\n            attempts.append(stats)\n            selected = int(candidate)\n            print(\n                f\"[batch probe] study_batch={candidate}: SELECTED at \"\n                f\"{100 * stats['max_reserved_fraction']:.1f}% peak reserved VRAM\"\n            )\n            break\n        except (torch.cuda.OutOfMemoryError, RuntimeError) as exc:\n            message = str(exc)\n            if not isinstance(exc, torch.cuda.OutOfMemoryError) and \"out of memory\" not in message.lower():\n                raise\n            attempts.append({\n                \"study_batch\": int(candidate),\n                \"triplet_images_total\": int(candidate * sum(cfg.selected_tokens_per_slot)),\n                \"status\": \"oom\",\n                \"seconds\": round(time.time() - started, 3),\n                \"error\": message[:500],\n            })\n            print(f\"[batch probe] study_batch={candidate}: OOM; trying the next lower candidate\")\n        finally:\n            gc.collect()\n            torch.cuda.empty_cache()\n\n    if selected is None:\n        raise RuntimeError(f\"No safe study batch found from {cfg.batch_candidates}\")\n    cfg.study_batch_size = selected\n    result = {\n        \"mode\": \"auto_worst_case_optimizer_step\",\n        \"candidates\": list(cfg.batch_candidates),\n        \"selection_rule\": \"first_candidate_without_cuda_oom\",\n        \"selected_study_batch\": selected,\n        \"triplets_per_study\": sum(cfg.selected_tokens_per_slot),\n        \"selected_triplet_images_total\": selected * sum(cfg.selected_tokens_per_slot),\n        \"grad_accum\": int(cfg.grad_accum),\n        \"effective_study_batch\": int(selected * cfg.grad_accum),\n        \"attempts\": attempts,\n    }\n    print(\"BATCH_TUNING\", json.dumps(result, indent=2))\n    return result\n\n\ndef set_cosine_lr(optimizer, epoch_index: int, total_epochs: int):\n    factor = 0.05 + 0.95 * 0.5 * (\n        1.0 + math.cos(math.pi * epoch_index / max(total_epochs - 1, 1))\n    )\n    for group in optimizer.param_groups:\n        base = cfg.encoder_lr if group[\"name\"] == \"encoder\" else cfg.head_lr\n        group[\"lr\"] = base * factor\n\n\ndef compute_pos_weight(indices: np.ndarray) -> torch.Tensor:\n    positive = (W[indices] * Y[indices]).sum(0)\n    negative = (W[indices] * (1.0 - Y[indices])).sum(0)\n    value = np.clip(negative / np.maximum(positive, 1e-6), 1.0, cfg.max_pos_weight)\n    return torch.tensor(value, dtype=torch.float32, device=DEVICE)\n\n\ndef train_one_epoch(\n    encoder: nn.Module,\n    head: nn.Module,\n    runner,\n    optimizer,\n    scaler,\n    indices: np.ndarray,\n    pos_weight: torch.Tensor,\n    epoch: int,\n    total_epochs: int,\n) -> Tuple[float, dict]:\n    encoder.train()\n    head.train()\n    set_cosine_lr(optimizer, epoch - 1, total_epochs)\n    batch_used = int(cfg.study_batch_size)\n    accum_used = int(cfg.grad_accum)\n    reset_cuda_peak_stats()\n    loader = make_loader(indices, training=True, epoch=epoch)\n    optimizer.zero_grad(set_to_none=True)\n    losses = []\n    for step, batch in enumerate(loader):\n        row_indices = batch[\"index\"].numpy().astype(np.int64)\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        labels = torch.from_numpy(Y[row_indices]).to(DEVICE)\n        weights = torch.from_numpy(W[row_indices]).to(DEVICE)\n        with autocast_context():\n            features = encode_studies(runner, images, mask)\n            logits = head(features, mask, slot)\n            cell_loss = F.binary_cross_entropy_with_logits(\n                logits, labels, reduction=\"none\", pos_weight=pos_weight\n            )\n            per_study = (cell_loss * weights).sum(1) / weights.sum(1).clamp_min(1.0)\n            loss = per_study.mean()\n        group_start = (step // accum_used) * accum_used\n        group_size = min(accum_used, len(loader) - group_start)\n        scaler.scale(loss / group_size).backward()\n        if (step + 1) % accum_used == 0 or step + 1 == len(loader):\n            scaler.unscale_(optimizer)\n            nn.utils.clip_grad_norm_(encoder.parameters(), 5.0)\n            nn.utils.clip_grad_norm_(head.parameters(), 5.0)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad(set_to_none=True)\n        losses.append(float(loss.detach()))\n        if (step + 1) % 100 == 0:\n            print(f\"epoch={epoch} step={step + 1}/{len(loader)} loss={np.mean(losses[-100:]):.4f}\")\n    memory = cuda_peak_stats()\n    memory.update({\n        \"study_batch_size\": batch_used,\n        \"grad_accum\": accum_used,\n        \"effective_study_batch\": batch_used * accum_used,\n        \"triplets_per_study\": sum(cfg.selected_tokens_per_slot),\n    })\n    return float(np.mean(losses)), memory\n\n\n@torch.inference_mode()\ndef predict_indices(encoder: nn.Module, head: nn.Module, runner, indices: np.ndarray):\n    encoder.eval()\n    head.eval()\n    loader = make_loader(indices, training=False)\n    predictions, row_order, valid_slot_counts = [], [], []\n    for batch in 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(runner, 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_slot_counts.append(np.stack([\n            ((batch[\"mask\"] & batch[\"slot\"].eq(s)).sum(1)).numpy()\n            for s in range(len(SLOT_SPECS))\n        ], axis=1))\n    return (\n        np.asarray(row_order, np.int64),\n        np.concatenate(predictions, axis=0),\n        np.concatenate(valid_slot_counts, axis=0),\n    )\n\n\ndef cpu_state(module: nn.Module, half: bool = False) -> dict:\n    output = {}\n    for key, value in module.state_dict().items():\n        tensor = value.detach().cpu().clone()\n        if half and tensor.is_floating_point():\n            tensor = tensor.half()\n        output[key] = tensor\n    return output\n\n\ndef atomic_torch_save(obj: dict, path: Path):\n    temp = path.with_suffix(path.suffix + \".tmp\")\n    torch.save(obj, temp)\n    os.replace(temp, path)\n\n\ndef atomic_json(obj: dict, path: Path):\n    temp = path.with_suffix(path.suffix + \".tmp\")\n    temp.write_text(json.dumps(obj, indent=2), encoding=\"utf-8\")\n    os.replace(temp, path)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"a7cdd558-54d0-4973-87e7-4ea7c8bdf475","cell_type":"markdown","source":"## 6. Approximate resume: best checkpoint → remaining pseudo epochs → clean-gold selection\n\nThe attached best checkpoint supplies encoder/head weights and its saved epoch. A fresh AdamW and GradScaler continue only the remaining epoch numbers; the 4,349 pseudo studies remain the sole gradient source, while all 58 gold studies remain validation-only. Before continuing, the notebook recomputes the checkpoint's gold AUC so a run with no later improvement can still export the original best model.\n","metadata":{}},{"id":"31bf4be9-2b3c-44ce-8adb-6f7c4de86769","cell_type":"code","source":"def resolve_resume_checkpoint() -> Optional[Path]:\n    if cfg.resume_checkpoint:\n        configured = Path(cfg.resume_checkpoint)\n        if configured.is_file():\n            matches = [configured]\n        elif configured.is_dir():\n            matches = list(configured.rglob(\"simple_resnet34_cv_best.pt\"))\n        else:\n            raise FileNotFoundError(f\"RSNA_RESUME_CHECKPOINT does not exist: {configured}\")\n    else:\n        input_root = Path(\"/kaggle/input\")\n        matches = list(input_root.rglob(\"simple_resnet34_cv_best.pt\")) if input_root.is_dir() else []\n    matches = sorted({path.resolve() for path in matches}, key=str)\n    if not matches:\n        if cfg.require_resume:\n            raise FileNotFoundError(\n                \"Attach the failed run output containing simple_resnet34_cv_best.pt, \"\n                \"or set RSNA_RESUME_CHECKPOINT to that file/folder.\"\n            )\n        return None\n    if len(matches) != 1:\n        raise RuntimeError(\n            f\"Found {len(matches)} resume checkpoints; set RSNA_RESUME_CHECKPOINT explicitly: {matches}\"\n        )\n    return matches[0]\n\n\ndef read_resume_checkpoint(path: Path) -> dict:\n    try:\n        checkpoint = torch.load(path, map_location=\"cpu\", weights_only=False)\n    except TypeError:\n        checkpoint = torch.load(path, map_location=\"cpu\")\n    required = {\"encoder_state\", \"head_state\", \"epoch\", \"gold_macro_auc\"}\n    missing = sorted(required.difference(checkpoint))\n    if missing:\n        raise RuntimeError(f\"Resume checkpoint is missing keys: {missing}\")\n    saved_epoch = int(checkpoint[\"epoch\"])\n    if saved_epoch < 1 or saved_epoch > cfg.cv_epochs:\n        raise RuntimeError(\n            f\"Resume epoch {saved_epoch} is outside the configured 1..{cfg.cv_epochs} range\"\n        )\n    return checkpoint\n\n\nresume_path = resolve_resume_checkpoint()\nresume_checkpoint = read_resume_checkpoint(resume_path) if resume_path is not None else None\nbatch_tuning = configure_study_batch()\natomic_json(batch_tuning, OUT / \"simple_batch_tuning.json\")\ncv_encoder, cv_head, cv_runner = build_modules(load_pretrained=resume_checkpoint is None)\nif resume_checkpoint is not None:\n    cv_encoder.load_state_dict(resume_checkpoint[\"encoder_state\"], strict=True)\n    cv_head.load_state_dict(resume_checkpoint[\"head_state\"], strict=True)\n    cv_encoder.pretrained_source = f\"resume:{resume_path.name}\"\nprint(\"model source:\", cv_encoder.pretrained_source)\n\n# Approximate resume: model/BatchNorm state is restored; optimizer, scaler, and RNG are fresh.\ncv_optimizer = make_optimizer(cv_encoder, cv_head)\ncv_scaler = make_scaler()\ncv_pos_weight = compute_pos_weight(cv_train_indices)\ncv_started = time.time()\nhistory = []\n\nif resume_checkpoint is not None:\n    saved_epoch = int(resume_checkpoint[\"epoch\"])\n    rows, prediction, slot_counts = predict_indices(cv_encoder, cv_head, cv_runner, val_indices)\n    recomputed_auc, recomputed_by_target = weighted_macro_auc(\n        GOLD_Y[rows].astype(np.int64), prediction, None\n    )\n    recorded_auc = float(resume_checkpoint[\"gold_macro_auc\"])\n    if abs(recomputed_auc - recorded_auc) > 1e-8:\n        print(\n            f\"WARNING: checkpoint gold AUC recomputed as {recomputed_auc:.10f}, \"\n            f\"recorded value was {recorded_auc:.10f}; using the recomputed value.\"\n        )\n    best_auc = float(recomputed_auc)\n    best_epoch = saved_epoch\n    best_rows = rows.copy()\n    best_prediction = prediction.copy()\n    best_by_target = dict(recomputed_by_target)\n    best_state = {\n        \"encoder_state\": cpu_state(cv_encoder),\n        \"head_state\": cpu_state(cv_head),\n    }\n    start_epoch = saved_epoch + 1\n    resume_info = {\n        \"enabled\": True,\n        \"source_file\": resume_path.name,\n        \"source_epoch\": saved_epoch,\n        \"recorded_gold_macro_auc\": recorded_auc,\n        \"recomputed_gold_macro_auc\": best_auc,\n        \"optimizer_restored\": False,\n        \"scaler_restored\": False,\n        \"rng_restored\": False,\n        \"continued_from_epoch\": start_epoch,\n    }\n    history.append({\n        \"epoch\": saved_epoch,\n        \"event\": \"resume_checkpoint_baseline\",\n        \"loss\": None,\n        \"val_gold_macro_auc\": best_auc,\n        \"val_auc_by_target\": best_by_target,\n        \"valid_tokens_per_slot_mean\": slot_counts.mean(0).round(4).tolist(),\n    })\n    atomic_torch_save(\n        {\n            \"version\": \"rsna-knee-w01-d0-resnet34-gold-best-v2\",\n            \"epoch\": best_epoch,\n            \"gold_macro_auc\": best_auc,\n            \"gold_auc_by_target\": best_by_target,\n            \"resume_info\": resume_info,\n            **best_state,\n        },\n        OUT / \"simple_resnet34_cv_best.pt\",\n    )\n    print(\"RESUME_INFO\", json.dumps(resume_info, indent=2))\nelse:\n    best_auc = -float(\"inf\")\n    best_epoch = 0\n    best_state = None\n    best_rows = None\n    best_prediction = None\n    best_by_target = None\n    start_epoch = 1\n    resume_info = {\"enabled\": False}\n\nfor epoch in range(start_epoch, cfg.cv_epochs + 1):\n    epoch_started = time.time()\n    loss, train_memory = train_one_epoch(\n        cv_encoder, cv_head, cv_runner, cv_optimizer, cv_scaler,\n        cv_train_indices, cv_pos_weight, epoch, cfg.cv_epochs,\n    )\n    rows, prediction, slot_counts = predict_indices(cv_encoder, cv_head, cv_runner, val_indices)\n    macro_auc, by_target = weighted_macro_auc(GOLD_Y[rows].astype(np.int64), prediction, None)\n    record = {\n        \"epoch\": epoch,\n        \"loss\": loss,\n        \"train_memory\": train_memory,\n        \"val_gold_macro_auc\": macro_auc,\n        \"val_auc_by_target\": by_target,\n        \"valid_tokens_per_slot_mean\": slot_counts.mean(0).round(4).tolist(),\n        \"epoch_minutes\": (time.time() - epoch_started) / 60,\n    }\n    history.append(record)\n    print(json.dumps(record, indent=2))\n    if macro_auc > best_auc:\n        best_auc = macro_auc\n        best_epoch = epoch\n        best_rows = rows.copy()\n        best_prediction = prediction.copy()\n        best_by_target = dict(by_target)\n        best_state = {\n            \"encoder_state\": cpu_state(cv_encoder),\n            \"head_state\": cpu_state(cv_head),\n        }\n        atomic_torch_save(\n            {\n                \"version\": \"rsna-knee-w01-d0-resnet34-gold-best-v2\",\n                \"epoch\": best_epoch,\n                \"gold_macro_auc\": best_auc,\n                \"gold_auc_by_target\": best_by_target,\n                \"resume_info\": resume_info,\n                **best_state,\n            },\n            OUT / \"simple_resnet34_cv_best.pt\",\n        )\n\nif best_rows is None or best_prediction is None:\n    raise RuntimeError(\"Gold validation did not produce a valid best prediction\")\ncv_prediction_table = pd.DataFrame({UID: np.asarray(ordered_uids)[best_rows]})\nfor j, target in enumerate(TARGETS):\n    cv_prediction_table[f\"truth__{target}\"] = GOLD_Y[best_rows, j].astype(np.int64)\n    cv_prediction_table[f\"pred__{target}\"] = best_prediction[:, j]\ncv_prediction_table.to_csv(OUT / \"simple_gold_predictions_best_epoch.csv\", index=False)\natomic_json(history, OUT / \"simple_gold_history.json\")\ngold_auc_delta = (\n    None if cfg.reference_gold_auc is None else float(best_auc - cfg.reference_gold_auc)\n)\nworth_submit = (\n    None if gold_auc_delta is None else bool(gold_auc_delta >= cfg.decision_min_delta)\n)\ngold_decision = {\n    \"best_epoch\": best_epoch,\n    \"best_gold_macro_auc\": best_auc,\n    \"reference_gold_auc\": cfg.reference_gold_auc,\n    \"gold_auc_delta\": gold_auc_delta,\n    \"decision_min_delta\": cfg.decision_min_delta,\n    \"worth_one_lb_submission\": worth_submit,\n    \"elapsed_minutes\": (time.time() - cv_started) / 60,\n    \"resume_info\": resume_info,\n}\natomic_json(gold_decision, OUT / \"simple_gold_decision.json\")\nprint(\"GOLD_DECISION\", json.dumps(gold_decision, indent=2))\n","metadata":{},"outputs":[],"execution_count":null},{"id":"fe0cb06c-9954-4d85-a421-7abbf5448a22","cell_type":"markdown","source":"## 7. Export the best pseudo-trained checkpoint\n\nThe saved bundle is exactly the highest clean-gold checkpoint across the attached pre-crash best model and the remaining resumed epochs. Its filename and model schema are unchanged, so the existing inference notebook loads it without modification. There is no refit and no gold oversampling.\n","metadata":{}},{"id":"4aa3f434-4b65-4b45-9863-8d2235f78de8","cell_type":"code","source":"if best_state is None or best_epoch < 1:\n    raise RuntimeError(\"CV did not produce a valid checkpoint\")\n\nexported_encoder_state = {\n    key: value.half() if value.is_floating_point() else value\n    for key, value in best_state[\"encoder_state\"].items()\n}\nexported_head_state = {\n    key: value.half() if value.is_floating_point() else value\n    for key, value in best_state[\"head_state\"].items()\n}\nexported_training_scope = \"pseudo_only_best_gold_selected_checkpoint\"\n\nbundle = {\n    \"version\": \"rsna-knee-w01-d0-resnet34-v1\",\n    \"model_schema\": \"resnet34-2p5d-sixslot-meanmax-mlp-v1\",\n    \"targets\": TARGETS,\n    \"uid\": UID,\n    \"slot_specs\": SLOT_SPECS,\n    \"preprocess_signature\": PREPROCESS_SIGNATURE,\n    \"preprocess_payload\": PREPROCESS_PAYLOAD,\n    \"cache_image_size\": cfg.cache_image_size,\n    \"model_image_size\": cfg.model_image_size,\n    \"dense_tokens_per_slot\": cfg.dense_tokens_per_slot,\n    \"selected_tokens_per_slot\": cfg.selected_tokens_per_slot,\n    \"normalization\": {\n        \"mean\": [0.485, 0.456, 0.406],\n        \"std\": [0.229, 0.224, 0.225],\n    },\n    \"hidden_dim\": cfg.hidden_dim,\n    \"dropout\": cfg.dropout,\n    \"encoder_state\": exported_encoder_state,\n    \"head_state\": exported_head_state,\n    \"training\": {\n        \"seed\": cfg.seed,\n        \"single_model\": True,\n        \"diagnostic_training_source\": \"all 4349 W01 pseudo-labeled studies; zero gold studies\",\n        \"validation_source\": \"all 58 official gold studies from train.csv; unweighted\",\n        \"diagnostic_train_studies\": int(len(cv_train_indices)),\n        \"gold_validation_studies\": int(len(val_indices)),\n        \"best_epoch\": best_epoch,\n        \"best_gold_macro_auc\": best_auc,\n        \"reference_gold_auc\": cfg.reference_gold_auc,\n        \"gold_auc_delta\": gold_auc_delta,\n        \"worth_one_lb_submission\": worth_submit,\n        \"decision_min_delta\": cfg.decision_min_delta,\n        \"gold_used_for_gradient_updates\": False,\n        \"refit_performed\": False,\n        \"resume\": resume_info,\n        \"batch_tuning\": batch_tuning,\n        \"final_study_batch_size\": cfg.study_batch_size,\n        \"final_grad_accum\": cfg.grad_accum,\n        \"scope\": exported_training_scope,\n        \"cv_history\": history,\n    },\n}\n\nbundle_path = OUT / \"simple_resnet34_bundle.pt\"\natomic_torch_save(bundle, bundle_path)\nmanifest = {\n    \"version\": bundle[\"version\"],\n    \"bundle\": bundle_path.name,\n    \"bundle_sha256\": sha256_file(bundle_path),\n    \"training_scope\": exported_training_scope,\n    \"best_epoch\": best_epoch,\n    \"best_gold_macro_auc\": best_auc,\n    \"reference_gold_auc\": cfg.reference_gold_auc,\n    \"gold_auc_delta\": gold_auc_delta,\n    \"worth_one_lb_submission\": worth_submit,\n    \"preprocess_signature\": PREPROCESS_SIGNATURE,\n    \"selected_tokens_per_slot\": list(cfg.selected_tokens_per_slot),\n    \"selected_study_batch\": batch_tuning[\"selected_study_batch\"],\n    \"final_study_batch_size\": cfg.study_batch_size,\n    \"final_grad_accum\": cfg.grad_accum,\n    \"single_model\": True,\n    \"resume\": resume_info,\n}\natomic_json(manifest, OUT / \"simple_resnet34_manifest.json\")\nprint(json.dumps(manifest, indent=2))\nprint(\"SIMPLE_RESNET34_TRAIN_COMPLETE\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"b5fcd76c-95ef-446a-9236-37f7689cf030","cell_type":"markdown","source":"## Kaggle run contract\n\n1. Attach competition data, the Stage 0 output containing `w01_supervision.csv`, every completed Stage 0.5 cache part, and the failed run output containing `simple_resnet34_cv_best.pt`. ImageNet weights are not needed when resume succeeds.\n2. Run all cells with 2×T4. The notebook requires exactly one resume checkpoint by default; if several are attached, set `RSNA_RESUME_CHECKPOINT` to the intended file or containing folder.\n3. The default study batch is 64 and no batch-128 probe is performed. Override `RSNA_STUDY_BATCH_SIZE` only deliberately.\n4. The checkpoint epoch is validated on all 58 gold studies, then training continues from `saved_epoch + 1` through epoch 10 with a fresh optimizer/scaler.\n5. Save the completed output as a Kaggle Dataset and attach it to the existing inference notebook. It still resolves exactly one `simple_resnet34_bundle.pt`; no inference edit is required.\n6. Keep `RSNA_SIMPLE_TOKENS_PER_SLOT=2` so this remains the same D0 experiment.\n","metadata":{}}]}