{"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":[{"id":"48d75d95-3238-4af1-b5e9-897090472711","cell_type":"markdown","source":"# RSNA Knee W01 — Stage 0.5: resumable preprocessed cache\n\nPreprocesses every unique training study once using the exact W01 Stage 1 geometry, series selection, 140 mm crop, 1–99 percentile normalization and 64-candidate layout. The cache is written as study-chunked compressed HDF5 shards.\n\nEach Kaggle run creates one output part. To continue, save the notebook output as a Kaggle Dataset, attach it to the next run together with the competition data, then run this notebook again. Compatible prior manifests are discovered automatically; completed studies are skipped.\n","metadata":{}},{"id":"87a70cb9-bcae-44f6-9114-59b953d7447c","cell_type":"code","source":"from __future__ import annotations\n\nimport hashlib, json, os, random, re, shutil, 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 pydicom\nimport torch\nimport torch.nn.functional as F\nfrom pydicom.dataset import FileDataset, FileMetaDataset\nfrom pydicom.uid import ExplicitVRLittleEndian, generate_uid\nfrom torch.utils.data import Dataset\n\nUID = \"StudyInstanceUID\"\nCACHE_VERSION = \"w01-stage05-cache-v1\"\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    previous_cache_dirs: Optional[str] = os.getenv(\"RSNA_CACHE_DIRS\")\n    output_dir: str = os.getenv(\"RSNA_CACHE_OUTPUT\", \"/kaggle/working/w01_stage05_cache\")\n    image_size: int = 336\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    tokens_per_slot: Tuple[int, ...] = (16, 12, 12, 8, 10, 6)\n    studies_per_shard: int = int(os.getenv(\"RSNA_STUDIES_PER_SHARD\", \"32\"))\n    compression: str = os.getenv(\"RSNA_CACHE_COMPRESSION\", \"gzip\")\n    gzip_level: int = int(os.getenv(\"RSNA_CACHE_GZIP_LEVEL\", \"1\"))\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\", \"12\"))\n    max_output_gib: float = float(os.getenv(\"RSNA_MAX_OUTPUT_GIB\", \"48\"))\n    min_free_gib: float = float(os.getenv(\"RSNA_MIN_FREE_GIB\", \"4\"))\n\ncfg = CFG()\nif cfg.smoke:\n    cfg.output_dir = os.getenv(\"RSNA_CACHE_OUTPUT\", str(Path.cwd() / \"smoke_w01_stage05_cache\"))\n    cfg.image_size = 64\n    cfg.tokens_per_slot = (2, 2, 2, 2, 2, 2)\n    cfg.studies_per_shard = 4\n    cfg.time_budget_minutes = 30\n    cfg.stop_margin_minutes = 1\n    cfg.max_output_gib = 2\n    cfg.min_free_gib = 0.1\nOUT = Path(cfg.output_dir)\nOUT.mkdir(parents=True, exist_ok=True)\nrandom.seed(cfg.seed); np.random.seed(cfg.seed)\nprint(json.dumps(asdict(cfg), indent=2))\n","metadata":{},"outputs":[],"execution_count":null},{"id":"645e2501-f4d7-4ad2-90e2-ac177feec977","cell_type":"markdown","source":"## Competition discovery and synthetic smoke corpus\n","metadata":{}},{"id":"e90ac9c3-40bf-4498-b381-6b4d84e8f911","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)\nif not cfg.smoke:\n    assert sum(cfg.tokens_per_slot) == 64\n\nTARGETS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n    \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\",\n    \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\",\n]\n\ndef _write_synthetic_dicom(path: Path, value: int, z: float, plane: str, laterality: str) -> None:\n    path.parent.mkdir(parents=True, exist_ok=True)\n    meta = FileMetaDataset()\n    meta.MediaStorageSOPClassUID = generate_uid()\n    meta.MediaStorageSOPInstanceUID = generate_uid()\n    meta.TransferSyntaxUID = ExplicitVRLittleEndian\n    ds = FileDataset(str(path), {}, file_meta=meta, preamble=b\"\\0\" * 128)\n    ds.SOPClassUID = meta.MediaStorageSOPClassUID\n    ds.SOPInstanceUID = meta.MediaStorageSOPInstanceUID\n    ds.Rows = ds.Columns = 64\n    ds.SamplesPerPixel = 1\n    ds.PhotometricInterpretation = \"MONOCHROME2\"\n    ds.BitsAllocated = ds.BitsStored = 16\n    ds.HighBit = 15\n    ds.PixelRepresentation = 0\n    ds.PixelSpacing = [0.6, 0.6]\n    ds.InstanceNumber = int(z + 20)\n    ds.Laterality = laterality\n    if plane == \"Sagittal\":\n        ds.ImageOrientationPatient = [0, 1, 0, 0, 0, -1]\n        ds.ImagePositionPatient = [z, 0, 0]\n    elif plane == \"Coronal\":\n        ds.ImageOrientationPatient = [1, 0, 0, 0, 0, -1]\n        ds.ImagePositionPatient = [0, z, 0]\n    else:\n        ds.ImageOrientationPatient = [1, 0, 0, 0, 1, 0]\n        ds.ImagePositionPatient = [0, 0, z]\n    yy, xx = np.mgrid[:64, :64]\n    arr = value + 300 * np.exp(-((xx - 32) ** 2 + (yy - 32) ** 2) / 250)\n    arr += 40 * np.sin((xx + z) / 6)\n    ds.PixelData = np.clip(arr, 0, 65535).astype(np.uint16).tobytes()\n    ds.save_as(path, write_like_original=False)\n\ndef make_smoke_competition(root: Path, n_train: int = 16) -> Path:\n    root.mkdir(parents=True, exist_ok=True)\n    train_rows, train_series = [], []\n    reports = [\n        \"Complete ACL tear with moderate joint effusion. No fracture.\",\n        \"Medial meniscus tear. Mild medial compartment osteoarthritis.\",\n        \"No ligament or meniscal injury. Small Baker cyst.\",\n        \"Bone marrow contusion and synovitis. No fracture.\",\n    ]\n    for i in range(n_train):\n        study = f\"train.study.{i:03d}\"\n        row = {UID: study, \"Report\": reports[i % len(reports)] + f\" template {i // 4}\"}\n        vals = np.array([(i + j) % 3 == 0 for j in range(len(TARGETS))], np.float32)\n        for j, target in enumerate(TARGETS):\n            row[target] = float(vals[j]) if i < 8 else np.nan\n        train_rows.append(row)\n        for s, (_, plane, desired) in enumerate(SLOT_SPECS):\n            fluid = int(desired == \"fluid\")\n            series = f\"{study}.series.{s}\"\n            train_series.append({\n                UID: study, \"SeriesInstanceUID\": series,\n                \"Fluid_Sensitive\": fluid, \"Fat_Suppression\": fluid,\n                \"Anatomical_Plane\": plane,\n            })\n            folder = root / \"train_series\" / study / series\n            for k in range(7):\n                _write_synthetic_dicom(folder / f\"{k:03d}.dcm\", 100 + 7 * i + 11 * s,\n                                       float(k), plane, \"R\" if i % 2 else \"L\")\n    pd.DataFrame(train_rows).to_csv(root / \"train.csv\", index=False)\n    pd.DataFrame(train_series).to_csv(root / \"train_series.csv\", index=False)\n    return root\n\ndef find_comp_root() -> Path:\n    if cfg.comp_root:\n        p = Path(cfg.comp_root)\n        if (p / \"train.csv\").is_file() and (p / \"train_series.csv\").is_file():\n            return p\n        raise FileNotFoundError(f\"Invalid RSNA_COMP_ROOT: {p}\")\n    if cfg.smoke:\n        return make_smoke_competition(Path.cwd() / \"smoke_rsna_data\")\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() and (p / \"train_series.csv\").is_file():\n            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() and (p / \"train_series\").is_dir():\n                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})\nseries_df = pd.read_csv(ROOT / \"train_series.csv\", dtype={UID: str, \"SeriesInstanceUID\": str})\nif train_df[UID].duplicated().any():\n    raise RuntimeError(\"train.csv contains duplicate StudyInstanceUID values\")\nprint(\"root:\", ROOT, \"train studies / series:\", len(train_df), len(series_df))\n","metadata":{},"outputs":[],"execution_count":null},{"id":"4d19d6b3-eef9-484e-b9cf-66b417a7cf56","cell_type":"markdown","source":"## Exact W01 preprocessing\n","metadata":{}},{"id":"33116881-7361-4c8d-9488-e5e76f39a6d5","cell_type":"code","source":"def dcm_files(folder: Path) -> List[Path]:\n    if not folder.is_dir():\n        return []\n    files = list(folder.glob(\"*.dcm\"))\n    return files if files else [p for p in folder.iterdir() if p.is_file()]\n\n\ndef header_key(path: Path) -> Tuple[float, Optional[np.ndarray], Optional[np.ndarray], object]:\n    tags = [\"ImagePositionPatient\", \"ImageOrientationPatient\", \"InstanceNumber\",\n            \"PixelSpacing\", \"Laterality\", \"ImageLaterality\"]\n    ds = pydicom.dcmread(path, stop_before_pixels=True, force=True, specific_tags=tags)\n    try:\n        iop = np.asarray(ds.ImageOrientationPatient, np.float64)\n        ipp = np.asarray(ds.ImagePositionPatient, np.float64)\n        key = float(np.dot(ipp, np.cross(iop[:3], iop[3:])))\n    except Exception:\n        iop = ipp = None\n        key = float(getattr(ds, \"InstanceNumber\", 0))\n    return key, iop, ipp, ds\n\n\ndef sorted_series_files(folder: Path) -> Tuple[List[Path], Optional[np.ndarray], str]:\n    recs = []\n    for p in dcm_files(folder):\n        try:\n            recs.append((p,) + header_key(p))\n        except Exception:\n            pass\n    if not recs:\n        return [], None, \"\"\n    recs.sort(key=lambda x: x[1])\n    # Remove duplicate physical locations (multi-echo/interleaved stacks).\n    dedup, seen = [], set()\n    for rec in recs:\n        k = round(float(rec[1]), 3)\n        if k not in seen:\n            dedup.append(rec); seen.add(k)\n    first_ds = dedup[0][4]\n    lat = str(getattr(first_ds, \"ImageLaterality\", \"\") or getattr(first_ds, \"Laterality\", \"\")).upper()\n    if lat not in (\"L\", \"R\"):\n        ipps = [r[3] for r in dedup if r[3] is not None]\n        if ipps:\n            x = float(np.median([p[0] for p in ipps]))\n            if abs(x) > 20:\n                lat = \"L\" if x > 0 else \"R\"\n    return [r[0] for r in dedup], dedup[0][2], lat\n\n\nTARGET_AXES = {\n    \"Sagittal\": (np.array([0, -1, 0.]), np.array([0, 0, -1.])),\n    \"Coronal\": (np.array([1, 0, 0.]), np.array([0, 0, -1.])),\n    \"Axial\": (np.array([1, 0, 0.]), np.array([0, 1, 0.])),\n}  # desired (column/right axis, row/down axis) in DICOM LPS\n\n\ndef canonicalize(arr: np.ndarray, ds, plane: str) -> Tuple[np.ndarray, Tuple[float, float]]:\n    spacing = list(map(float, getattr(ds, \"PixelSpacing\", [1.0, 1.0])))\n    try:\n        iop = np.asarray(ds.ImageOrientationPatient, np.float64)\n        col_axis, row_axis = iop[:3], iop[3:]\n        want_col, want_row = TARGET_AXES[plane]\n        if abs(np.dot(row_axis, want_col)) > abs(np.dot(col_axis, want_col)):\n            arr = arr.T\n            row_axis, col_axis = col_axis, row_axis\n            spacing = [spacing[1], spacing[0]]\n        if np.dot(col_axis, want_col) < 0:\n            arr = arr[:, ::-1]\n        if np.dot(row_axis, want_row) < 0:\n            arr = arr[::-1]\n    except Exception:\n        pass\n    return np.ascontiguousarray(arr), (spacing[0], spacing[1])\n\n\ndef crop_resize(arr: np.ndarray, spacing: Tuple[float, float], out: int) -> np.ndarray:\n    h, w = arr.shape\n    ch = min(h, max(16, int(round(cfg.crop_mm / max(spacing[0], 1e-3)))))\n    cw = min(w, max(16, int(round(cfg.crop_mm / max(spacing[1], 1e-3)))))\n    y0, x0 = max(0, (h - ch) // 2), max(0, (w - cw) // 2)\n    x = torch.from_numpy(np.ascontiguousarray(arr[y0:y0 + ch, x0:x0 + cw])).float()[None, None]\n    x = F.interpolate(x, (out, out), mode=\"bilinear\", align_corners=False)\n    return x[0, 0].numpy()\n\n\n_FATSAT_RX = re.compile(r\"\\bfs\\b|fatsat|fat sat|\\bstir\\b|\\bspair\\b|\\bspir\\b|water excit|fatsup\")\n_T1_RX = re.compile(r\"\\bt1\\b|\\bt1w\\b\")\n_T2_RX = re.compile(r\"\\bt2\\b|\\bt2w\\b\")\n_PD_RX = re.compile(r\"\\bpd\\b|\\bpdw\\b|proton|dens\")\n_GRE_RX = re.compile(r\"gradient|\\bgre\\b|\\bffe\\b|\\bflash\\b\")\n\n\ndef series_profile(folder: Path, row: pd.Series) -> dict:\n    files = dcm_files(folder)\n    csv_fluid = pd.to_numeric(pd.Series([row.get(\"Fluid_Sensitive\", np.nan)]), errors=\"coerce\").iloc[0]\n    csv_fs = pd.to_numeric(pd.Series([row.get(\"Fat_Suppression\", np.nan)]), errors=\"coerce\").iloc[0]\n    prof = {\n        \"n_files\": len(files), \"fluid\": bool(csv_fluid == 1),\n        \"fatsat\": bool(csv_fs == 1), \"header_ok\": False,\n    }\n    if not files:\n        return prof\n    tags = [\"SeriesDescription\", \"ProtocolName\", \"SequenceName\", \"ScanningSequence\",\n            \"SequenceVariant\", \"ScanOptions\", \"RepetitionTime\", \"EchoTime\"]\n    try:\n        ds = pydicom.dcmread(files[len(files) // 2], stop_before_pixels=True,\n                             force=True, specific_tags=tags)\n        text = \" \".join(str(getattr(ds, k, \"\")) for k in tags[:6]).lower().replace(\"_\", \" \")\n        tr = float(getattr(ds, \"RepetitionTime\", np.nan))\n        te = float(getattr(ds, \"EchoTime\", np.nan))\n        gre = bool(_GRE_RX.search(text))\n        fatsat = bool(_FATSAT_RX.search(text)) or prof[\"fatsat\"]\n        t1 = bool(_T1_RX.search(text)) or (np.isfinite(tr) and tr <= 800 and not gre)\n        t2 = bool(_T2_RX.search(text)) or (np.isfinite(tr) and tr > 800 and np.isfinite(te) and te >= 60)\n        pdw = bool(_PD_RX.search(text)) or (np.isfinite(tr) and tr > 800 and np.isfinite(te) and te < 60)\n        prof.update({\"fluid\": bool(t2 or pdw or (prof[\"fluid\"] and not t1)),\n                     \"fatsat\": fatsat, \"struct\": bool(t1 or pdw or gre),\n                     \"header_ok\": bool(text.strip() or np.isfinite(tr) or np.isfinite(te))})\n    except Exception:\n        pass\n    prof.setdefault(\"struct\", not prof[\"fluid\"])\n    return prof\n\n\ndef choose_series(rows: pd.DataFrame, image_root: Path) -> Dict[int, Path]:\n    chosen, used = {}, set()\n    records = []\n    for _, r in rows.iterrows():\n        folder = image_root / str(r[UID]) / str(r[\"SeriesInstanceUID\"])\n        rec = r.to_dict(); rec.update(series_profile(folder, r)); rec[\"folder\"] = folder\n        records.append(rec)\n    meta = pd.DataFrame(records)\n    if meta.empty:\n        return chosen\n    for slot, (_, plane, desired) in enumerate(SLOT_SPECS):\n        g = meta[(meta[\"Anatomical_Plane\"] == plane) & (~meta[\"SeriesInstanceUID\"].isin(used))].copy()\n        if g.empty:\n            continue\n        want_fluid = desired == \"fluid\"\n        match = np.where(want_fluid, g[\"fluid\"] & g[\"fatsat\"], g[\"struct\"] & ~g[\"fatsat\"])\n        fallback = np.where(want_fluid, g[\"fluid\"], g[\"struct\"])\n        # Prefer a diagnostic 2D stack near the corpus median. This avoids localizers\n        # and very long 3D acquisitions winning merely because they contain more files.\n        n = g[\"n_files\"].clip(lower=1).to_numpy(float)\n        stack_quality = -np.abs(np.log(n / 32.0)) - 2.0 * ((n < 10) | (n > 160))\n        g[\"quality\"] = 6.0 * match.astype(float) + 2.5 * fallback.astype(float) + stack_quality\n        r = g.sort_values([\"quality\", \"n_files\"], ascending=[False, False]).iloc[0]\n        used.add(r[\"SeriesInstanceUID\"]); chosen[slot] = Path(r[\"folder\"])\n    return chosen\n\n\ndef sample_fractions(n_tokens: int) -> np.ndarray:\n    if n_tokens <= 1:\n        return np.asarray([0.5], np.float32)\n    u = np.linspace(-1.0, 1.0, n_tokens)\n    lo, hi = cfg.coverage\n    half = (hi - lo) / 2.0\n    return (0.5 + half * np.sign(u) * np.abs(u) ** 1.6).astype(np.float32)\n\n\ndef read_series_triplets(folder: Path, plane: str, n_tokens: int) -> Tuple[np.ndarray, np.ndarray]:\n    files, iop, lat = sorted_series_files(folder)\n    if not files:\n        return np.zeros((n_tokens, 3, cfg.image_size, cfg.image_size), np.uint8), np.zeros(n_tokens, bool)\n    if lat == \"R\" and plane == \"Sagittal\":\n        files = files[::-1]\n    centers = (sample_fractions(n_tokens) * (len(files) - 1)).round().astype(int)\n    if cfg.input_mode == \"repeat_center\":\n        triplets = [np.asarray([c, c, c], dtype=int) for c in centers]\n    else:\n        triplets = [np.clip([c - cfg.triplet_offset, c, c + cfg.triplet_offset], 0, len(files) - 1) for c in centers]\n    unique = sorted(set(int(i) for tri in triplets for i in tri))\n    decoded: Dict[int, Tuple[np.ndarray, Tuple[float, float]]] = {}\n    sample_values = []\n    for i in unique:\n        try:\n            ds = pydicom.dcmread(files[i], force=True)\n            arr = ds.pixel_array.astype(np.float32)\n            if arr.ndim == 3:\n                arr = arr[len(arr) // 2]\n            arr = arr * float(getattr(ds, \"RescaleSlope\", 1) or 1) + float(getattr(ds, \"RescaleIntercept\", 0) or 0)\n            if str(getattr(ds, \"PhotometricInterpretation\", \"\")).upper() == \"MONOCHROME1\":\n                arr = float(arr.max() + arr.min()) - arr\n            arr, spacing = canonicalize(arr, ds, plane)\n            if lat == \"R\" and plane in (\"Coronal\", \"Axial\"):\n                arr = arr[:, ::-1]\n            decoded[i] = (arr, spacing)\n            sample_values.append(arr[::4, ::4].reshape(-1))\n        except Exception:\n            pass\n    if not decoded:\n        return np.zeros((n_tokens, 3, cfg.image_size, cfg.image_size), np.uint8), np.zeros(n_tokens, bool)\n    q = np.concatenate(sample_values)\n    qlo, qhi = np.percentile(q, [1, 99])\n    out = np.zeros((n_tokens, 3, cfg.image_size, cfg.image_size), np.uint8)\n    mask = np.zeros(n_tokens, bool)\n    available = sorted(decoded)\n    for k, tri in enumerate(triplets):\n        planes = []\n        for i in tri:\n            j = min(available, key=lambda a: abs(a - int(i)))\n            arr, spacing = decoded[j]\n            arr = np.clip((arr - qlo) / max(qhi - qlo, 1e-6), 0, 1)\n            planes.append(crop_resize(arr, spacing, cfg.image_size))\n        out[k] = np.clip(np.stack(planes) * 255, 0, 255).round().astype(np.uint8)\n        mask[k] = True\n    return out, mask\n\n\nclass StudyPixels(Dataset):\n    def __init__(self, studies: pd.DataFrame, series: pd.DataFrame, image_root: Path):\n        self.studies = studies.reset_index(drop=True)\n        self.series_groups = {k: g.copy() for k, g in series.groupby(UID)}\n        self.image_root = image_root\n        self.max_tokens = sum(cfg.tokens_per_slot)\n\n    def __len__(self):\n        return len(self.studies)\n\n    def __getitem__(self, idx: int):\n        study = str(self.studies.iloc[idx][UID])\n        rows = self.series_groups.get(study, pd.DataFrame(columns=series_df.columns))\n        selected = choose_series(rows, self.image_root)\n        images, masks, slots, zpos = [], [], [], []\n        for slot, n_tok in enumerate(cfg.tokens_per_slot):\n            if slot in selected:\n                x, m = read_series_triplets(selected[slot], SLOT_SPECS[slot][1], n_tok)\n            else:\n                x = np.zeros((n_tok, 3, cfg.image_size, cfg.image_size), np.uint8)\n                m = np.zeros(n_tok, bool)\n            images.append(x); masks.append(m)\n            slots.extend([slot] * n_tok)\n            zpos.extend(sample_fractions(n_tok).tolist())\n        return {\n            \"index\": idx,\n            \"uid\": study,\n            \"images\": torch.from_numpy(np.concatenate(images)),\n            \"mask\": torch.from_numpy(np.concatenate(masks)),\n            \"slot\": torch.tensor(slots, dtype=torch.long),\n            \"z\": torch.tensor(zpos, dtype=torch.float32),\n        }\n","metadata":{},"outputs":[],"execution_count":null},{"id":"fb4c687e-76ca-4d0e-b12b-309fb0508f61","cell_type":"markdown","source":"## Cache contract and resumable shard writer\n","metadata":{}},{"id":"065935ee-4bbf-4536-9005-bb99d45fc77d","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","metadata":{},"outputs":[],"execution_count":null},{"id":"e04e3d66-06c8-4ac4-8dc3-0001ad0be607","cell_type":"code","source":"MANIFEST_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 atomic_json(obj: dict, path: Path) -> None:\n    tmp = path.with_suffix(path.suffix + \".tmp\")\n    tmp.write_text(json.dumps(obj, indent=2), encoding=\"utf-8\")\n    os.replace(tmp, path)\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\ndef configured_roots(value: Optional[str]) -> List[Path]:\n    if not value:\n        return []\n    return [Path(x.strip()) for x in re.split(r\"[;,]\", value) if x.strip()]\n\ndef discover_manifests() -> List[Path]:\n    roots = configured_roots(cfg.previous_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    current = OUT / MANIFEST_NAME\n    if current.is_file():\n        found.append(current)\n    return sorted(set(p.resolve() for p in found))\n\ndef load_compatible_manifests() -> Tuple[List[Tuple[Path, dict]], List[Path]]:\n    compatible, ignored = [], []\n    for path in discover_manifests():\n        try:\n            obj = json.loads(path.read_text(encoding=\"utf-8\"))\n            if (obj.get(\"preprocess_signature\") == PREPROCESS_SIGNATURE and\n                    obj.get(\"train_uid_sha256\") == train_uid_sha256):\n                compatible.append((path, obj))\n            else:\n                ignored.append(path)\n        except Exception:\n            ignored.append(path)\n    return compatible, ignored\n\ncompatible, ignored = load_compatible_manifests()\nif ignored:\n    print(\"ignored incompatible/unreadable manifests:\", [str(p) for p in ignored])\n\ncompleted_locations: Dict[str, Tuple[Path, str]] = {}\nmax_part = -1\ncurrent_manifest = None\nfor manifest_path, manifest in compatible:\n    max_part = max(max_part, int(manifest.get(\"part_index\", -1)))\n    if manifest_path == (OUT / MANIFEST_NAME).resolve():\n        current_manifest = manifest\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\"manifest references missing shard: {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 completed_locations and completed_locations[uid] != location:\n                raise RuntimeError(f\"duplicate cached study across parts: {uid}\")\n            completed_locations[uid] = location\n\nif current_manifest is None:\n    part_index = max_part + 1\n    current_manifest = {\n        \"version\": CACHE_VERSION,\n        \"part_index\": part_index,\n        \"preprocess_signature\": PREPROCESS_SIGNATURE,\n        \"preprocess_payload\": PREPROCESS_PAYLOAD,\n        \"train_uid_sha256\": train_uid_sha256,\n        \"total_expected_studies\": len(ordered_uids),\n        \"shards\": [], \"failed_studies\": [],\n    }\nelse:\n    part_index = int(current_manifest[\"part_index\"])\n\ndataset = StudyPixels(train_df[[UID]], series_df, ROOT / \"train_series\")\nremaining_indices = [i for i, uid in enumerate(ordered_uids) if uid not in completed_locations]\nprint(f\"compatible parts={len(compatible)} completed={len(completed_locations)} remaining={len(remaining_indices)} part={part_index}\")\n\ndef output_bytes() -> int:\n    return sum(p.stat().st_size for p in OUT.rglob(\"*\") if p.is_file())\n\ndef free_gib() -> float:\n    return shutil.disk_usage(OUT).free / (1024 ** 3)\n\ndef write_shard(indices: List[int], local_shard_id: int, started: float) -> Tuple[Optional[dict], List[dict], bool]:\n    final = OUT / f\"w01_cache_p{part_index:03d}_s{local_shard_id:04d}.h5\"\n    temp = final.with_suffix(\".partial.h5\")\n    if temp.exists():\n        temp.unlink()\n    studies, failures, stop = [], [], False\n    compression = None if cfg.compression.lower() in (\"\", \"none\") else cfg.compression.lower()\n    compression_opts = cfg.gzip_level if compression == \"gzip\" else None\n    with h5py.File(temp, \"w\") as h5:\n        h5.attrs[\"version\"] = CACHE_VERSION\n        h5.attrs[\"preprocess_signature\"] = PREPROCESS_SIGNATURE\n        for idx in indices:\n            elapsed_min = (time.time() - started) / 60\n            if studies and elapsed_min >= cfg.time_budget_minutes - cfg.stop_margin_minutes:\n                stop = True; break\n            if studies and free_gib() <= cfg.min_free_gib:\n                stop = True; break\n            uid = ordered_uids[idx]\n            try:\n                rec = dataset[idx]\n                key = f\"s{idx:05d}\"\n                group = h5.create_group(key)\n                group.attrs[\"uid\"] = uid\n                images = rec[\"images\"].numpy()\n                group.create_dataset(\n                    \"images\", data=images,\n                    chunks=(1, 3, cfg.image_size, cfg.image_size),\n                    compression=compression, compression_opts=compression_opts,\n                    shuffle=bool(compression),\n                )\n                group.create_dataset(\"mask\", data=rec[\"mask\"].numpy(), compression=None)\n                group.create_dataset(\"slot\", data=rec[\"slot\"].numpy(), compression=None)\n                group.create_dataset(\"z\", data=rec[\"z\"].numpy(), compression=None)\n                studies.append({\"uid\": uid, \"index\": idx, \"key\": key})\n                if len(studies) % 8 == 0:\n                    h5.flush()\n                    print(f\"shard={local_shard_id} studies={len(studies)}/{len(indices)} uid={uid}\")\n            except Exception as exc:\n                failures.append({\"uid\": uid, \"index\": idx, \"error\": repr(exc)})\n    if not studies:\n        if temp.exists(): temp.unlink()\n        return None, failures, stop\n    os.replace(temp, final)\n    entry = {\n        \"file\": final.name,\n        \"study_count\": len(studies),\n        \"studies\": studies,\n        \"size_bytes\": final.stat().st_size,\n        \"sha256\": sha256_file(final),\n    }\n    return entry, failures, stop\n\nstarted = time.time()\nlocal_shard_id = len(current_manifest.get(\"shards\", []))\ncursor = 0\nwhile cursor < len(remaining_indices):\n    if output_bytes() >= cfg.max_output_gib * (1024 ** 3) or free_gib() <= cfg.min_free_gib:\n        print(\"OUTPUT_BUDGET_GRACEFUL_STOP\")\n        break\n    batch_indices = remaining_indices[cursor:cursor + cfg.studies_per_shard]\n    entry, failures, stop = write_shard(batch_indices, local_shard_id, started)\n    current_manifest[\"failed_studies\"].extend(failures)\n    if entry is not None:\n        current_manifest[\"shards\"].append(entry)\n        for item in entry[\"studies\"]:\n            completed_locations[item[\"uid\"]] = (OUT / entry[\"file\"], item[\"key\"])\n        cursor += len(entry[\"studies\"]) + len(failures)\n        local_shard_id += 1\n    else:\n        cursor += len(failures)\n    current_manifest.update({\n        \"completed_global_studies\": len(completed_locations),\n        \"remaining_global_studies\": len(ordered_uids) - len(completed_locations),\n        \"complete_global\": len(completed_locations) == len(ordered_uids),\n        \"elapsed_minutes_this_run\": (time.time() - started) / 60,\n        \"output_size_gib\": output_bytes() / (1024 ** 3),\n    })\n    atomic_json(current_manifest, OUT / MANIFEST_NAME)\n    print(f\"SHARD_SAVED: {entry['file'] if entry else 'none'} global={len(completed_locations)}/{len(ordered_uids)}\")\n    if stop:\n        print(\"TIME_OR_DISK_GRACEFUL_STOP\")\n        break\n\ncurrent_manifest.update({\n    \"completed_global_studies\": len(completed_locations),\n    \"remaining_global_studies\": len(ordered_uids) - len(completed_locations),\n    \"complete_global\": len(completed_locations) == len(ordered_uids),\n    \"elapsed_minutes_this_run\": (time.time() - started) / 60,\n    \"output_size_gib\": output_bytes() / (1024 ** 3),\n})\natomic_json(current_manifest, OUT / MANIFEST_NAME)\nprint(json.dumps({k: current_manifest[k] for k in (\n    \"part_index\", \"completed_global_studies\", \"remaining_global_studies\",\n    \"complete_global\", \"elapsed_minutes_this_run\", \"output_size_gib\"\n)}, indent=2))\n\nif cfg.smoke:\n    assert current_manifest[\"shards\"]\n    first = current_manifest[\"shards\"][0]\n    with h5py.File(OUT / first[\"file\"], \"r\") as h5:\n        item = first[\"studies\"][0]\n        group = h5[item[\"key\"]]\n        assert str(group.attrs[\"uid\"]) == item[\"uid\"]\n        assert group[\"images\"].shape == (sum(cfg.tokens_per_slot), 3, cfg.image_size, cfg.image_size)\n        assert group[\"images\"].dtype == np.uint8\n    print(\"STAGE05_SMOKE_PASS\" if current_manifest[\"complete_global\"] else \"STAGE05_SMOKE_PARTIAL_PASS\")\nprint(\"STAGE05_COMPLETE\" if current_manifest[\"complete_global\"] else \"STAGE05_PARTIAL_RESUME\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"b453d851-7dc5-479c-b733-1f67e299aa2f","cell_type":"markdown","source":"## Kaggle resume usage\n\n- First run: attach the competition data and run all cells.\n- Save the output as a Kaggle Dataset.\n- Next run: attach the competition data plus every previous Stage 0.5 output part; run this notebook again.\n- Repeat until the final line is `STAGE05_COMPLETE`.\n- Stage 1 must be given all Stage 0.5 parts. It validates exact UID coverage and refuses incomplete or incompatible caches.\n","metadata":{}}]}