{"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"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# RSNA Knee — baseline training notebook. Paste into a Kaggle Notebook cell and run.\n#\n# Setup on the Kaggle side (one-time, in the notebook's right sidebar):\n#   1. Add Input -> \"rsna-knee-abnormality-detection\" (competition data) if not attached.\n#   2. Add Input -> \"kairavsapru/rsna-knee-cv-folds\" (the small CV-fold file we uploaded\n#      from the local project) if not attached.\n#   3. Settings -> Accelerator -> GPU T4 x2 (or P100) -> Save.\n#   4. Settings -> Internet -> On (training needs it to download the pretrained backbone;\n#      the SEPARATE submission notebook that scores on the leaderboard must NOT use\n#      internet, but this training notebook isn't that one).\n#\n# What this is: the v1 \"get something through the pipeline\" baseline per the project plan\n# -- geometric DICOM slice sorting, one representative series per anatomical plane\n# (Sagittal/Coronal/Axial), a 3-slice RGB triplet per plane, DINOv2-Small backbone (per\n# community testing, bigger encoders haven't paid off -- see discussion #735154), masked\n# multi-label BCE (0.5 = \"report never addressed this\", excluded from the loss, not a\n# training target), early stopping, then a held-out check against the 58 GOLD-labeled\n# studies (never used to pick the config, per discussion #734055's finding that 58 studies\n# is too few/noisy to rank models -- it's a final sanity check only).\n#\n# This is deliberately NOT the final architecture -- see module docstrings below for what's\n# simplified here vs. the fuller community approach (6 series \"slots\" instead of 3, 9\n# slices/3 anchors instead of 1 anchor). Get this working end-to-end first.\n#\n# Same logic as this project's local src/ modules (dicom_utils.py, series_selection.py,\n# dataset.py, model.py, losses.py, metrics.py, train.py), all of which are unit-tested\n# against synthetic data locally -- inlined here since Kaggle Notebooks can't import this\n# repo's local package. Only the backbone (real pretrained DINOv2-Small instead of a tiny\n# random test stand-in) and the DICOM error handling (defensive try/except around real\n# files we can't test against locally) are genuinely new versus the tested local code.\n\nimport subprocess\nimport sys\n\nsubprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"-q\",\n                 \"pylibjpeg\", \"pylibjpeg-libjpeg\", \"pylibjpeg-openjpeg\"], check=True)\n\nimport copy\nimport json\nimport random\nimport warnings\nfrom pathlib import Path\n\nimport cv2\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.amp import autocast, GradScaler\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import roc_auc_score\nfrom transformers import AutoModel\n\nwarnings.filterwarnings(\"ignore\")\n\n# ============================================================\n# Config\n# ============================================================\nCOMPETITION_DIR = Path(\"/kaggle/input/competitions/rsna-knee-abnormality-detection\")\nCV_FOLDS_PATH = Path(\"/kaggle/input/datasets/kairavsapru/rsna-knee-cv-folds/cv_folds.csv\")\nTRAIN_SERIES_DIR = COMPETITION_DIR / \"train_series\"\nOUT_DIR = Path(\"/kaggle/working\")\n\nLABEL_COLUMNS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\", \"Medial OA\", \"Lateral OA\",\n    \"PF OA\", \"Effusion\", \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\",\n]\n\nVAL_FOLD = 0          # fold 0 held out for validation; the rest (1-4) train\nCROP_MM = 160.0        # matches this dataset's median field-of-view (community testing)\nOUTPUT_SIZE = 224\nBATCH_SIZE = 16\nEPOCHS = 10  # lowered from 15 as a safety margin against a real ~7hr GPU-quota budget\nPATIENCE = 4\n# Separate learning rates for the pretrained backbone vs. the freshly-initialized head.\n# First real run showed val AUC peaking at epoch 2 then declining -- the classic signature\n# of a random head's large early gradients dragging an already-good pretrained backbone\n# off course when both train at the same rate. HEAD_LR stays fast since it's learning from\n# scratch; BACKBONE_LR is 100x gentler so DINOv2's pretrained features shift slowly instead\n# of getting scrambled before the head has learned anything useful.\nHEAD_LR = 1e-3\nBACKBONE_LR = 1e-5\nBACKBONE_NAME = \"facebook/dinov2-small\"\nSEED = 42\n\nDEVICE = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", DEVICE)\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\n\nset_seed(SEED)\n\n# ============================================================\n# DICOM geometry: slice sorting, anchor sampling, physical-scale crop\n# (identical logic to src/dicom_utils.py -- see that file for the \"why\" on each choice)\n# ============================================================\nTIER_RANK = {\"geometry\": 0, \"slice_location\": 1, \"instance_number\": 2, \"filename\": 3}\n\n\ndef compute_slice_position(ds):\n    iop = ds.ImageOrientationPatient\n    ipp = ds.ImagePositionPatient\n    row = np.array(iop[0:3], dtype=float)\n    col = np.array(iop[3:6], dtype=float)\n    normal = np.cross(row, col)\n    return float(np.dot(normal, np.array(ipp, dtype=float)))\n\n\ndef _tier_for(ds):\n    if hasattr(ds, \"ImageOrientationPatient\") and hasattr(ds, \"ImagePositionPatient\"):\n        try:\n            compute_slice_position(ds)\n            return \"geometry\"\n        except Exception:\n            pass\n    if hasattr(ds, \"SliceLocation\"):\n        return \"slice_location\"\n    if hasattr(ds, \"InstanceNumber\"):\n        return \"instance_number\"\n    return \"filename\"\n\n\ndef sort_series_files(dicom_paths):\n    datasets = [(p, pydicom.dcmread(p, stop_before_pixels=True)) for p in dicom_paths]\n    tiers = [_tier_for(ds) for _, ds in datasets]\n    common_tier = max(tiers, key=lambda t: TIER_RANK[t])\n\n    def key_for(path, ds):\n        if common_tier == \"geometry\":\n            return compute_slice_position(ds)\n        elif common_tier == \"slice_location\":\n            return float(ds.SliceLocation)\n        elif common_tier == \"instance_number\":\n            return int(ds.InstanceNumber)\n        return str(path)\n\n    keyed = [(key_for(p, ds), p) for p, ds in datasets]\n    keyed.sort(key=lambda x: x[0])\n    return [p for _, p in keyed]\n\n\ndef select_anchor_triplets(n_slices, n_anchors=1, edge_clip_frac=0.15):\n    if n_slices < 3:\n        raise ValueError(f\"need >=3 slices, got {n_slices}\")\n    clip = int(round(n_slices * edge_clip_frac))\n    lo, hi = clip, n_slices - 1 - clip\n    if hi <= lo:\n        lo, hi = 0, n_slices - 1\n    anchors = np.round(np.linspace(lo, hi, n_anchors)).astype(int)\n    triplets = []\n    for a in anchors:\n        start = max(0, min(int(a) - 1, n_slices - 3))\n        triplets.append([start, start + 1, start + 2])\n    return triplets\n\n\ndef read_pixel_array(path):\n    ds = pydicom.dcmread(path)\n    arr = ds.pixel_array.astype(np.float32)\n    slope = float(getattr(ds, \"RescaleSlope\", 1.0))\n    intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n    return arr * slope + intercept\n\n\ndef normalize_intensity(image, low_pct=0.5, high_pct=99.5):\n    lo, hi = np.percentile(image, [low_pct, high_pct])\n    if hi <= lo:\n        return np.zeros_like(image, dtype=np.uint8)\n    clipped = np.clip(image, lo, hi)\n    normalized = (clipped - lo) / (hi - lo)\n    return (normalized * 255).astype(np.uint8)\n\n\ndef physical_crop_and_resize(image, pixel_spacing, crop_mm, output_size):\n    h, w = image.shape[:2]\n    row_spacing, col_spacing = pixel_spacing\n    crop_h_px = max(1, int(round(crop_mm / row_spacing)))\n    crop_w_px = max(1, int(round(crop_mm / col_spacing)))\n    center_y, center_x = h // 2, w // 2\n    y0 = max(0, center_y - crop_h_px // 2)\n    y1 = min(h, y0 + crop_h_px)\n    x0 = max(0, center_x - crop_w_px // 2)\n    x1 = min(w, x0 + crop_w_px)\n    cropped = image[y0:y1, x0:x1]\n    pad_h = crop_h_px - cropped.shape[0]\n    pad_w = crop_w_px - cropped.shape[1]\n    if pad_h > 0 or pad_w > 0:\n        cropped = np.pad(cropped, ((pad_h // 2, pad_h - pad_h // 2),\n                                    (pad_w // 2, pad_w - pad_w // 2)), mode=\"constant\")\n    return cv2.resize(cropped, (output_size, output_size), interpolation=cv2.INTER_AREA)\n\n\n# ============================================================\n# Series selection: one representative series per anatomical plane\n# ============================================================\nPLANES = [\"Sagittal\", \"Coronal\", \"Axial\"]\n\n\ndef select_representative_series(study_series: pd.DataFrame) -> dict:\n    selection = {}\n    for plane in PLANES:\n        candidates = study_series[study_series[\"Anatomical_Plane\"] == plane]\n        if len(candidates) == 0:\n            selection[plane] = None\n            continue\n        fluid_sensitive = candidates[candidates[\"Fluid_Sensitive\"] == 1]\n        chosen = fluid_sensitive.iloc[0] if len(fluid_sensitive) > 0 else candidates.iloc[0]\n        selection[plane] = chosen[\"SeriesInstanceUID\"]\n    return selection\n\n\n# ============================================================\n# Dataset\n# ============================================================\nclass KneeStudyDataset(Dataset):\n    def __init__(self, studies_df, series_df, series_dir_fn, label_columns,\n                 crop_mm=CROP_MM, output_size=OUTPUT_SIZE, cache_dir=None):\n        self.studies = studies_df.reset_index(drop=True)\n        self.series_df = series_df\n        self.series_dir_fn = series_dir_fn\n        self.label_columns = label_columns\n        self.crop_mm = crop_mm\n        self.output_size = output_size\n        # DataLoader workers (num_workers>0) run in separate processes and get fresh,\n        # re-spawned every epoch by default -- an in-memory cache on self wouldn't survive\n        # between epochs. A DISK cache does: every study's decoded/cropped image is\n        # computed from raw DICOM once, on whichever epoch first needs it, then every\n        # later epoch (and every other worker process) just loads a small saved .npy\n        # instead of re-decoding DICOM. Without this, N epochs means paying the full\n        # (expensive) DICOM decode cost N times over.\n        self.cache_dir = Path(cache_dir) if cache_dir else None\n        if self.cache_dir:\n            self.cache_dir.mkdir(parents=True, exist_ok=True)\n\n    def __len__(self):\n        return len(self.studies)\n\n    def _load_plane_image(self, study_id, series_id):\n        empty = (np.zeros((self.output_size, self.output_size, 3), dtype=np.uint8), 0.0)\n        if series_id is None:\n            return empty\n\n        cache_path = self.cache_dir / f\"{study_id}_{series_id}.npy\" if self.cache_dir else None\n        if cache_path is not None and cache_path.exists():\n            try:\n                return np.load(cache_path), 1.0\n            except Exception:\n                pass  # corrupt/partial cache file (e.g. from an interrupted write) -- recompute\n\n        try:\n            series_dir = self.series_dir_fn(study_id, series_id)\n            dicom_paths = list(Path(series_dir).glob(\"*.dcm\"))\n            if len(dicom_paths) < 3:\n                return empty\n\n            sorted_paths = sort_series_files(dicom_paths)\n            triplet = select_anchor_triplets(len(sorted_paths), n_anchors=1)[0]\n\n            ref_ds = pydicom.dcmread(sorted_paths[triplet[0]], stop_before_pixels=True)\n            pixel_spacing = tuple(float(x) for x in getattr(ref_ds, \"PixelSpacing\", (1.0, 1.0)))\n\n            channels = []\n            for idx in triplet:\n                raw = read_pixel_array(sorted_paths[idx])\n                norm = normalize_intensity(raw)\n                cropped = physical_crop_and_resize(norm, pixel_spacing, self.crop_mm,\n                                                    self.output_size)\n                channels.append(cropped)\n            image = np.stack(channels, axis=-1)\n\n            if cache_path is not None:\n                # write to a temp name then rename -- atomic on the same filesystem, so a\n                # worker process killed mid-write can never leave behind a half-written\n                # file that a later epoch would try (and fail) to load. np.save() always\n                # appends \".npy\" to whatever name it's given, so the temp name must\n                # already end in \".npy\" or this ends up pointing at the wrong file.\n                tmp_path = cache_path.with_name(cache_path.name + \".tmp.npy\")\n                np.save(tmp_path, image)\n                tmp_path.replace(cache_path)\n\n            return image, 1.0\n        except Exception as e:\n            # Real DICOM data has edge cases synthetic test data can't cover (corrupt\n            # files, unsupported transfer syntaxes, missing tags) -- one bad series\n            # should degrade to \"missing\" for that plane, not crash the whole run.\n            print(f\"  [warn] failed to load {study_id}/{series_id}: {e}\")\n            return empty\n\n    def __getitem__(self, idx):\n        row = self.studies.iloc[idx]\n        study_id = row[\"StudyInstanceUID\"]\n        study_series = self.series_df[self.series_df[\"StudyInstanceUID\"] == study_id]\n        selection = select_representative_series(study_series)\n\n        images, mask = [], []\n        for plane in PLANES:\n            image, present = self._load_plane_image(study_id, selection[plane])\n            images.append(image)\n            mask.append(present)\n\n        stacked = np.stack(images, axis=0).astype(np.float32) / 255.0\n        stacked = np.transpose(stacked, (0, 3, 1, 2))  # (3 planes, 3ch, H, W)\n        labels = row[self.label_columns].to_numpy(dtype=np.float32)\n\n        return {\n            \"image\": torch.from_numpy(stacked),\n            \"mask\": torch.tensor(mask, dtype=torch.float32),\n            \"label\": torch.from_numpy(labels),\n            \"study_id\": study_id,\n        }\n\n\ndef train_series_dir_fn(study_id, series_id):\n    return TRAIN_SERIES_DIR / study_id / series_id\n\n\n# ============================================================\n# Model: DINOv2-Small backbone (real pretrained weights) + masked-mean-pool + linear head\n# ============================================================\nclass Dinov2Backbone(nn.Module):\n    \"\"\"Wraps HF's DINOv2 to the (N,3,H,W) -> (N,embed_dim) signature the pooling/head code\n    expects. Uses the CLS token (index 0 of last_hidden_state) as the per-image embedding,\n    DINOv2's standard usage for downstream classification.\"\"\"\n    def __init__(self, model_name=BACKBONE_NAME):\n        super().__init__()\n        self.model = AutoModel.from_pretrained(model_name)\n        self.embed_dim = self.model.config.hidden_size\n        # DINOv2 was pretrained on ImageNet-normalized RGB; our images are already [0,1]\n        # single-triplet-as-RGB, so we still apply ImageNet mean/std to match its\n        # pretraining distribution (see discussion #735154's RAD-DINO postmortem on why a\n        # mismatched preprocessing contract silently produces a plausible, wrong result).\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, x):\n        x = (x - self.mean) / self.std\n        out = self.model(pixel_values=x)\n        return out.last_hidden_state[:, 0, :]  # CLS token\n\n\nclass KneeMultiPlaneModel(nn.Module):\n    def __init__(self, backbone, embed_dim, num_labels=len(LABEL_COLUMNS)):\n        super().__init__()\n        self.backbone = backbone\n        self.head = nn.Linear(embed_dim, num_labels)\n\n    def forward(self, images, mask):\n        B, P, C, H, W = images.shape\n        flat = images.view(B * P, C, H, W)\n        feats = self.backbone(flat).view(B, P, -1)\n        mask_expanded = mask.unsqueeze(-1)\n        summed = (feats * mask_expanded).sum(dim=1)\n        count = mask_expanded.sum(dim=1).clamp(min=1.0)\n        pooled = summed / count\n        return self.head(pooled)\n\n\ndef masked_bce_loss(logits, targets, ignore_value=0.5, eps=1e-6):\n    mask = (targets - ignore_value).abs() > eps\n    if mask.sum() == 0:\n        return logits.sum() * 0.0\n    per_cell = F.binary_cross_entropy_with_logits(logits, targets, reduction=\"none\")\n    return (per_cell * mask).sum() / mask.sum()\n\n\ndef masked_macro_auc(y_true, y_pred, label_columns, ignore_value=0.5, eps=1e-6):\n    per_label = {}\n    for i, col in enumerate(label_columns):\n        col_true, col_pred = y_true[:, i], y_pred[:, i]\n        known = np.abs(col_true - ignore_value) > eps\n        if known.sum() == 0:\n            per_label[col] = None\n            continue\n        binary_true = (col_true[known] > 0.5).astype(int)\n        if len(np.unique(binary_true)) < 2:\n            per_label[col] = None\n            continue\n        per_label[col] = float(roc_auc_score(binary_true, col_pred[known]))\n    valid = [v for v in per_label.values() if v is not None]\n    macro_auc = float(np.mean(valid)) if valid else None\n    return macro_auc, per_label\n\n\n# ============================================================\n# Load labels + build train/val split from the pre-computed CV folds\n# ============================================================\ncv_folds = pd.read_csv(CV_FOLDS_PATH)\ntrain_series_csv = pd.read_csv(COMPETITION_DIR / \"train_series.csv\")\n\ntrain_df = cv_folds[(cv_folds[\"fold\"] != VAL_FOLD) & (cv_folds[\"fold\"] >= 0)].reset_index(drop=True)\nval_df = cv_folds[cv_folds[\"fold\"] == VAL_FOLD].reset_index(drop=True)\ngold_df = cv_folds[cv_folds[\"label_source\"] == \"gold\"].reset_index(drop=True)\n\nprint(f\"train: {len(train_df)} | val (fold {VAL_FOLD}): {len(val_df)} | \"\n      f\"gold sanity-check: {len(gold_df)}\")\n\nCACHE_DIR = OUT_DIR / \"plane_cache\"\ntrain_ds = KneeStudyDataset(train_df, train_series_csv, train_series_dir_fn, LABEL_COLUMNS,\n                             cache_dir=CACHE_DIR)\nval_ds = KneeStudyDataset(val_df, train_series_csv, train_series_dir_fn, LABEL_COLUMNS,\n                           cache_dir=CACHE_DIR)\ngold_ds = KneeStudyDataset(gold_df, train_series_csv, train_series_dir_fn, LABEL_COLUMNS,\n                            cache_dir=CACHE_DIR)\n\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True,\n                           num_workers=4, pin_memory=True, drop_last=True)\nval_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False,\n                         num_workers=4, pin_memory=True)\ngold_loader = DataLoader(gold_ds, batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=2, pin_memory=True)\n\n# ============================================================\n# Train\n# ============================================================\nbackbone = Dinov2Backbone(BACKBONE_NAME)\nmodel = KneeMultiPlaneModel(backbone, embed_dim=backbone.embed_dim).to(DEVICE)\n\noptimizer = torch.optim.AdamW([\n    {\"params\": model.backbone.parameters(), \"lr\": BACKBONE_LR},\n    {\"params\": model.head.parameters(), \"lr\": HEAD_LR},\n], weight_decay=1e-4)\nscaler = GradScaler(enabled=(DEVICE.type == \"cuda\"))\n\n\ndef run_epoch(loader, train_mode):\n    model.train() if train_mode else model.eval()\n    total_loss, n_batches = 0.0, 0\n    all_preds, all_labels = [], []\n    context = torch.enable_grad() if train_mode else torch.no_grad()\n    with context:\n        for batch in loader:\n            images = batch[\"image\"].to(DEVICE, non_blocking=True)\n            mask = batch[\"mask\"].to(DEVICE, non_blocking=True)\n            labels = batch[\"label\"].to(DEVICE, non_blocking=True)\n\n            if train_mode:\n                optimizer.zero_grad()\n            with autocast(device_type=DEVICE.type, enabled=(DEVICE.type == \"cuda\")):\n                logits = model(images, mask)\n                loss = masked_bce_loss(logits, labels)\n\n            if train_mode:\n                scaler.scale(loss).backward()\n                scaler.step(optimizer)\n                scaler.update()\n\n            total_loss += loss.item()\n            n_batches += 1\n            all_preds.append(torch.sigmoid(logits).detach().float().cpu().numpy())\n            all_labels.append(labels.detach().cpu().numpy())\n\n    y_pred = np.concatenate(all_preds, axis=0)\n    y_true = np.concatenate(all_labels, axis=0)\n    macro_auc, per_label_auc = masked_macro_auc(y_true, y_pred, LABEL_COLUMNS)\n    return total_loss / max(n_batches, 1), macro_auc, per_label_auc\n\n\nbest_auc, best_state, counter = -1.0, None, 0\nfor epoch in range(EPOCHS):\n    train_loss, train_auc, _ = run_epoch(train_loader, train_mode=True)\n    val_loss, val_auc, val_per_label = run_epoch(val_loader, train_mode=False)\n\n    val_auc_str = f\"{val_auc:.4f}\" if val_auc is not None else \"N/A\"\n    print(f\"Epoch {epoch + 1:02d} | train_loss={train_loss:.4f} | val_loss={val_loss:.4f} \"\n          f\"| val_macro_auc={val_auc_str}\")\n\n    if val_auc is not None and val_auc > best_auc:\n        best_auc, best_state, counter = val_auc, copy.deepcopy(model.state_dict()), 0\n        # Save to disk the moment we have a new best, not just after the whole loop ends\n        # -- an overnight run that gets cut off by a timeout/quota limit should still\n        # leave behind whatever the best checkpoint was at that point, not nothing.\n        torch.save(best_state, OUT_DIR / \"model.pt\")\n        with open(OUT_DIR / \"model_config.json\", \"w\") as f:\n            json.dump({\n                \"backbone_name\": BACKBONE_NAME, \"embed_dim\": backbone.embed_dim,\n                \"label_columns\": LABEL_COLUMNS, \"crop_mm\": CROP_MM,\n                \"output_size\": OUTPUT_SIZE, \"val_fold\": VAL_FOLD,\n                \"val_macro_auc\": best_auc, \"gold_macro_auc\": None,\n                \"epoch\": epoch + 1, \"status\": \"in_progress\",\n            }, f, indent=2)\n        print(f\"  New best (val_macro_auc={best_auc:.4f}) -- saved checkpoint to \"\n              f\"{OUT_DIR / 'model.pt'}\")\n    else:\n        counter += 1\n        if counter >= PATIENCE:\n            print(f\"Early stopping at epoch {epoch + 1}.\")\n            break\n\nif best_state is not None:\n    model.load_state_dict(best_state)\nprint(f\"\\nBest val_macro_auc (fold {VAL_FOLD}, weak labels): {best_auc:.4f}\")\n\n# ============================================================\n# Final sanity check against the 58 GOLD studies (never used to pick this config)\n# ============================================================\n_, gold_auc, gold_per_label = run_epoch(gold_loader, train_mode=False)\ngold_auc_str = f\"{gold_auc:.4f}\" if gold_auc is not None else \"N/A\"\nprint(f\"\\nGold-check macro AUC (58 held-out expert-labeled studies): {gold_auc_str}\")\nprint(\"Per-label gold AUC:\")\nfor k, v in gold_per_label.items():\n    print(f\"  {k:<20s} {v:.4f}\" if v is not None else f\"  {k:<20s} N/A\")\n\n# ============================================================\n# Save weights + a small metadata file for the (separate) offline submission notebook\n# ============================================================\ntorch.save(model.state_dict(), OUT_DIR / \"model.pt\")\nwith open(OUT_DIR / \"model_config.json\", \"w\") as f:\n    json.dump({\n        \"backbone_name\": BACKBONE_NAME,\n        \"embed_dim\": backbone.embed_dim,\n        \"label_columns\": LABEL_COLUMNS,\n        \"crop_mm\": CROP_MM,\n        \"output_size\": OUTPUT_SIZE,\n        \"val_fold\": VAL_FOLD,\n        \"val_macro_auc\": best_auc,\n        \"gold_macro_auc\": gold_auc,\n        \"status\": \"done\",\n    }, f, indent=2)\n\nprint(f\"\\nSaved model.pt and model_config.json to {OUT_DIR}\")\nprint(\"Now: Save Version, then either download these two files or (better) create a \"\n      \"Kaggle Model/Dataset from this notebook's output so the submission notebook can \"\n      \"attach them as an input without re-downloading DINOv2 weights from the internet \"\n      \"(which the scored submission run isn't allowed to do).\")","metadata":{"_uuid":"7674bbf1-d763-4079-a3f1-d3686a9980e9","_cell_guid":"8d8c11bd-029b-4f35-95f2-3ca0c5135fbd","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-08-23T21:47:36.10628Z","iopub.execute_input":"2026-08-23T21:47:36.106722Z"}},"outputs":[],"execution_count":null}]}