{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# 1. Imports, reproducibility and user settings\nfrom __future__ import annotations\n\nimport gc\nimport hashlib\nimport json\nimport math\nimport os\nimport re\nimport time\nimport traceback\nimport warnings\nfrom collections import defaultdict\nfrom concurrent.futures import ThreadPoolExecutor, as_completed\nfrom dataclasses import dataclass\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom transformers import Dinov2Config, Dinov2Model\n\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\nos.environ.setdefault(\"TOKENIZERS_PARALLELISM\", \"false\")\nfor _var in (\"OMP_NUM_THREADS\", \"OPENBLAS_NUM_THREADS\", \"MKL_NUM_THREADS\"):\n    os.environ.setdefault(_var, \"4\")\n\nSEED = 2026\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\n# Safe defaults. Package/member metadata overrides image, slice and band values.\nCROP_MM = 130.0\nDEFAULT_IMG = 336\nDEFAULT_GROUP = 3\nDEFAULT_SLICES = 6\nDEFAULT_BAND = (0.20, 0.80)\nMAX_TTA_WINDOWS = 4\nMAX_MEMBERS = 20          # set lower only for a quick smoke test\nREQUIRE_GPU = True\nDECODE_WORKERS = min(12, max(2, os.cpu_count() or 2))\nHEADER_WORKERS = min(16, max(2, os.cpu_count() or 2))\n\nDEVICE = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nGPU_COUNT = torch.cuda.device_count() if DEVICE.type == \"cuda\" else 0\nif REQUIRE_GPU and DEVICE.type != \"cuda\":\n    raise RuntimeError(\"GPU is not enabled. In Kaggle: Settings → Accelerator → GPU, then restart and Run All.\")\n\nif DEVICE.type == \"cuda\":\n    torch.backends.cuda.matmul.allow_tf32 = True\n    torch.backends.cudnn.allow_tf32 = True\n    torch.backends.cudnn.benchmark = True\n    gpu_gb = torch.cuda.get_device_properties(0).total_memory / 1024**3\n    BATCH_SIZE = 8 if gpu_gb >= 20 else 4\nelse:\n    gpu_gb = 0\n    BATCH_SIZE = 1\n\nT0 = time.time()\ndef log(message):\n    print(f\"[{time.time() - T0:7.1f}s] {message}\", flush=True)\n\nlog(f\"torch={torch.__version__} device={DEVICE} GPUs={GPU_COUNT} batch={BATCH_SIZE}\")\nif DEVICE.type == \"cuda\":\n    log(f\"GPU={torch.cuda.get_device_name(0)} memory={gpu_gb:.1f} GiB\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T08:28:55.298713Z","iopub.execute_input":"2026-08-17T08:28:55.299929Z","iopub.status.idle":"2026-08-17T08:29:04.85625Z","shell.execute_reply.started":"2026-08-17T08:28:55.299875Z","shell.execute_reply":"2026-08-17T08:29:04.855106Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2. Competition discovery and schema validation\ndef find_competition_root():\n    direct = [\n        Path(\"/kaggle/input/competitions/rsna-knee-abnormality-detection\"),\n        Path(\"/kaggle/input/rsna-knee-abnormality-detection\"),\n    ]\n    for path in direct:\n        if (path / \"test.csv\").is_file() and (path / \"sample_submission.csv\").is_file():\n            return path\n    base = Path(\"/kaggle/input\")\n    if base.is_dir():\n        for root, dirs, files in os.walk(base):\n            dirs[:] = [d for d in dirs if d not in {\"train_series\", \"test_series\"}]\n            if \"test.csv\" in files and \"sample_submission.csv\" in files:\n                return Path(root)\n    raise FileNotFoundError(\n        \"Competition data was not found. Add the RSNA Knee Abnormality Detection competition as a notebook input.\"\n    )\n\nROOT = find_competition_root()\nTEST_DF = pd.read_csv(ROOT / \"test.csv\")\nSAMPLE = pd.read_csv(ROOT / \"sample_submission.csv\")\n\nif \"StudyInstanceUID\" not in TEST_DF or \"StudyInstanceUID\" not in SAMPLE:\n    raise ValueError(\"StudyInstanceUID is missing from test.csv or sample_submission.csv\")\n\nTEST_DF[\"StudyInstanceUID\"] = TEST_DF[\"StudyInstanceUID\"].astype(str)\nSAMPLE[\"StudyInstanceUID\"] = SAMPLE[\"StudyInstanceUID\"].astype(str)\nTARGETS = [column for column in SAMPLE.columns if column != \"StudyInstanceUID\"]\nif len(TARGETS) != 12:\n    raise ValueError(f\"Expected 12 target columns, found {len(TARGETS)}: {TARGETS}\")\n\nTEST_SERIES_DIR = ROOT / \"test_series\"\nTEST_SERIES_CSV = ROOT / \"test_series.csv\"\nif not TEST_SERIES_DIR.is_dir():\n    raise FileNotFoundError(f\"Missing test_series directory: {TEST_SERIES_DIR}\")\n\nSERIES_TABLE = pd.read_csv(TEST_SERIES_CSV) if TEST_SERIES_CSV.is_file() else pd.DataFrame()\nlog(f\"competition={ROOT}\")\nlog(f\"test studies={len(TEST_DF)} targets={TARGETS}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T08:29:08.529258Z","iopub.execute_input":"2026-08-17T08:29:08.53011Z","iopub.status.idle":"2026-08-17T08:29:08.555813Z","shell.execute_reply.started":"2026-08-17T08:29:08.530071Z","shell.execute_reply":"2026-08-17T08:29:08.554976Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 3. Discover and expand standalone checkpoints or packaged ensembles\n@dataclass\nclass LoadedMember:\n    state: dict\n    config: dict\n    name: str\n    source: Path\n\nCONFIG_KEYS = {\n    \"img\", \"img_size\", \"slices\", \"group\", \"n_group\", \"band\", \"pool\", \"prior\",\n    \"hidden\", \"weight\", \"target_weights\", \"gammas\", \"model_variant\", \"crop_mm\",\n}\n\ndef package_config(obj):\n    if not isinstance(obj, dict):\n        return {}\n    cfg = {}\n    if isinstance(obj.get(\"config\"), dict):\n        cfg.update(obj[\"config\"])\n    for key in CONFIG_KEYS:\n        if key in obj:\n            cfg[key] = obj[key]\n    if \"img\" not in cfg and \"img_size\" in cfg:\n        cfg[\"img\"] = cfg[\"img_size\"]\n    if \"slices\" not in cfg and \"n_group\" in cfg:\n        cfg[\"slices\"] = int(cfg.get(\"group\", DEFAULT_GROUP)) * int(cfg[\"n_group\"])\n    return cfg\n\ndef clean_state_dict(obj):\n    if not isinstance(obj, dict):\n        return {}\n    clean = {}\n    for key, value in obj.items():\n        if not isinstance(key, str) or not torch.is_tensor(value):\n            continue\n        while key.startswith((\"module.\", \"model.\", \"_orig_mod.\")):\n            key = key.split(\".\", 1)[1]\n        if key == \"head.slot_embedding\":\n            key = \"head.slot_emb\"\n        clean[key] = value\n    return clean\n\ndef looks_like_model_state(obj):\n    state = clean_state_dict(obj)\n    keys = list(state)\n    return (\n        len(state) >= 10\n        and any(key.startswith(\"backbone.\") for key in keys)\n        and any(key.startswith(\"head.\") for key in keys)\n    )\n\ndef unwrap_model_state(obj):\n    if not isinstance(obj, dict):\n        raise TypeError(f\"Expected checkpoint dictionary, received {type(obj).__name__}\")\n    if looks_like_model_state(obj):\n        return clean_state_dict(obj)\n    for key in (\"state\", \"state_dict\", \"model_state_dict\", \"model\", \"weights\", \"net\", \"network\"):\n        value = obj.get(key)\n        if isinstance(value, dict):\n            try:\n                return unwrap_model_state(value)\n            except (TypeError, RuntimeError):\n                pass\n    for key, value in obj.items():\n        if str(key).lower() in {\"optimizer\", \"optimizer_state_dict\", \"scheduler\", \"scaler\"}:\n            continue\n        if isinstance(value, dict):\n            try:\n                return unwrap_model_state(value)\n            except (TypeError, RuntimeError):\n                pass\n    raise RuntimeError(f\"No full backbone+head state_dict found; keys={list(obj)[:30]}\")\n\ndef state_identity(state):\n    # Cheap deterministic identity used only to prevent duplicate members.\n    pieces = []\n    for key in sorted(state):\n        tensor = state[key]\n        pieces.append(f\"{key}:{tuple(tensor.shape)}\")\n    return hashlib.sha1(\"|\".join(pieces).encode()).hexdigest()\n\ndef candidate_checkpoint_files():\n    base = Path(\"/kaggle/input\")\n    candidates = []\n    if not base.is_dir():\n        return candidates\n    for path in base.rglob(\"*\"):\n        if not path.is_file() or path.suffix.lower() not in {\".pth\", \".pt\", \".ckpt\", \".bin\"}:\n            continue\n        low = str(path).lower()\n        if any(part in low for part in (\"train_series\", \"test_series\", \"dinov2\", \"metaresearch\")):\n            continue\n        if path.name.lower() in {\"pytorch_model.bin\", \"optimizer.pt\", \"scheduler.pt\", \"training_args.bin\"}:\n            continue\n        if path.stat().st_size >= 1_000_000 and (\"rsna\" in low or \"knee\" in low):\n            candidates.append(path)\n    return sorted(candidates, key=lambda p: (\"pilkwang\" not in str(p).lower(), str(p)))\n\ndef package_children(package):\n    if not isinstance(package, dict):\n        return None\n    for key in (\"members\", \"models\", \"checkpoints\", \"folds\"):\n        if isinstance(package.get(key), (list, tuple, dict)):\n            return package[key]\n    return None\n\ndef expand_file(path):\n    log(f\"opening checkpoint package: {path}\")\n    package = torch.load(path, map_location=\"cpu\", weights_only=False)\n    top_cfg = package_config(package)\n    if isinstance(package, dict) and package.get(\"targets\") is not None:\n        if list(package[\"targets\"]) != list(TARGETS):\n            raise RuntimeError(f\"Target order mismatch in {path.name}: {package['targets']} != {TARGETS}\")\n    children = package_children(package)\n    if children is None:\n        iterable = [(path.stem, package)]\n    elif isinstance(children, dict):\n        iterable = list(children.items())\n    else:\n        iterable = list(enumerate(children))\n\n    output = []\n    for index, child in iterable:\n        cfg = dict(top_cfg)\n        cfg.update(package_config(child))\n        name = str(index)\n        obj = child\n        if isinstance(child, dict):\n            name = str(child.get(\"id\", child.get(\"name\", child.get(\"run_name\", index))))\n            # Some packages point to a separate member file.\n            if \"file\" in child and not looks_like_model_state(child):\n                member_path = path.parent / str(child[\"file\"])\n                if member_path.is_file():\n                    obj = torch.load(member_path, map_location=\"cpu\", weights_only=False)\n        state = unwrap_model_state(obj)\n        output.append(LoadedMember(state=state, config=cfg, name=f\"{path.stem}/{name}\", source=path))\n    del package\n    gc.collect()\n    log(f\"expanded {path.name}: {len(output)} model member(s)\")\n    return output\n\nfiles = candidate_checkpoint_files()\nif not files:\n    attached = [str(path) for path in sorted(Path(\"/kaggle/input\").iterdir()) if path.is_dir()]\n    raise FileNotFoundError(\n        \"No RSNA-trained checkpoint package was found. Add Input → Datasets → \"\n        \"pilkwang/rsna-knee-weights, restart the session, and Run All. \"\n        f\"Attached inputs: {attached}\"\n    )\n\n# Prefer the first package that successfully yields complete models. This prevents\n# accidental double counting when two datasets contain copies of the same ensemble.\nLOADED_MEMBERS = []\nerrors = []\nfor checkpoint_path in files:\n    try:\n        expanded = expand_file(checkpoint_path)\n        if expanded:\n            LOADED_MEMBERS.extend(expanded)\n            # A packaged ensemble is sufficient; do not scan duplicate packages.\n            if len(expanded) > 1:\n                break\n    except Exception as exc:\n        errors.append(f\"{checkpoint_path}: {exc}\")\n        log(f\"ignored incompatible file: {checkpoint_path.name}: {exc}\")\n\n# Remove exact architecture duplicates only when they came from repeated discovery.\nunique = {}\nfor member in LOADED_MEMBERS:\n    identity = (member.name, state_identity(member.state))\n    unique[identity] = member\nLOADED_MEMBERS = list(unique.values())[:MAX_MEMBERS]\n\nif not LOADED_MEMBERS:\n    raise RuntimeError(\"Checkpoint files were found but no complete model could be extracted:\\n\" + \"\\n\".join(errors[:10]))\n\nlog(f\"usable ensemble members={len(LOADED_MEMBERS)}\")\nfor member in LOADED_MEMBERS[:5]:\n    log(f\"  {member.name}: config={member.config}, tensors={len(member.state)}\")\nif len(LOADED_MEMBERS) > 5:\n    log(f\"  ... plus {len(LOADED_MEMBERS)-5} more\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T08:29:15.016223Z","iopub.execute_input":"2026-08-17T08:29:15.016542Z","iopub.status.idle":"2026-08-17T08:36:30.948081Z","shell.execute_reply.started":"2026-08-17T08:29:15.016516Z","shell.execute_reply":"2026-08-17T08:36:30.947298Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 4. MRI series discovery, sequence classification and slot assignment\nSLOTS = [\n    (\"SAG_FLUID_FS\", \"Sagittal\", \"fluid\", True),\n    (\"COR_FLUID_FS\", \"Coronal\", \"fluid\", True),\n    (\"AX_FLUID_FS\", \"Axial\", \"fluid\", True),\n    (\"SAG_FLUID_NOFS\", \"Sagittal\", \"fluid\", False),\n    (\"COR_T1\", \"Coronal\", \"t1\", False),\n    (\"SAG_T1\", \"Sagittal\", \"t1\", False),\n]\nN_SLOT = len(SLOTS)\n\n_SEP = re.compile(r\"[_\\-.]+\")\n_FATSAT = re.compile(r\"\\bfs\\b|fat ?sat|fatsup|\\bstir\\b|\\bspair\\b|\\bspir\\b|water excit|\\btirm\\b\", re.I)\n_T1 = re.compile(r\"\\bt1\\b|\\bt1w\\b\", re.I)\n_T2 = re.compile(r\"\\bt2\\b|\\bt2w\\b\", re.I)\n_PD = re.compile(r\"\\bpd\\b|\\bpdw\\b|proton|density|\\bdp\\b\", re.I)\n_LOCALIZER = re.compile(r\"locali[sz]er|scout|survey|calibration|three[ -]?plane|3[ -]?plane|pilot\", re.I)\n\ndef as_float(value, default=np.nan):\n    try:\n        return float(value)\n    except Exception:\n        return default\n\nplane_map = {}\nif not SERIES_TABLE.empty and {\"SeriesInstanceUID\", \"Anatomical_Plane\"}.issubset(SERIES_TABLE.columns):\n    plane_map = dict(zip(SERIES_TABLE.SeriesInstanceUID.astype(str), SERIES_TABLE.Anatomical_Plane.astype(str)))\n\ndef first_header(series_dir):\n    files = sorted(path for path in series_dir.iterdir() if path.is_file())\n    if not files:\n        return None, files\n    for path in (files[0], files[len(files)//2], files[-1]):\n        try:\n            return pydicom.dcmread(str(path), stop_before_pixels=True, force=True), files\n        except Exception:\n            pass\n    return None, files\n\ndef infer_plane(ds, table_value=\"\"):\n    value = str(table_value).strip().capitalize()\n    if value in {\"Sagittal\", \"Coronal\", \"Axial\"}:\n        return value\n    try:\n        iop = np.asarray(ds.ImageOrientationPatient, float)\n        normal = np.cross(iop[:3], iop[3:6])\n        return (\"Sagittal\", \"Coronal\", \"Axial\")[int(np.argmax(np.abs(normal)))]\n    except Exception:\n        return \"Unknown\"\n\ndef infer_laterality(ds):\n    for key in (\"ImageLaterality\", \"Laterality\"):\n        value = str(getattr(ds, key, \"\")).strip().upper()\n        if value[:1] in {\"L\", \"R\"}:\n            return value[:1]\n    try:\n        ipp = np.asarray(ds.ImagePositionPatient, float)\n        iop = np.asarray(ds.ImageOrientationPatient, float)\n        spacing = np.asarray(ds.PixelSpacing, float)\n        centre = ipp + iop[:3] * spacing[1] * float(ds.Columns)/2 + iop[3:6] * spacing[0] * float(ds.Rows)/2\n        if abs(centre[0]) >= 12:\n            return \"R\" if centre[0] < 0 else \"L\"\n    except Exception:\n        pass\n    return \"U\"\n\ndef contrast_info(ds):\n    desc = \" \".join(str(getattr(ds, key, \"\")) for key in (\"SeriesDescription\", \"ProtocolName\", \"SequenceName\"))\n    clean = _SEP.sub(\" \", desc.lower())\n    options = str(getattr(ds, \"ScanOptions\", \"\")).upper()\n    fatsat = bool(_FATSAT.search(clean)) or any(token in options for token in (\"FS\", \"FATSAT\", \"FSAT\"))\n    tr = as_float(getattr(ds, \"RepetitionTime\", np.nan))\n    te = as_float(getattr(ds, \"EchoTime\", np.nan))\n    if _T1.search(clean) and not _T2.search(clean) and not _PD.search(clean):\n        contrast = \"t1\"\n    elif _T2.search(clean) or _PD.search(clean):\n        contrast = \"fluid\"\n    elif np.isfinite(tr) and tr < 900 and (not np.isfinite(te) or te < 40):\n        contrast = \"t1\"\n    elif (np.isfinite(te) and te > 45) or (np.isfinite(tr) and tr >= 900):\n        contrast = \"fluid\"\n    else:\n        contrast = \"unknown\"\n    return clean, contrast, fatsat\n\ndef scan_study(uid):\n    study_dir = TEST_SERIES_DIR / str(uid)\n    records = []\n    if not study_dir.is_dir():\n        return str(uid), records\n    for series_dir in sorted(path for path in study_dir.iterdir() if path.is_dir()):\n        ds, files = first_header(series_dir)\n        if ds is None or not files:\n            continue\n        desc, contrast, fatsat = contrast_info(ds)\n        records.append({\n            \"uid\": series_dir.name, \"files\": files, \"n_files\": len(files), \"desc\": desc,\n            \"contrast\": contrast, \"fatsat\": fatsat,\n            \"plane\": infer_plane(ds, plane_map.get(series_dir.name, \"\")),\n            \"laterality\": infer_laterality(ds),\n        })\n    return str(uid), records\n\ndef slot_score(record, plane, contrast, fatsat):\n    score = 0.0\n    score += 8.0 if record[\"plane\"] == plane else -20.0\n    score += 4.0 if record[\"contrast\"] == contrast else (-1.5 if record[\"contrast\"] == \"unknown\" else -5.0)\n    score += 2.5 if record[\"fatsat\"] == fatsat else -2.0\n    score += min(math.log1p(record[\"n_files\"]), 4.0)\n    if _LOCALIZER.search(record[\"desc\"]):\n        score -= 30.0\n    return score\n\nSLOT_MAP, LATERALITY = {}, {}\nwith ThreadPoolExecutor(max_workers=HEADER_WORKERS) as executor:\n    futures = [executor.submit(scan_study, uid) for uid in TEST_DF.StudyInstanceUID]\n    for future in as_completed(futures):\n        uid, records = future.result()\n        chosen = {}\n        for slot_name, plane, contrast, fatsat in SLOTS:\n            candidates = sorted(records, key=lambda rec: slot_score(rec, plane, contrast, fatsat), reverse=True)\n            if candidates and slot_score(candidates[0], plane, contrast, fatsat) > -5:\n                chosen[slot_name] = candidates[0]\n        SLOT_MAP[uid] = chosen\n        votes = [rec[\"laterality\"] for rec in records if rec[\"laterality\"] in {\"L\", \"R\"}]\n        LATERALITY[uid] = max(set(votes), key=votes.count) if votes else \"U\"\n\ncoverage = {name: float(np.mean([name in SLOT_MAP.get(uid, {}) for uid in TEST_DF.StudyInstanceUID])) for name, *_ in SLOTS}\nlog(\"slot coverage: \" + \", \".join(f\"{key}={value:.1%}\" for key, value in coverage.items()))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T08:38:19.920432Z","iopub.execute_input":"2026-08-17T08:38:19.9208Z","iopub.status.idle":"2026-08-17T08:38:20.288556Z","shell.execute_reply.started":"2026-08-17T08:38:19.920771Z","shell.execute_reply":"2026-08-17T08:38:20.28771Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 5. Geometry-correct DICOM decoding and configuration cache\ndef natural_key(path):\n    return tuple(int(x) if x.isdigit() else x.lower() for x in re.split(r\"(\\d+)\", path.name))\n\ndef ordered_files(files):\n    rows = []\n    for path in files:\n        coordinate = instance = np.nan\n        try:\n            ds = pydicom.dcmread(str(path), stop_before_pixels=True, force=True)\n            ipp = np.asarray(getattr(ds, \"ImagePositionPatient\", []), float)\n            iop = np.asarray(getattr(ds, \"ImageOrientationPatient\", []), float)\n            if len(ipp) >= 3 and len(iop) >= 6:\n                coordinate = float(np.dot(ipp[:3], np.cross(iop[:3], iop[3:6])))\n            instance = as_float(getattr(ds, \"InstanceNumber\", np.nan))\n        except Exception:\n            pass\n        rows.append((path, coordinate, instance))\n    if np.mean([np.isfinite(row[1]) for row in rows]) >= 0.8:\n        return [row[0] for row in sorted(rows, key=lambda row: (row[1] if np.isfinite(row[1]) else np.inf, natural_key(row[0])))]\n    if np.mean([np.isfinite(row[2]) for row in rows]) >= 0.8:\n        return [row[0] for row in sorted(rows, key=lambda row: (row[2] if np.isfinite(row[2]) else np.inf, natural_key(row[0])))]\n    return sorted(files, key=natural_key)\n\ndef decode_series(record, n_slices, image_size, band, crop_mm):\n    files = ordered_files(record[\"files\"])\n    if not files:\n        return None\n    lo_index = int(round(float(band[0]) * max(0, len(files)-1)))\n    hi_index = int(round(float(band[1]) * max(0, len(files)-1)))\n    indices = np.rint(np.linspace(lo_index, max(lo_index, hi_index), n_slices)).astype(int)\n    planes, headers = [], []\n    for index in indices:\n        try:\n            ds = pydicom.dcmread(str(files[int(index)]), force=True)\n            arr = ds.pixel_array.astype(np.float32)\n            arr = arr * as_float(getattr(ds, \"RescaleSlope\", 1), 1) + as_float(getattr(ds, \"RescaleIntercept\", 0), 0)\n            planes.append(arr)\n            headers.append(ds)\n        except Exception:\n            planes.append(None)\n            headers.append(None)\n    good = [index for index, arr in enumerate(planes) if arr is not None]\n    if not good:\n        return None\n    for index, arr in enumerate(planes):\n        if arr is None:\n            nearest = min(good, key=lambda item: abs(item-index))\n            planes[index], headers[index] = planes[nearest], headers[nearest]\n    shape = planes[good[0]].shape\n    planes = [arr if arr.shape == shape else np.zeros(shape, np.float32) for arr in planes]\n    volume = np.stack(planes)\n    ds0 = headers[good[0]]\n    try:\n        spacing = np.asarray(ds0.PixelSpacing, float)\n        crop_h = min(int(round(crop_mm / spacing[0])), volume.shape[-2])\n        crop_w = min(int(round(crop_mm / spacing[1])), volume.shape[-1])\n        y0 = (volume.shape[-2] - crop_h)//2\n        x0 = (volume.shape[-1] - crop_w)//2\n        volume = volume[:, y0:y0+crop_h, x0:x0+crop_w]\n    except Exception:\n        pass\n    finite = volume[np.isfinite(volume)]\n    if finite.size == 0:\n        return None\n    low, high = np.percentile(finite, [1, 99])\n    volume = np.nan_to_num((volume-low)/max(high-low, 1e-6), nan=0.0, posinf=1.0, neginf=0.0)\n    tensor = torch.from_numpy(np.ascontiguousarray(volume)).unsqueeze(0).float()\n    tensor = F.interpolate(tensor, size=(image_size, image_size), mode=\"bilinear\", align_corners=False).squeeze(0)\n    if record.get(\"right_knee\", False):\n        tensor = torch.flip(tensor, [-1] if record[\"plane\"] in {\"Coronal\", \"Axial\"} else [0])\n    return tensor.mul(255).round().clamp(0, 255).byte().numpy()\n\nclass StudyCache:\n    def __init__(self, image_size, n_slices, band, crop_mm):\n        self.image_size = int(image_size)\n        self.n_slices = int(n_slices)\n        self.band = tuple(float(value) for value in band)\n        self.crop_mm = float(crop_mm)\n        self.studies = TEST_DF.StudyInstanceUID.astype(str).tolist()\n        shape = (len(self.studies), N_SLOT, self.n_slices, self.image_size, self.image_size)\n        required_gib = np.prod(shape, dtype=np.int64) / 1024**3\n        try:\n            import psutil\n            available_gib = psutil.virtual_memory().available / 1024**3\n        except Exception:\n            available_gib = 24.0\n        self.memmap_path = None\n        if required_gib <= max(2.0, available_gib * 0.55):\n            self.pixels = np.zeros(shape, np.uint8)\n            storage = \"RAM\"\n        else:\n            self.memmap_path = Path(\"/kaggle/working\") / f\"rsna_cache_{self.image_size}_{self.n_slices}.dat\"\n            self.pixels = np.memmap(self.memmap_path, mode=\"w+\", dtype=np.uint8, shape=shape)\n            self.pixels[:] = 0\n            storage = \"working-disk memmap\"\n        self.mask = np.zeros((len(self.studies), N_SLOT), np.float32)\n        log(f\"cache={shape}, {required_gib:.2f} GiB, storage={storage}\")\n\n        jobs = []\n        for row, uid in enumerate(self.studies):\n            for slot, (slot_name, *_rest) in enumerate(SLOTS):\n                record = SLOT_MAP.get(uid, {}).get(slot_name)\n                if record is not None:\n                    record = dict(record)\n                    record[\"right_knee\"] = LATERALITY.get(uid) == \"R\"\n                    jobs.append((row, slot, record))\n\n        def worker(job):\n            row, slot, record = job\n            arr = decode_series(record, self.n_slices, self.image_size, self.band, self.crop_mm)\n            return row, slot, arr\n\n        failed = 0\n        with ThreadPoolExecutor(max_workers=DECODE_WORKERS) as executor:\n            futures = [executor.submit(worker, job) for job in jobs]\n            for count, future in enumerate(as_completed(futures), 1):\n                try:\n                    row, slot, arr = future.result()\n                    if arr is not None:\n                        self.pixels[row, slot] = arr\n                        self.mask[row, slot] = 1.0\n                    else:\n                        failed += 1\n                except Exception:\n                    failed += 1\n                if count % 100 == 0:\n                    log(f\"decoded {count}/{len(jobs)} slot-series\")\n        empty = int(np.sum(self.mask.sum(1) == 0))\n        log(f\"cache ready: failed slots={failed}, studies without decoded slots={empty}\")\n\n    def close(self):\n        if isinstance(self.pixels, np.memmap):\n            self.pixels.flush()\n            mmap = getattr(self.pixels, \"_mmap\", None)\n            if mmap is not None:\n                mmap.close()\n        del self.pixels\n        if self.memmap_path is not None and self.memmap_path.exists():\n            self.memmap_path.unlink()\n\nclass CacheDataset(Dataset):\n    def __init__(self, cache): self.cache = cache\n    def __len__(self): return len(self.cache.studies)\n    def __getitem__(self, index):\n        return torch.from_numpy(np.asarray(self.cache.pixels[index])), torch.from_numpy(self.cache.mask[index])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T08:38:37.109983Z","iopub.execute_input":"2026-08-17T08:38:37.110391Z","iopub.status.idle":"2026-08-17T08:38:37.143421Z","shell.execute_reply.started":"2026-08-17T08:38:37.110355Z","shell.execute_reply":"2026-08-17T08:38:37.142523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 6. Reconstruct DINOv2 and the slot-attention classification model from checkpoint shapes\nSLOT_PRIOR_STRENGTH = 0.55\nSLOT_PRIOR_TABLE = {\n    \"ACL\": (0, 3, 5), \"MCL\": (1, 4),\n    \"Medial Meniscus\": (0, 1, 3, 4), \"Lateral Meniscus\": (0, 1, 3, 4),\n    \"Medial OA\": (1, 4, 5), \"Lateral OA\": (1, 4, 5), \"PF OA\": (0, 2, 5),\n    \"Effusion\": (0, 2), \"Synovitis\": (0, 2), \"Baker's\": (0,),\n    \"Contusion\": (0, 1, 2), \"Fracture\": (0, 1, 2, 4, 5),\n}\n\nclass SlotHead(nn.Module):\n    def __init__(self, dim, n_slot, n_out, hidden=256, dropout=0.2, prior=False):\n        super().__init__()\n        self.proj = nn.Sequential(nn.LayerNorm(dim), nn.Linear(dim, hidden), nn.GELU())\n        self.slot_emb = nn.Parameter(torch.randn(n_slot, hidden) * 0.02)\n        self.query = nn.Parameter(torch.randn(n_out, hidden) * 0.02)\n        self.drop = nn.Dropout(dropout)\n        self.out = nn.Linear(hidden, n_out)\n        self.hidden = hidden\n        self.prior = prior\n        if prior:\n            table = torch.zeros(n_out, n_slot)\n            if n_slot == N_SLOT:\n                for target, slots in SLOT_PRIOR_TABLE.items():\n                    if target in TARGETS:\n                        table[TARGETS.index(target), list(slots)] = SLOT_PRIOR_STRENGTH\n            self.register_buffer(\"slot_prior\", table)\n\n    def forward(self, features, mask):\n        hidden = self.proj(features) + self.slot_emb\n        attention = torch.einsum(\"bsh,th->bts\", hidden, self.query) / math.sqrt(self.hidden)\n        if self.prior:\n            attention = attention + self.slot_prior.unsqueeze(0)\n        attention = attention.masked_fill(mask.unsqueeze(1) < 0.5, -1e4).softmax(-1)\n        context = self.drop(torch.einsum(\"bts,bsh->bth\", attention, hidden))\n        return (context * self.out.weight.unsqueeze(0)).sum(-1) + self.out.bias\n\nclass KneeModel(nn.Module):\n    def __init__(self, backbone, pool, n_slot, hidden, prior):\n        super().__init__()\n        self.backbone = backbone\n        self.pool = pool\n        parts = 3 if pool == \"cls_mean_focal\" else 2\n        self.head = SlotHead(backbone.config.hidden_size * parts, n_slot, len(TARGETS), hidden, prior=prior)\n        self.register_buffer(\"mean\", torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1))\n        self.register_buffer(\"std\", torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1))\n\n    def forward(self, images, mask, gamma=1.0):\n        batch, slots = images.shape[:2]\n        pixels = images.reshape(batch*slots, *images.shape[2:]).float().div_(255.0)\n        if gamma != 1.0:\n            pixels = pixels.clamp(0, 1).pow(float(gamma))\n        pixels = (pixels-self.mean)/self.std\n        tokens = self.backbone(pixel_values=pixels).last_hidden_state\n        patch = tokens[:, 1:]\n        parts = [tokens[:, 0], patch.mean(1)]\n        if self.pool == \"cls_mean_focal\":\n            k = max(1, patch.shape[1]//8)\n            parts.append(patch.topk(k, dim=1).values.mean(1))\n        features = torch.cat(parts, 1).reshape(batch, slots, -1)\n        return self.head(features, mask)\n\ndef model_spec(state, cfg):\n    patch_key = \"backbone.embeddings.patch_embeddings.projection.weight\"\n    if patch_key not in state:\n        raise RuntimeError(f\"Full DINOv2 backbone is missing from checkpoint; absent key={patch_key}\")\n    patch_weight = state[patch_key]\n    hidden_size = int(patch_weight.shape[0])\n    patch_size = int(patch_weight.shape[-1])\n    layer_ids = []\n    for key in state:\n        match = re.match(r\"backbone\\.encoder\\.layer\\.(\\d+)\\.\", key)\n        if match:\n            layer_ids.append(int(match.group(1)))\n    num_layers = max(layer_ids)+1 if layer_ids else {384: 12, 768: 12, 1024: 24}.get(hidden_size, 12)\n    fc1_key = \"backbone.encoder.layer.0.mlp.fc1.weight\"\n    intermediate_size = int(state[fc1_key].shape[0]) if fc1_key in state else hidden_size*4\n    heads = {384: 6, 768: 12, 1024: 16, 1536: 24}.get(hidden_size)\n    if heads is None:\n        raise RuntimeError(f\"Unsupported DINOv2 hidden size: {hidden_size}\")\n    position_key = \"backbone.embeddings.position_embeddings\"\n    if position_key in state:\n        n_patch = max(1, int(state[position_key].shape[1])-1)\n        pretrain_image = int(round(math.sqrt(n_patch))) * patch_size\n    else:\n        pretrain_image = 224\n\n    slot_key = \"head.slot_emb\"\n    if slot_key not in state:\n        raise RuntimeError(\"Checkpoint is missing head.slot_emb\")\n    n_slot, head_hidden = map(int, state[slot_key].shape)\n    proj_key = \"head.proj.1.weight\"\n    if proj_key not in state:\n        raise RuntimeError(\"Checkpoint is missing head.proj.1.weight\")\n    pooled_dim = int(state[proj_key].shape[1])\n    parts = int(round(pooled_dim/hidden_size))\n    if parts not in {2, 3}:\n        raise RuntimeError(f\"Unsupported pooling width: {pooled_dim}/{hidden_size}\")\n    pool = \"cls_mean_focal\" if parts == 3 else \"cls_mean\"\n    prior = \"head.slot_prior\" in state\n    return {\n        \"hidden_size\": hidden_size, \"patch_size\": patch_size, \"num_layers\": num_layers,\n        \"intermediate_size\": intermediate_size, \"num_heads\": heads, \"pretrain_image\": pretrain_image,\n        \"n_slot\": n_slot, \"head_hidden\": head_hidden, \"pool\": pool, \"prior\": prior,\n    }\n\ndef build_model(state, cfg):\n    spec = model_spec(state, cfg)\n    if spec[\"n_slot\"] != N_SLOT:\n        raise RuntimeError(f\"Checkpoint has {spec['n_slot']} slots but notebook defines {N_SLOT}\")\n    backbone_cfg = Dinov2Config(\n        image_size=spec[\"pretrain_image\"], patch_size=spec[\"patch_size\"], num_channels=3,\n        hidden_size=spec[\"hidden_size\"], num_hidden_layers=spec[\"num_layers\"],\n        num_attention_heads=spec[\"num_heads\"], intermediate_size=spec[\"intermediate_size\"],\n        hidden_dropout_prob=0.0, attention_probs_dropout_prob=0.0, drop_path_rate=0.0,\n        layerscale_value=1.0, use_mask_token=True,\n    )\n    backbone = Dinov2Model(backbone_cfg)\n    model = KneeModel(backbone, spec[\"pool\"], spec[\"n_slot\"], spec[\"head_hidden\"], spec[\"prior\"])\n    missing, unexpected = model.load_state_dict(state, strict=False)\n    critical_missing = [key for key in missing if (key.startswith(\"backbone.\") or key.startswith(\"head.\")) and not key.endswith(\"position_ids\")]\n    critical_unexpected = [key for key in unexpected if key.startswith(\"backbone.\") or key.startswith(\"head.\")]\n    if critical_missing or critical_unexpected:\n        raise RuntimeError(\n            f\"Checkpoint architecture mismatch; missing={critical_missing[:12]}, unexpected={critical_unexpected[:12]}\"\n        )\n    return model, spec","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T08:39:01.408454Z","iopub.execute_input":"2026-08-17T08:39:01.409375Z","iopub.status.idle":"2026-08-17T08:39:01.43295Z","shell.execute_reply.started":"2026-08-17T08:39:01.409342Z","shell.execute_reply":"2026-08-17T08:39:01.43198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 7. GPU inference, TTA and rank ensemble\nTTA_POOL = {\n    \"Fracture\": \"max\", \"Contusion\": \"max\", \"Medial Meniscus\": \"max\",\n    \"Lateral Meniscus\": \"max\", \"ACL\": \"top2\", \"MCL\": \"top2\", \"Baker's\": \"max\",\n}\n\ndef resolve_config(cfg):\n    cfg = cfg or {}\n    image_size = int(cfg.get(\"img\", cfg.get(\"img_size\", DEFAULT_IMG)))\n    group = int(cfg.get(\"group\", DEFAULT_GROUP))\n    n_slices = int(cfg.get(\"slices\", group*int(cfg[\"n_group\"]) if \"n_group\" in cfg else DEFAULT_SLICES))\n    n_slices = max(n_slices, group)\n    band = cfg.get(\"band\", DEFAULT_BAND)\n    if isinstance(band, str):\n        band = [item.strip() for item in band.strip(\"()[]\").split(\",\")]\n    band = tuple(float(value) for value in band)\n    if len(band) != 2 or not 0 <= band[0] <= band[1] <= 1:\n        raise ValueError(f\"Invalid band={band}\")\n    if group != 3:\n        raise ValueError(f\"The DINOv2 input requires group=3; checkpoint requested group={group}\")\n    crop_mm = float(cfg.get(\"crop_mm\", CROP_MM))\n    return image_size, n_slices, group, band, crop_mm\n\ndef window_starts(n_slices, group):\n    if n_slices <= group:\n        return [0]\n    starts = np.arange(n_slices-group+1)\n    if len(starts) <= MAX_TTA_WINDOWS:\n        return starts.tolist()\n    return sorted(set(np.rint(np.linspace(0, len(starts)-1, MAX_TTA_WINDOWS)).astype(int).tolist()))\n\ndef predict(model, cache, group, gammas):\n    loader = DataLoader(\n        CacheDataset(cache), batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=min(4, max(0, os.cpu_count() or 2)), pin_memory=True,\n        persistent_workers=False,\n    )\n    starts = window_starts(cache.n_slices, group)\n    target_index = {target: index for index, target in enumerate(TARGETS)}\n    output = []\n    model.eval()\n    with torch.inference_mode():\n        for pixels, mask in loader:\n            pixels = pixels.to(DEVICE, non_blocking=True)\n            mask = mask.to(DEVICE, non_blocking=True)\n            tta = []\n            for start in starts:\n                image = pixels[:, :, start:start+group]\n                if image.shape[2] < group:\n                    image = torch.cat([image, image[:, :, -1:].expand(-1, -1, group-image.shape[2], -1, -1)], 2)\n                for gamma in gammas:\n                    with torch.autocast(device_type=\"cuda\", dtype=torch.float16, enabled=True):\n                        tta.append(torch.sigmoid(model(image, mask, gamma=gamma).float()))\n            stack = torch.stack(tta, 0)\n            pooled = stack.mean(0)\n            for target, mode in TTA_POOL.items():\n                if target not in target_index:\n                    continue\n                column = target_index[target]\n                if mode == \"max\":\n                    pooled[:, column] = stack[:, :, column].max(0).values\n                elif mode.startswith(\"top\"):\n                    k = min(max(1, int(mode[3:])), stack.shape[0])\n                    pooled[:, column] = stack[:, :, column].topk(k, 0).values.mean(0)\n            output.append(pooled.cpu().numpy())\n    return np.concatenate(output, 0)\n\ndef ensemble_weights(cfg):\n    base = float((cfg or {}).get(\"weight\", 1.0))\n    values = (cfg or {}).get(\"target_weights\")\n    if isinstance(values, dict):\n        result = np.asarray([base*float(values.get(target, 1.0)) for target in TARGETS], np.float64)\n    elif isinstance(values, (list, tuple, np.ndarray)) and len(values) == len(TARGETS):\n        result = base*np.asarray(values, np.float64)\n    else:\n        result = np.full(len(TARGETS), base, np.float64)\n    if not np.isfinite(result).all() or np.any(result <= 0):\n        raise ValueError(f\"Invalid ensemble weights: {result}\")\n    return result\n\ngroups = defaultdict(list)\nfor member in LOADED_MEMBERS:\n    groups[resolve_config(member.config)].append(member)\n\nrank_sum = np.zeros((len(TEST_DF), len(TARGETS)), np.float64)\nweight_sum = np.zeros(len(TARGETS), np.float64)\naudit = []\n\nfor group_number, configuration in enumerate(sorted(groups), 1):\n    image_size, n_slices, group, band, crop_mm = configuration\n    members = groups[configuration]\n    log(\n        f\"configuration {group_number}/{len(groups)}: img={image_size}, slices={n_slices}, \"\n        f\"band={band}, crop={crop_mm}mm, members={len(members)}\"\n    )\n    cache = StudyCache(image_size, n_slices, band, crop_mm)\n    for member_number, member in enumerate(members, 1):\n        log(f\"member {member_number}/{len(members)}: {member.name}\")\n        model = predictions = ranks = None\n        try:\n            model, spec = build_model(member.state, member.config)\n            model = model.to(DEVICE)\n            if GPU_COUNT > 1:\n                model = nn.DataParallel(model)\n            gamma_values = member.config.get(\"gammas\", [1.0])\n            if isinstance(gamma_values, (int, float)):\n                gamma_values = [gamma_values]\n            gammas = tuple(float(value) for value in gamma_values)\n            predictions = predict(model, cache, group, gammas)\n            if predictions.shape != (len(TEST_DF), len(TARGETS)) or not np.isfinite(predictions).all():\n                raise RuntimeError(f\"Invalid prediction matrix: {predictions.shape}\")\n            weights = ensemble_weights(member.config)\n            ranks = pd.DataFrame(predictions, columns=TARGETS).rank(method=\"average\", pct=True).to_numpy(np.float64)\n            rank_sum += ranks*weights[None, :]\n            weight_sum += weights\n            audit.append({\n                \"member\": member.name, \"img\": image_size, \"slices\": n_slices,\n                \"pool\": spec[\"pool\"], \"hidden\": spec[\"hidden_size\"],\n                \"min\": float(predictions.min()), \"max\": float(predictions.max()),\n            })\n            log(f\"complete: range=[{predictions.min():.4f}, {predictions.max():.4f}]\")\n        except torch.cuda.OutOfMemoryError:\n            raise RuntimeError(\n                f\"GPU memory exhausted for {member.name}. Set BATCH_SIZE=2 in cell 1 and restart Run All.\"\n            )\n        finally:\n            model = None\n            predictions = None\n            ranks = None\n            member.state = None\n            gc.collect()\n            torch.cuda.empty_cache()\n    cache.close()\n    del cache\n    gc.collect()\n\nif not audit or np.any(weight_sum <= 0):\n    raise RuntimeError(\"No model completed inference successfully\")\n\nfinal = rank_sum/weight_sum[None, :]\naudit_df = pd.DataFrame(audit)\ndisplay(audit_df)\nlog(f\"ensemble complete: {len(audit_df)} members\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T08:39:18.725294Z","iopub.execute_input":"2026-08-17T08:39:18.726164Z","iopub.status.idle":"2026-08-17T08:39:51.86468Z","shell.execute_reply.started":"2026-08-17T08:39:18.726127Z","shell.execute_reply":"2026-08-17T08:39:51.863864Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 8. Strict submission validation\nsubmission = TEST_DF[[\"StudyInstanceUID\"]].copy()\nfor column, target in enumerate(TARGETS):\n    submission[target] = final[:, column]\n\nsubmission = SAMPLE[[\"StudyInstanceUID\"]].merge(\n    submission, on=\"StudyInstanceUID\", how=\"left\", validate=\"one_to_one\"\n)\nif list(submission.columns) != list(SAMPLE.columns):\n    raise RuntimeError(f\"Submission columns differ from sample: {submission.columns.tolist()}\")\nif len(submission) != len(SAMPLE) or not submission.StudyInstanceUID.is_unique:\n    raise RuntimeError(\"Submission row count or StudyInstanceUID uniqueness is invalid\")\n\nvalues = submission[TARGETS].apply(pd.to_numeric, errors=\"raise\").to_numpy(float)\nif not np.isfinite(values).all() or np.any(values < 0) or np.any(values > 1):\n    raise RuntimeError(\"Submission contains invalid probabilities\")\nif len(submission) > 3:\n    constant = [target for target in TARGETS if submission[target].nunique() <= 1]\n    if constant:\n        raise RuntimeError(f\"Constant predictions detected for: {constant}\")\n\nsubmission.to_csv(\"submission.csv\", index=False)\nlog(f\"SUCCESS: submission.csv {submission.shape}, range=[{values.min():.5f}, {values.max():.5f}]\")\ndisplay(submission.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T08:41:18.691436Z","iopub.execute_input":"2026-08-17T08:41:18.692261Z","iopub.status.idle":"2026-08-17T08:41:18.732901Z","shell.execute_reply.started":"2026-08-17T08:41:18.692224Z","shell.execute_reply":"2026-08-17T08:41:18.732136Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}