{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11"},"w01":{"stage":"RSNA_Knee_W01_01_DINOv2S_FullData_20Epoch_AUC_Checkpoints_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 01: cached DINOv2-S full-data 20-epoch adaptation\n\nTrains from scratch on the complete Stage 0.5 cache, with hard gold labels x8. Stage 0.5 preserves the original 64-candidate preprocessing. At training time this notebook samples two candidates per slot **before reading image chunks**, batches four studies (48 selected images) for DINOv2-S, and keeps the original effective batch of eight studies via gradient accumulation 2.\n\nEpoch checkpoints contain optimizer, scaler and RNG state. On a later Kaggle run, attach the previous Stage 1 output and all Stage 0.5 parts; compatible `w01_last.pt` state is discovered automatically.\n","metadata":{}},{"cell_type":"code","source":"from __future__ import annotations\n\nimport gc, hashlib, json, math, os, random, re, time\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\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\"\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    supervision_dir: Optional[str] = os.getenv(\"RSNA_SUPERVISION_DIR\")\n    cache_dirs: Optional[str] = os.getenv(\"RSNA_CACHE_DIRS\")\n    resume_checkpoint: Optional[str] = os.getenv(\"RSNA_RESUME_CHECKPOINT\")\n    auto_resume: bool = os.getenv(\"RSNA_AUTO_RESUME\", \"1\") == \"1\"\n    output_dir: str = os.getenv(\"RSNA_ADAPT_OUTPUT\", \"/kaggle/working/w01_adapt\")\n    dinov2_model: str = \"vit_small_patch14_dinov2.lvd142m\"\n    dinov2_weights: str = os.getenv(\"RSNA_DINOV2_WEIGHTS\", \"\")\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    tune_tokens_per_slot: Tuple[int, ...] = (2, 2, 2, 2, 2, 2)\n    study_batch_size: int = int(os.getenv(\"RSNA_STUDY_BATCH_SIZE\", \"8\"))\n    tune_encode_batch: int = int(os.getenv(\"RSNA_TUNE_ENCODE_BATCH\", \"96\"))\n    cache_workers: int = int(os.getenv(\"RSNA_CACHE_WORKERS\", \"4\"))\n    hidden_dim: int = 256\n    head_dropout: float = 0.22\n    series_layers: int = 2\n    series_heads: int = 8\n    series_dropout: float = 0.12\n    finetune_epochs: int = 30\n    checkpoint_epochs: Tuple[int, ...] = (10, 20, 30)\n    grad_accum: int = int(os.getenv(\"RSNA_GRAD_ACCUM\", \"2\"))\n    head_lr: float = 2e-4\n    projection_lr: float = 1e-4\n    backbone_lr: float = 1e-5\n    weight_decay: float = 1e-2\n    unfreeze_blocks: int = 6\n    max_pos_weight: float = 6.0\n    gold_oversample: int = 8\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\", \"15\"))\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_ADAPT_OUTPUT\", str(Path.cwd() / \"smoke_w01_adapt\"))\n    cfg.image_size = 64; cfg.spatial_grid = 7; cfg.projection_dim = 32\n    cfg.tokens_per_slot = (2, 2, 2, 2, 2, 2)\n    cfg.tune_tokens_per_slot = (1, 1, 1, 1, 1, 1)\n    cfg.study_batch_size = 2; cfg.tune_encode_batch = 12; cfg.cache_workers = 0\n    cfg.hidden_dim = 48; cfg.series_layers = 1; cfg.series_heads = 4\n    cfg.finetune_epochs = 1; cfg.checkpoint_epochs = (1,)\n    cfg.grad_accum = 1; cfg.gold_oversample = 2\n    cfg.require_two_gpus = False; cfg.time_budget_minutes = 30; cfg.stop_margin_minutes = 1\nOUT = Path(cfg.output_dir); OUT.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\ndef seed_everything(seed: int):\n    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n    if torch.cuda.is_available(): torch.backends.cudnn.benchmark = True\n\nseed_everything(cfg.seed)\nprint(json.dumps(asdict(cfg), indent=2)); print(\"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 01\")\nif cfg.study_batch_size * sum(cfg.tune_tokens_per_slot) < cfg.tune_encode_batch:\n    print(\"note: tune_encode_batch exceeds available selected images; the actual call will be smaller\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Competition UID order\n","metadata":{}},{"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) == len(cfg.tune_tokens_per_slot)\n\ndef find_comp_root() -> Path:\n    if cfg.comp_root:\n        p = Path(cfg.comp_root)\n        if (p / \"train.csv\").is_file(): return p\n        raise FileNotFoundError(f\"Invalid RSNA_COMP_ROOT: {p}\")\n    if cfg.smoke:\n        p = Path.cwd() / \"smoke_rsna_data\"\n        if (p / \"train.csv\").is_file(): return p\n        raise FileNotFoundError(\"Run Stage 0.5 smoke before Stage 1 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 p in candidates:\n        if (p / \"train.csv\").is_file(): return p\n    base = Path(\"/kaggle/input\")\n    if base.is_dir():\n        for p in base.iterdir():\n            if p.is_dir() and (p / \"train.csv\").is_file(): return p\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})\nfor target in TARGETS:\n    train_df[target] = pd.to_numeric(train_df.get(target), errors=\"coerce\")\nif train_df[UID].duplicated().any():\n    raise RuntimeError(\"train.csv contains duplicate StudyInstanceUID values\")\nprint(\"root:\", ROOT, \"train studies:\", len(train_df))\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load immutable Stage 00 supervision\n","metadata":{}},{"cell_type":"code","source":"def write_smoke_supervision() -> Path:\n    folder = Path.cwd() / \"smoke_w01_supervision\"\n    folder.mkdir(parents=True, exist_ok=True)\n    path = folder / \"w01_supervision.csv\"\n    out = pd.DataFrame({UID: train_df[UID].astype(str)})\n    gold = train_df[TARGETS].notna().all(axis=1).to_numpy()\n    row_id = np.arange(len(train_df))\n    for j, target in enumerate(TARGETS):\n        hard = train_df[target].to_numpy(np.float32)\n        pseudo = (0.15 + 0.70 * (((row_id + j) % 3) == 0)).astype(np.float32)\n        y = np.where(gold, hard, pseudo)\n        out[f\"y__{target}\"] = y\n        out[f\"w__{target}\"] = 1.0\n        out[f\"conf__{target}\"] = np.where(gold, 1.0, 0.8)\n        out[f\"gold__{target}\"] = np.where(gold, hard, np.nan)\n    out[\"is_gold\"] = gold\n    out[\"fold\"] = row_id % 2\n    out.to_csv(path, index=False)\n    return path\n\ndef find_unique(filename: str, configured_dir: Optional[str] = None) -> Path:\n    if configured_dir:\n        path = Path(configured_dir)\n        path = path if path.is_file() else path / filename\n        if path.is_file(): return path\n        raise FileNotFoundError(path)\n    if cfg.smoke:\n        return write_smoke_supervision()\n    roots = [Path.cwd(), Path(\"/kaggle/input\")]\n    found = []\n    for base in roots:\n        if not base.exists(): continue\n        found.extend(p for p in base.rglob(filename)\n                     if \"train_series\" not in p.parts and \"test_series\" not in p.parts)\n    found = sorted(set(found), key=str)\n    if len(found) != 1:\n        raise RuntimeError(f\"Expected exactly one {filename}; found {found}. Set RSNA_SUPERVISION_DIR.\")\n    return found[0]\n\nsupervision_path = find_unique(\"w01_supervision.csv\", cfg.supervision_dir)\nsupervision = pd.read_csv(supervision_path, dtype={UID: str})\nif supervision[UID].tolist() != train_df[UID].tolist():\n    raise RuntimeError(\"Stage00 supervision UID order differs from competition train.csv\")\nY = supervision[[f\"y__{t}\" for t in TARGETS]].to_numpy(np.float32)\nW = supervision[[f\"w__{t}\" for t in TARGETS]].to_numpy(np.float32)\nC = supervision[[f\"conf__{t}\" for t in TARGETS]].to_numpy(np.float32)\ngold_mask = supervision[\"is_gold\"].to_numpy(bool)\nfold_ids = supervision[\"fold\"].to_numpy(np.int64)\nprint(\"supervision:\", supervision_path, \"gold:\", int(gold_mask.sum()))\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Validate Stage 0.5 parts and read only selected token chunks\n","metadata":{}},{"cell_type":"code","source":"def 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\nMANIFEST_NAME = \"w01_stage05_manifest.json\"\nPREPROCESS_SIGNATURE, PREPROCESS_PAYLOAD = preprocessing_signature()\nordered_uids = train_df[UID].astype(str).tolist()\ntrain_uid_sha256 = hashlib.sha256(\"\\n\".join(ordered_uids).encode()).hexdigest()\n\ndef configured_roots(value: Optional[str]) -> List[Path]:\n    if not value: return []\n    return [Path(x.strip()) for x in re.split(r\"[;,]\", value) if x.strip()]\n\ndef discover_cache_manifests() -> List[Path]:\n    roots = configured_roots(cfg.cache_dirs)\n    if not roots:\n        roots = [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(p.resolve() for p in found))\n\ndef sha256_file(path: Path) -> str:\n    h = hashlib.sha256()\n    with path.open(\"rb\") as f:\n        for block in iter(lambda: f.read(8 << 20), b\"\"):\n            h.update(block)\n    return h.hexdigest()\n\ncache_index: Dict[str, Tuple[Path, str]] = {}\ncompatible_manifests = []\nignored_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 cache 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)} expected={len(ordered_uids)} \"\n        f\"missing={missing[:10]} extra={extra[:10]}. Attach every Stage 0.5 output part.\"\n    )\nif ignored_manifests:\n    print(\"ignored incompatible cache manifests:\", ignored_manifests)\nprint(f\"cache ready: {len(cache_index)} studies from {len(compatible_manifests)} part manifest(s)\")\n\nclass CacheAdaptPixels(Dataset):\n    def __init__(self, indices: np.ndarray, training: bool):\n        self.indices = np.asarray(indices, dtype=np.int64)\n        self.training = training\n        self._handles = {}\n        self.n_out = sum(cfg.tune_tokens_per_slot)\n\n    def __len__(self):\n        return len(self.indices)\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            # Shards are immutable and fully closed before they appear in a manifest.\n            # Plain read-only opens are safe across DataLoader worker processes and\n            # avoid imposing HDF5 SWMR-format requirements on the cache writer.\n            handle = h5py.File(path, \"r\")\n            if handle.attrs.get(\"preprocess_signature\", \"\") != PREPROCESS_SIGNATURE:\n                handle.close(); raise RuntimeError(f\"cache signature mismatch: {path}\")\n            self._handles[key] = handle\n        return handle\n\n    def __getitem__(self, i):\n        j = int(self.indices[i]); uid = ordered_uids[j]\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        dense_z = group[\"z\"][:].astype(np.float32)\n        selected, out_mask, out_slot, out_z = [], [], [], []\n        for slot_id, quota in enumerate(cfg.tune_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                    pos = np.linspace(0, len(candidates) - 1, count).round().astype(np.int64)\n                    chosen = candidates[pos]\n                selected.extend(chosen.tolist())\n                out_mask.extend([True] * count)\n                out_slot.extend([slot_id] * count)\n                out_z.extend(dense_z[chosen].tolist())\n            for _ in range(int(quota) - count):\n                selected.append(-1); out_mask.append(False); out_slot.append(slot_id); out_z.append(0.5)\n        images = np.zeros((self.n_out, 3, cfg.image_size, cfg.image_size), np.uint8)\n        read_positions = [(out_pos, dense_pos) for out_pos, dense_pos in enumerate(selected) if dense_pos >= 0]\n        if read_positions:\n            dense_positions = [dense_pos for _, dense_pos in read_positions]\n            values = group[\"images\"][dense_positions]\n            for (out_pos, _), value in zip(read_positions, values):\n                images[out_pos] = value\n        return {\n            \"index\": j, \"uid\": uid,\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            \"z\": torch.tensor(out_z, dtype=torch.float32),\n        }\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## DINOv2-S 336/14 encoder — final six transformer blocks trainable\n","metadata":{}},{"cell_type":"code","source":"class 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    model = timm.create_model(\n        cfg.dinov2_model, pretrained=pretrained, img_size=cfg.image_size,\n        num_classes=0, dynamic_img_size=False,\n    )\n    if cfg.dinov2_weights:\n        path = Path(cfg.dinov2_weights)\n        if not path.is_file():\n            raise FileNotFoundError(path)\n        obj = torch.load(path, map_location=\"cpu\", weights_only=False)\n        state = _unwrap_state(obj)\n        clean = {}\n        own = model.state_dict()\n        for key, value in state.items():\n            key = str(key)\n            for prefix in (\"module.\", \"model.\", \"backbone.\"):\n                if key.startswith(prefix):\n                    key = key[len(prefix):]\n            if key in own and own[key].shape == value.shape:\n                clean[key] = value\n        result = model.load_state_dict(clean, strict=False)\n        if len(clean) < int(.90 * len(own)):\n            raise RuntimeError(f\"DINOv2-S weights incompatible: {len(clean)}/{len(own)} tensors\")\n        print(\"loaded offline DINOv2-S:\", path, \"missing:\", len(result.missing_keys))\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=not bool(cfg.dinov2_weights)).to(DEVICE)\n    return {enc.name: enc}\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# DINOv2-S is loaded through timm; the complete tuned state is stored in the checkpoint.\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Label-specific spatial pooling + per-series Transformer (six slots)\n","metadata":{}},{"cell_type":"code","source":"class LabelMIL(nn.Module):\n    def __init__(self, region_dim: int, hidden: int, n_slots: int, dropout: float,\n                 series_layers: int, series_heads: int, series_dropout: float):\n        super().__init__()\n        if hidden % series_heads:\n            raise ValueError(f\"hidden={hidden} must be divisible by series_heads={series_heads}\")\n        self.region_dim = region_dim\n        self.hidden = hidden\n        self.n_slots = n_slots\n        self.proj = nn.Sequential(nn.LayerNorm(region_dim), nn.Linear(region_dim, hidden), nn.GELU())\n        self.slot_emb = nn.Embedding(n_slots, hidden)\n        self.z_mlp = nn.Sequential(nn.Linear(1, hidden), nn.Tanh(), nn.Linear(hidden, hidden))\n        self.spatial_a = nn.Linear(hidden, hidden); self.spatial_b = nn.Linear(hidden, hidden)\n        self.spatial_q = nn.Parameter(torch.randn(len(TARGETS), hidden) * 0.02)\n        layer = nn.TransformerEncoderLayer(\n            d_model=hidden, nhead=series_heads, dim_feedforward=hidden * 3,\n            dropout=series_dropout, activation=\"gelu\", batch_first=True, norm_first=True,\n        )\n        self.series_encoder = nn.TransformerEncoder(layer, num_layers=series_layers)\n        self.series_norm = nn.LayerNorm(hidden)\n        self.depth_a = nn.Linear(hidden, hidden); self.depth_b = nn.Linear(hidden, hidden)\n        self.depth_q = nn.Parameter(torch.randn(len(TARGETS), hidden) * 0.02)\n        # slot order: sag fluid/struct, cor fluid/struct, axial fluid/struct\n        prior = torch.tensor([\n            [1.0, .8, .5, .4, .1, .2], [.4, .4, 1.0, .8, .1, .2],\n            [.8, 1.0, .7, .8, .1, .2], [.8, 1.0, .7, .8, .1, .2],\n            [.4, .5, .8, 1.0, .2, .3], [.4, .5, .8, 1.0, .2, .3],\n            [.2, .2, .2, .2, 1.0, .9], [.6, .2, .4, .2, 1.0, .8],\n            [.6, .2, .4, .2, 1.0, .7], [1.0, .3, .4, .2, .7, .5],\n            [.8, .3, .8, .3, .8, .7], [.6, .6, .6, .6, .6, .7],\n        ], dtype=torch.float32)\n        if n_slots != prior.shape[1]:\n            raise ValueError(f\"W01 requires six slots, got {n_slots}\")\n        prior = prior - prior.mean(1, keepdim=True)\n        self.slot_bias = nn.Parameter(0.35 * prior)\n        self.cls_weight = nn.Parameter(torch.randn(len(TARGETS), hidden) * 0.02)\n        self.cls_bias = nn.Parameter(torch.zeros(len(TARGETS)))\n        self.mean_residual = nn.Linear(hidden, len(TARGETS))\n        self.dropout = nn.Dropout(dropout)\n\n    def forward(self, feat, mask, slot, z):\n        # First perform target-specific spatial pooling inside every 2D slice.\n        base = self.slot_emb(slot) + self.z_mlp(z.unsqueeze(-1))\n        h = self.proj(feat) + base[:, :, None]\n        spatial_gate = torch.tanh(self.spatial_a(h)) * torch.sigmoid(self.spatial_b(h))\n        spatial_score = torch.einsum(\"btrh,lh->bltr\", spatial_gate, self.spatial_q) / math.sqrt(self.hidden)\n        spatial_att = spatial_score.softmax(-1)\n        token_ctx = torch.einsum(\"bltr,btrh->blth\", spatial_att, h)\n        token_ctx = token_ctx * mask[:, None, :, None]\n\n        # V4A: true slice-to-slice interaction. Each diagnostic series is encoded\n        # independently, preserving its ordered z positions and padding mask.\n        batch, labels, _, hidden = token_ctx.shape\n        sequence_ctx = torch.zeros_like(token_ctx)\n        for slot_id in range(self.n_slots):\n            positions = torch.nonzero(slot[0] == slot_id, as_tuple=False).flatten()\n            if not len(positions):\n                continue\n            value = token_ctx.index_select(2, positions)\n            n_depth = value.shape[2]\n            value = value.reshape(batch * labels, n_depth, hidden)\n            padding = (~mask.index_select(1, positions)).unsqueeze(1)\n            padding = padding.expand(batch, labels, n_depth).reshape(batch * labels, n_depth)\n            valid_rows = (~padding).any(1)\n            encoded = torch.zeros_like(value)\n            if valid_rows.any():\n                valid_index = torch.nonzero(valid_rows, as_tuple=False).flatten()\n                valid_encoded = self.series_encoder(\n                    value.index_select(0, valid_index),\n                    src_key_padding_mask=padding.index_select(0, valid_index),\n                )\n                encoded = encoded.index_copy(0, valid_index, valid_encoded)\n            # CUDA autocast may return FP32 from LayerNorm while token_ctx is FP16.\n            # index_copy requires an exact dtype match on GPU.\n            encoded = self.series_norm(encoded).to(sequence_ctx.dtype)\n            encoded = encoded.reshape(batch, labels, n_depth, hidden)\n            sequence_ctx = sequence_ctx.index_copy(2, positions, encoded)\n        sequence_ctx = sequence_ctx * mask[:, None, :, None]\n\n        # Target-specific pooling now consumes contextualized slice features.\n        depth_gate = torch.tanh(self.depth_a(sequence_ctx)) * torch.sigmoid(self.depth_b(sequence_ctx))\n        depth_score = torch.einsum(\"blth,lh->blt\", depth_gate, self.depth_q) / math.sqrt(self.hidden)\n        depth_score = depth_score + self.slot_bias[:, slot].permute(1, 0, 2)\n        depth_score = depth_score.masked_fill(~mask[:, None], -1e4)\n        depth_att = depth_score.softmax(-1)\n        ctx = torch.einsum(\"blt,blth->blh\", depth_att, sequence_ctx)\n        logits = (self.dropout(ctx) * self.cls_weight[None]).sum(-1) + self.cls_bias\n        denom = (mask.sum(1, keepdim=True) * feat.shape[2]).clamp_min(1)\n        mean = (h * mask[:, :, None, None]).sum((1, 2)) / denom\n        return logits + 0.25 * self.mean_residual(self.dropout(mean))\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Resumable 20-epoch adaptation\n","metadata":{}},{"cell_type":"code","source":"def set_trainable(enc: nn.Module):\n    for parameter in enc.parameters(): parameter.requires_grad = False\n    for module in (enc.projector, enc.project_norm):\n        for parameter in module.parameters(): parameter.requires_grad = True\n    if enc.name == \"tiny\": return\n    blocks = getattr(enc.model, \"blocks\", None)\n    if blocks is None or len(blocks) < cfg.unfreeze_blocks:\n        raise RuntimeError(\"DINOv2-S block list is incompatible\")\n    for block in blocks[-cfg.unfreeze_blocks:]:\n        for parameter in block.parameters(): parameter.requires_grad = True\n    if getattr(enc.model, \"norm\", None) is not None:\n        for parameter in enc.model.norm.parameters(): parameter.requires_grad = True\n\ndef load_pt(path: Path):\n    try: return torch.load(path, map_location=\"cpu\", weights_only=False)\n    except TypeError: return torch.load(path, map_location=\"cpu\")\n\ndef atomic_torch_save(obj: dict, path: Path):\n    tmp = path.with_suffix(path.suffix + \".tmp\")\n    torch.save(obj, tmp); os.replace(tmp, path)\n\ndef atomic_json(obj: dict, path: Path):\n    tmp = path.with_suffix(path.suffix + \".tmp\")\n    tmp.write_text(json.dumps(obj, indent=2), encoding=\"utf-8\"); os.replace(tmp, path)\n\ndef full_encoder_state(enc: nn.Module, half: bool = False):\n    state = {}\n    for name, value in enc.state_dict().items():\n        value = value.detach().cpu()\n        state[name] = value.half() if half and value.is_floating_point() else value.clone()\n    return state\n\ndef weighted_macro_auc(y_true: np.ndarray, y_prob: np.ndarray, weight: Optional[np.ndarray] = None):\n    by_target, values = {}, []\n    binary = (np.asarray(y_true) >= 0.5).astype(np.int8)\n    probability = np.asarray(y_prob, dtype=np.float64)\n    weight = np.ones_like(probability, dtype=np.float64) if weight is None else np.asarray(weight, dtype=np.float64)\n    for target_id, target in enumerate(TARGETS):\n        keep = np.isfinite(probability[:, target_id]) & np.isfinite(weight[:, target_id]) & (weight[:, target_id] > 0)\n        labels = binary[keep, target_id]\n        if len(labels) < 2 or np.unique(labels).size < 2:\n            by_target[target] = None; continue\n        score = float(roc_auc_score(labels, probability[keep, target_id], sample_weight=weight[keep, target_id]))\n        by_target[target] = score; values.append(score)\n    return (float(np.mean(values)) if values else float(\"nan\")), by_target\n\ndef encode_batch(runner, enc, images, mask_cpu):\n    batch, tokens = mask_cpu.shape\n    flat_mask = mask_cpu.reshape(-1)\n    valid = torch.nonzero(flat_mask, as_tuple=False).flatten()\n    full = torch.zeros((batch * tokens, int(enc.n_regions), int(enc.region_dim)), device=DEVICE)\n    flat_images = images.reshape(batch * tokens, *images.shape[2:])\n    for start in range(0, len(valid), cfg.tune_encode_batch):\n        pos = valid[start:start + cfg.tune_encode_batch]\n        x = flat_images.index_select(0, pos).to(DEVICE, non_blocking=True)\n        with torch.autocast(\"cuda\", enabled=DEVICE.type == \"cuda\"):\n            value = runner(x).float()\n        full = full.index_copy(0, pos.to(DEVICE), value)\n    return full.reshape(batch, tokens, int(enc.n_regions), int(enc.region_dim))\n\nenc = next(iter(build_encoders().values()))\nset_trainable(enc)\nhead = LabelMIL(int(enc.region_dim), cfg.hidden_dim, len(SLOT_SPECS), cfg.head_dropout,\n                cfg.series_layers, cfg.series_heads, cfg.series_dropout).to(DEVICE)\nrunner = nn.DataParallel(enc, device_ids=[0, 1]) if torch.cuda.device_count() >= 2 else enc\nprint(f\"encoder={enc.name} study_batch={cfg.study_batch_size} selected/study={sum(cfg.tune_tokens_per_slot)} \"\n      f\"DINO-call<={cfg.tune_encode_batch} DataParallel={isinstance(runner, nn.DataParallel)}\")\n\ndef seed_worker(worker_id: int):\n    worker_seed = torch.initial_seed() % (2 ** 32)\n    np.random.seed(worker_seed); random.seed(worker_seed)\n\n@torch.inference_mode()\ndef audit_gold_train_auc(indices: np.ndarray):\n    enc.eval(); head.eval()\n    loader = DataLoader(\n        CacheAdaptPixels(indices, training=False), batch_size=cfg.study_batch_size,\n        shuffle=False, num_workers=cfg.cache_workers,\n        pin_memory=DEVICE.type == \"cuda\", persistent_workers=False,\n        worker_init_fn=seed_worker,\n    )\n    predictions, row_indices = [], []\n    for batch in loader:\n        mask_cpu = batch[\"mask\"]\n        feature = encode_batch(runner, enc, batch[\"images\"], mask_cpu)\n        with torch.autocast(\"cuda\", enabled=DEVICE.type == \"cuda\"):\n            logits = head(feature, mask_cpu.to(DEVICE), batch[\"slot\"].to(DEVICE), batch[\"z\"].to(DEVICE))\n        predictions.append(torch.sigmoid(logits.float()).cpu().numpy())\n        row_indices.extend(batch[\"index\"].numpy().astype(int).tolist())\n    if not predictions: return float(\"nan\"), {}\n    rows = np.asarray(row_indices, dtype=np.int64)\n    return weighted_macro_auc(Y[rows], np.concatenate(predictions), None)\n\ngold_indices = np.flatnonzero(gold_mask)\npseudo_indices = np.flatnonzero(~gold_mask)\ntrain_indices = np.concatenate([pseudo_indices, np.repeat(gold_indices, cfg.gold_oversample)])\nif not cfg.smoke and (len(gold_indices) != 58 or len(pseudo_indices) != 4349):\n    raise RuntimeError(f\"expected 58 gold + 4349 pseudo; got {len(gold_indices)} + {len(pseudo_indices)}\")\nif not np.allclose(Y[gold_mask], supervision.loc[gold_mask, [f\"gold__{t}\" for t in TARGETS]].to_numpy(np.float32)):\n    raise RuntimeError(\"Stage00 did not write hard gold labels into y__ columns\")\n\ntrainable = [(name, parameter) for name, parameter in enc.named_parameters() if parameter.requires_grad]\nprojection_ids = {id(p) for p in list(enc.projector.parameters()) + list(enc.project_norm.parameters())}\nprojection_params = [p for _, p in trainable if id(p) in projection_ids]\nbackbone_params = [p for _, p in trainable if id(p) not in projection_ids]\noptimizer = torch.optim.AdamW([\n    {\"params\": head.parameters(), \"lr\": cfg.head_lr, \"name\": \"head\"},\n    {\"params\": projection_params, \"lr\": cfg.projection_lr, \"name\": \"projection\"},\n    {\"params\": backbone_params, \"lr\": cfg.backbone_lr, \"name\": \"backbone\"},\n], weight_decay=cfg.weight_decay)\nscaler = torch.amp.GradScaler(\"cuda\", enabled=DEVICE.type == \"cuda\")\n\nsignature_payload = {\n    \"version\": \"w01-dinov2s-cache-adapt-1\", \"preprocess_signature\": PREPROCESS_SIGNATURE,\n    \"epochs\": cfg.finetune_epochs, \"unfreeze_blocks\": cfg.unfreeze_blocks,\n    \"gold_oversample\": cfg.gold_oversample, \"checkpoint_epochs\": list(cfg.checkpoint_epochs),\n    \"slots\": SLOT_SPECS, \"tokens\": cfg.tokens_per_slot,\n    \"tune_tokens\": cfg.tune_tokens_per_slot, \"projection_dim\": cfg.projection_dim,\n    \"study_batch_size\": cfg.study_batch_size, \"grad_accum\": cfg.grad_accum,\n}\ntraining_signature = hashlib.sha256(json.dumps(signature_payload, sort_keys=True).encode()).hexdigest()\nlast_path = OUT / \"w01_last.pt\"\n\ndef discover_resume_checkpoint() -> Optional[Path]:\n    if cfg.resume_checkpoint:\n        path = Path(cfg.resume_checkpoint)\n        path = path / \"w01_last.pt\" if path.is_dir() else path\n        if not path.is_file(): raise FileNotFoundError(path)\n        return path\n    if last_path.is_file(): return last_path\n    if not cfg.auto_resume: return None\n    candidates = []\n    for root in (Path(\"/kaggle/input\"), Path.cwd()):\n        if not root.exists(): continue\n        for manifest_path in root.rglob(\"w01_manifest.json\"):\n            try:\n                manifest = json.loads(manifest_path.read_text(encoding=\"utf-8\"))\n                checkpoint_path = manifest_path.parent / \"w01_last.pt\"\n                if manifest.get(\"training_signature\") == training_signature and checkpoint_path.is_file():\n                    candidates.append((int(manifest.get(\"epoch_completed\", 0)), checkpoint_path))\n            except Exception:\n                pass\n    if not candidates: return None\n    candidates.sort(key=lambda x: (x[0], str(x[1])))\n    best_epoch = candidates[-1][0]\n    best = [path for epoch, path in candidates if epoch == best_epoch]\n    if len(best) > 1:\n        raise RuntimeError(f\"multiple equally advanced resume checkpoints: {best}. Set RSNA_RESUME_CHECKPOINT.\")\n    return best[0]\n\ndef move_optimizer_to_device(opt):\n    for state in opt.state.values():\n        for key, value in list(state.items()):\n            if torch.is_tensor(value): state[key] = value.to(DEVICE)\n\nstart_epoch, history = 0, []\nresume_path = discover_resume_checkpoint()\nif resume_path is not None:\n    checkpoint = load_pt(resume_path)\n    if checkpoint.get(\"training_signature\") != training_signature:\n        raise RuntimeError(\"refusing incompatible Stage01 resume checkpoint\")\n    enc.load_state_dict(checkpoint[\"encoder_state\"], strict=True)\n    head.load_state_dict(checkpoint[\"head_state\"], strict=True)\n    optimizer.load_state_dict(checkpoint[\"optimizer_state\"]); move_optimizer_to_device(optimizer)\n    if checkpoint.get(\"scaler_state\"): scaler.load_state_dict(checkpoint[\"scaler_state\"])\n    start_epoch = int(checkpoint[\"epoch_completed\"]); history = checkpoint.get(\"history\", [])\n    rng = checkpoint.get(\"rng_state\", {})\n    if rng:\n        random.setstate(rng[\"python\"]); np.random.set_state(rng[\"numpy\"]); torch.set_rng_state(rng[\"torch\"])\n        if torch.cuda.is_available() and rng.get(\"cuda\"):\n            try: torch.cuda.set_rng_state_all(rng[\"cuda\"])\n            except Exception as exc: print(\"warning: CUDA RNG state not restored:\", repr(exc))\n    print(\"resuming:\", resume_path, \"after epoch\", start_epoch)\nelse:\n    print(\"starting Stage01 from scratch\")\n\npos = (W[train_indices] * Y[train_indices]).sum(0)\nneg = (W[train_indices] * (1 - Y[train_indices])).sum(0)\npos_weight = torch.tensor(np.clip(neg / np.maximum(pos, 1e-6), 1, cfg.max_pos_weight), device=DEVICE)\nstarted = time.time()\n\nfor epoch in range(start_epoch, cfg.finetune_epochs):\n    completed_this_run = epoch - start_epoch\n    if completed_this_run > 0:\n        elapsed = (time.time() - started) / 60\n        mean_epoch = elapsed / completed_this_run\n        if elapsed + mean_epoch + cfg.stop_margin_minutes >= cfg.time_budget_minutes:\n            print(\"TIME_BUDGET_GRACEFUL_STOP_BEFORE_EPOCH\", epoch + 1)\n            break\n    epoch_started = time.time()\n    cosine = .5 * (1 + math.cos(math.pi * epoch / max(cfg.finetune_epochs - 1, 1)))\n    for group in optimizer.param_groups:\n        base_lr = {\"head\": cfg.head_lr, \"projection\": cfg.projection_lr, \"backbone\": cfg.backbone_lr}[group[\"name\"]]\n        group[\"lr\"] = base_lr * max(cosine, .10)\n    enc.train(); head.train(); optimizer.zero_grad(set_to_none=True)\n    epoch_prediction_sum = np.zeros_like(Y, dtype=np.float64)\n    epoch_prediction_count = np.zeros(len(Y), dtype=np.int32)\n    generator = torch.Generator().manual_seed(cfg.seed + epoch)\n    loader = DataLoader(\n        CacheAdaptPixels(train_indices, True), batch_size=cfg.study_batch_size, shuffle=True,\n        generator=generator, num_workers=cfg.cache_workers,\n        pin_memory=DEVICE.type == \"cuda\", persistent_workers=cfg.cache_workers > 0,\n        worker_init_fn=seed_worker, prefetch_factor=2 if cfg.cache_workers > 0 else None,\n    )\n    losses = []\n    for step, batch in enumerate(loader):\n        idxs = batch[\"index\"].numpy().astype(np.int64)\n        mask_cpu = batch[\"mask\"]\n        feat = encode_batch(runner, enc, batch[\"images\"], mask_cpu)\n        y = torch.from_numpy(Y[idxs]).to(DEVICE); w = torch.from_numpy(W[idxs]).to(DEVICE)\n        with torch.autocast(\"cuda\", enabled=DEVICE.type == \"cuda\"):\n            logits = head(feat, mask_cpu.to(DEVICE), batch[\"slot\"].to(DEVICE), batch[\"z\"].to(DEVICE))\n            cell = F.binary_cross_entropy_with_logits(logits, y, reduction=\"none\", pos_weight=pos_weight)\n            # Match the former batch_size=1 regime exactly: normalize each study by\n            # its own supervision weight, then average studies in the GPU batch.\n            per_study = (cell * w).sum(1) / w.sum(1).clamp_min(1)\n            loss = per_study.mean()\n        probability = torch.sigmoid(logits.detach().float()).cpu().numpy()\n        np.add.at(epoch_prediction_sum, idxs, probability)\n        np.add.at(epoch_prediction_count, idxs, 1)\n        scaler.scale(loss / cfg.grad_accum).backward()\n        if (step + 1) % cfg.grad_accum == 0 or step + 1 == len(loader):\n            scaler.unscale_(optimizer)\n            nn.utils.clip_grad_norm_(head.parameters(), 5.0)\n            nn.utils.clip_grad_norm_(projection_params, 2.0)\n            nn.utils.clip_grad_norm_(backbone_params, 1.0)\n            scaler.step(optimizer); scaler.update(); optimizer.zero_grad(set_to_none=True)\n        losses.append(float(loss.detach()))\n        if (step + 1) % 100 == 0:\n            print(f\"epoch={epoch + 1} step={step + 1}/{len(loader)} loss={np.mean(losses[-100:]):.4f}\")\n\n    seen = epoch_prediction_count > 0\n    averaged_prediction = np.zeros_like(epoch_prediction_sum)\n    averaged_prediction[seen] = epoch_prediction_sum[seen] / epoch_prediction_count[seen, None]\n    online_train_auc, online_train_auc_by_target = weighted_macro_auc(Y[seen], averaged_prediction[seen], W[seen])\n    online_pseudo_auc, _ = weighted_macro_auc(\n        Y[seen & ~gold_mask], averaged_prediction[seen & ~gold_mask], W[seen & ~gold_mask]\n    )\n    gold_train_auc, gold_train_auc_by_target = audit_gold_train_auc(gold_indices)\n    epoch_completed = epoch + 1\n    record = {\n        \"epoch\": epoch_completed, \"loss\": float(np.mean(losses)),\n        \"online_train_auc_proxy\": online_train_auc,\n        \"online_pseudo_auc_proxy\": online_pseudo_auc,\n        \"gold_train_auc_in_sample\": gold_train_auc,\n        \"online_train_auc_by_target\": online_train_auc_by_target,\n        \"gold_train_auc_by_target\": gold_train_auc_by_target,\n        \"pseudo_studies\": len(pseudo_indices), \"gold_studies\": len(gold_indices),\n        \"gold_oversample\": cfg.gold_oversample,\n        \"selection\": \"manual_from_epochs_10_15_20\",\n        \"validation_used\": False,\n        \"epoch_runtime_minutes\": (time.time() - epoch_started) / 60,\n        \"elapsed_minutes_this_run\": (time.time() - started) / 60,\n    }\n    history.append(record); print(json.dumps(record, indent=2))\n    rng_state = {\n        \"python\": random.getstate(), \"numpy\": np.random.get_state(),\n        \"torch\": torch.get_rng_state(),\n        \"cuda\": torch.cuda.get_rng_state_all() if torch.cuda.is_available() else [],\n    }\n    state = {\n        \"version\": \"w01-dinov2s-cache-adapt-last-1\", \"training_signature\": training_signature,\n        \"preprocess_signature\": PREPROCESS_SIGNATURE,\n        \"epoch_completed\": epoch_completed, \"encoder_state\": full_encoder_state(enc),\n        \"head_state\": {k: v.detach().cpu() for k, v in head.state_dict().items()},\n        \"optimizer_state\": optimizer.state_dict(), \"scaler_state\": scaler.state_dict(),\n        \"rng_state\": rng_state, \"history\": history, \"config\": asdict(cfg), \"slot_specs\": SLOT_SPECS,\n    }\n    atomic_torch_save(state, last_path)\n    if epoch_completed in set(cfg.checkpoint_epochs):\n        checkpoint_path = OUT / f\"w01_epoch{epoch_completed}.pt\"\n        exported = {\n            \"version\": \"w01-dinov2s-selectable-cache-adapt-1\",\n            \"training_signature\": training_signature, \"preprocess_signature\": PREPROCESS_SIGNATURE,\n            \"epoch_completed\": epoch_completed, \"encoder_state\": full_encoder_state(enc, half=True),\n            \"history\": history, \"config\": asdict(cfg), \"slot_specs\": SLOT_SPECS,\n            \"selection\": \"manual_checkpoint_candidate\", \"validation_used\": False,\n            \"early_stopping_used\": False,\n        }\n        atomic_torch_save(exported, checkpoint_path)\n        print(f\"CHECKPOINT_SAVED: {checkpoint_path.name}\")\n\ncomplete = bool(history and history[-1][\"epoch\"] == cfg.finetune_epochs)\navailable_checkpoints = [\n    f\"w01_epoch{checkpoint_epoch}.pt\" for checkpoint_epoch in cfg.checkpoint_epochs\n    if (OUT / f\"w01_epoch{checkpoint_epoch}.pt\").is_file()\n]\npreferred_final = OUT / f\"w01_epoch{cfg.finetune_epochs}.pt\"\ndefault_checkpoint = preferred_final.name if preferred_final.is_file() else (last_path.name if complete and last_path.is_file() else None)\nmanifest = {\n    \"version\": \"w01-stage01-cache-1\", \"complete\": complete,\n    \"epoch_completed\": history[-1][\"epoch\"] if history else start_epoch,\n    \"target_epochs\": cfg.finetune_epochs, \"checkpoint_epochs\": list(cfg.checkpoint_epochs),\n    \"available_checkpoints\": available_checkpoints,\n    \"default_checkpoint\": default_checkpoint,\n    \"backbone\": \"DINOv2-S\", \"unfreeze_blocks\": cfg.unfreeze_blocks,\n    \"study_batch_size\": cfg.study_batch_size, \"grad_accum\": cfg.grad_accum,\n    \"selected_tokens_per_study\": sum(cfg.tune_tokens_per_slot),\n    \"tune_encode_batch\": cfg.tune_encode_batch, \"full_data\": True,\n    \"gold_oversample\": cfg.gold_oversample, \"validation_used\": False,\n    \"early_stopping_used\": False, \"training_signature\": training_signature,\n    \"preprocess_signature\": PREPROCESS_SIGNATURE,\n}\natomic_json(manifest, OUT / \"w01_manifest.json\")\nprint(json.dumps(manifest, indent=2))\nif cfg.smoke:\n    assert last_path.is_file() and int(manifest[\"epoch_completed\"]) >= 1\n    if complete:\n        assert manifest[\"default_checkpoint\"] is not None\n        print(\"STAGE01_SMOKE_PASS\")\n    else:\n        print(\"STAGE01_SMOKE_PARTIAL_PASS\")\nprint(\"STAGE01_COMPLETE\" if complete else \"STAGE01_PARTIAL_RESUME\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Kaggle run contract\n\n1. Attach the competition data, `w01_supervision.csv`, and **all** Stage 0.5 output parts.\n2. First Stage 1 run starts from scratch unless a compatible Stage 1 output is attached.\n3. Save partial output after `STAGE01_PARTIAL_RESUME`.\n4. Attach that output on the next run. Auto-resume selects the compatible checkpoint with the highest completed epoch.\n5. Use `RSNA_AUTO_RESUME=0` to force a fresh run, or `RSNA_RESUME_CHECKPOINT=/exact/path/w01_last.pt` to select explicitly.\n","metadata":{}}]}