{"cells":[{"cell_type":"markdown","id":"e246f2bd","metadata":{},"source":"# Knee MRI Study Classifier — Two-Level MIL Training\n\nTrains a study-level classifier for the 12 knee-MRI findings (ACL, MCL,\nmedial/lateral meniscus, medial/lateral/PF osteoarthritis, effusion,\nsynovitis, Baker's cyst, contusion, fracture) directly from the DICOM\nimaging data.\n\nEach study is a variable-length bag of MRI series (different planes and\nsequences), and each series is itself a bag of slices. The model reflects\nthat structure directly: a shared 2D CNN backbone embeds sampled slices,\nan attention layer pools slices into a series embedding, and a second\nattention layer pools the variable number of series into a study\nembedding feeding a 12-way sigmoid head. Both attention layers are\nlearned end-to-end together with the backbone, rather than treating the\nbackbone as a fixed feature extractor.\n\nTraining targets come from a per-study soft-label table (a probability\nper finding, not a hard yes/no) covering the full training set - a small\nground-truth-labeled subset plus a much larger set of studies labeled by\nmining their radiology reports. Only the ground-truth subset is used for\nheld-out validation, via stratified k-fold cross-validation.\n"},{"cell_type":"code","execution_count":null,"id":"1ce971f4","metadata":{},"outputs":[],"source":"import os\n\n# Must be set before torch initializes its CUDA allocator.\nos.environ.setdefault(\"PYTORCH_ALLOC_CONF\", \"expandable_segments:True\")\nos.environ.setdefault(\"PYTORCH_CUDA_ALLOC_CONF\", \"expandable_segments:True\")\n\nimport gc\nimport glob\nimport pickle\nimport subprocess\nimport sys\nimport time\nfrom concurrent.futures import ThreadPoolExecutor\n\nimport numpy as np\nimport pandas as pd\nimport torch\n\nsubprocess.run(\n    [sys.executable, \"-m\", \"pip\", \"install\", \"-q\",\n     \"pylibjpeg\", \"pylibjpeg-libjpeg\", \"pylibjpeg-openjpeg\", \"python-gdcm\",\n     \"iterative-stratification\"],\n    check=False,\n)\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nif DEVICE != \"cuda\":\n    raise RuntimeError(\"No GPU detected - enable a GPU accelerator for this notebook.\")\n\n# The competition dataset and the silver-labels Dataset are both mounted\n# somewhere under /kaggle/input, but not at predictable flat paths -\n# search for whichever directory holds each expected file.\nKAGGLE_INPUT_ROOT = \"/kaggle/input\"\nprint(f\"{KAGGLE_INPUT_ROOT} top-level contents: {sorted(os.listdir(KAGGLE_INPUT_ROOT))}\")\n\nDATA_DIR = None\nSILVER_DIR = None\nPREBUILT_MANIFEST_DIR = None\nfor dirpath, dirnames, filenames in os.walk(KAGGLE_INPUT_ROOT):\n    if DATA_DIR is None and \"train.csv\" in filenames and \"train_series.csv\" in filenames:\n        DATA_DIR = dirpath\n    if SILVER_DIR is None and \"silver_labels.csv\" in filenames:\n        SILVER_DIR = dirpath\n    if PREBUILT_MANIFEST_DIR is None and \"series_manifest.pkl\" in filenames:\n        PREBUILT_MANIFEST_DIR = dirpath\n    if DATA_DIR is not None and SILVER_DIR is not None and PREBUILT_MANIFEST_DIR is not None:\n        break\n\nif DATA_DIR is None:\n    raise FileNotFoundError(f\"No train.csv/train_series.csv pair found under {KAGGLE_INPUT_ROOT}.\")\nif SILVER_DIR is None:\n    raise FileNotFoundError(\n        f\"No silver_labels.csv found under {KAGGLE_INPUT_ROOT}. \"\n        \"Make sure the knee-mri-silver-labels Dataset is attached to this notebook.\"\n    )\n# PREBUILT_MANIFEST_DIR is optional - falls back to building it fresh if absent.\n\nWORK_DIR = \"/kaggle/working\"\nos.makedirs(WORK_DIR, exist_ok=True)\n\nLABEL_COLS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\", \"Medial OA\",\n    \"Lateral OA\", \"PF OA\", \"Effusion\", \"Synovitis\", \"Baker\\'s\", \"Contusion\", \"Fracture\",\n]\n\nprint(f\"DATA_DIR={DATA_DIR}\")\nprint(f\"SILVER_DIR={SILVER_DIR}\")\nprint(f\"WORK_DIR={WORK_DIR}\")\n"},{"cell_type":"markdown","id":"175d08ea","metadata":{},"source":"## 1. Load labels and series metadata\n\nMerges the silver-label table (soft probability per finding, every\ntraining study) with the series-level metadata (which series belong to\nwhich study, and their plane/sequence-type flags).\n"},{"cell_type":"code","execution_count":null,"id":"67845875","metadata":{},"outputs":[],"source":"silver_labels = pd.read_csv(os.path.join(SILVER_DIR, \"silver_labels.csv\"))\ntrain_series = pd.read_csv(os.path.join(DATA_DIR, \"train_series.csv\"))\n\nprint(f\"silver_labels: {silver_labels.shape}, sources: {dict(silver_labels['source'].value_counts())}\")\nprint(f\"train_series: {train_series.shape}\")\n\n# One binary \"fluid-sensitive/fat-sat\" bit rather than two columns - the\n# EDA found Fluid_Sensitive and Fat_Suppression are always identical.\nseries_by_study = {}\nfor row in train_series.itertuples():\n    series_by_study.setdefault(row.StudyInstanceUID, []).append(\n        (row.SeriesInstanceUID, row.Anatomical_Plane, int(row.Fluid_Sensitive))\n    )\n\nplanes = sorted(train_series[\"Anatomical_Plane\"].dropna().unique().tolist())\nPLANE_VOCAB = {plane: i for i, plane in enumerate(planes)}\nPLANE_VOCAB[\"<unk>\"] = len(PLANE_VOCAB)\nN_PLANES = len(PLANE_VOCAB)\nprint(f\"plane vocab: {PLANE_VOCAB}\")\n\nsilver_labels = silver_labels.set_index(\"StudyInstanceUID\")\ngt_study_uids = silver_labels.index[silver_labels[\"source\"] == \"ground_truth\"].tolist()\nsilver_only_uids = silver_labels.index[silver_labels[\"source\"] != \"ground_truth\"].tolist()\nprint(f\"ground-truth studies: {len(gt_study_uids)}, silver-labeled studies: {len(silver_only_uids)}\")\n"},{"cell_type":"markdown","id":"5294c057","metadata":{},"source":"## 2. DICOM series manifest\n\nBuilding the sorted (by acquisition order) file list for a series\nrequires reading a header from every one of its DICOM files\n(`stop_before_pixels=True`, no pixel data touched here). Doing this for\nall ~24k series up front is a real cost (I/O-bound, parallelized with\nthreads), so it is deferred until after the smoke test below - the\nsmoke test instead builds a manifest scoped to just its own few studies,\nso a broken pipeline fails in seconds/minutes rather than after paying\nthe full manifest-building cost every debug iteration. The full\nmanifest is only built once, right before the dry run needs full\ncoverage, and is cached to disk in case this notebook session reuses it.\n"},{"cell_type":"code","execution_count":null,"id":"188d991d","metadata":{},"outputs":[],"source":"import pydicom\n\nMANIFEST_PATH = os.path.join(WORK_DIR, \"series_manifest.pkl\")\n\n\ndef _list_one_series(row):\n    pattern = os.path.join(DATA_DIR, \"train_series\", row.StudyInstanceUID, row.SeriesInstanceUID, \"*.dcm\")\n    fpaths = glob.glob(pattern)\n    entries = []\n    for fpath in fpaths:\n        try:\n            ds = pydicom.dcmread(fpath, stop_before_pixels=True)\n            inst_num = int(getattr(ds, \"InstanceNumber\", 0))\n        except Exception:\n            inst_num = 0\n        entries.append((inst_num, fpath))\n    entries.sort(key=lambda t: t[0])\n    return row.SeriesInstanceUID, [fpath for _, fpath in entries]\n\n\ndef build_series_manifest(study_uids=None, cache_path=None, max_workers=32):\n    if cache_path and os.path.exists(cache_path):\n        with open(cache_path, \"rb\") as f:\n            return pickle.load(f)\n\n    rows = list(train_series.itertuples())\n    if study_uids is not None:\n        study_uid_set = set(study_uids)\n        rows = [row for row in rows if row.StudyInstanceUID in study_uid_set]\n\n    t0 = time.time()\n    manifest = {}\n    with ThreadPoolExecutor(max_workers=max_workers) as executor:\n        for series_uid, fpaths in executor.map(_list_one_series, rows):\n            manifest[series_uid] = fpaths\n    print(f\"built manifest for {len(manifest)} series in {time.time() - t0:.0f}s\")\n\n    if cache_path:\n        with open(cache_path, \"wb\") as f:\n            pickle.dump(manifest, f)\n    return manifest\n\n\n# Populated for real just before the smoke test (small scope) and again\n# before the dry run (full scope) - sample_series_slices() below always\n# reads whatever this currently points to.\nSERIES_MANIFEST = {}\n"},{"cell_type":"markdown","id":"4a0f32a7","metadata":{},"source":"## 3. Slice sampling and preprocessing\n\nA fixed, small number of slices per series (rather than the full ~30\navailable) keeps memory and per-step compute tractable, since the\nbackbone is fine-tuned end-to-end - every sampled slice costs a full\nbackward pass, not just a cheap cached forward pass. Pixel values are\nrescaled with the DICOM slope/intercept, percentile-clipped and\nmin-max normalized (MRI has no fixed intensity scale like CT Hounsfield\nunits), replicated to 3 channels, resized, and ImageNet-normalized for\nthe pretrained backbone.\n"},{"cell_type":"code","execution_count":null,"id":"6e56a4d6","metadata":{},"outputs":[],"source":"import torch.nn.functional as F\n\nIMG_SIZE = 224\nN_SLICES_PER_SERIES = 5\n\nIMAGENET_MEAN = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)\nIMAGENET_STD = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)\n\n\ndef load_and_preprocess_slice(fpath):\n    ds = pydicom.dcmread(fpath)\n    arr = ds.pixel_array.astype(\"float32\")\n    slope = float(getattr(ds, \"RescaleSlope\", 1.0))\n    intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n    arr = arr * slope + intercept\n\n    lo, hi = np.percentile(arr, [1, 99])\n    if hi > lo:\n        arr = np.clip(arr, lo, hi)\n        arr = (arr - lo) / (hi - lo)\n    else:\n        arr = np.zeros_like(arr)\n    # np.percentile returns float64, which silently upcasts arr above -\n    # force back to float32 so this doesn't fight autocast's fp16 casting.\n    arr = arr.astype(\"float32\")\n\n    tensor = torch.from_numpy(arr).unsqueeze(0).unsqueeze(0)\n    tensor = F.interpolate(tensor, size=(IMG_SIZE, IMG_SIZE), mode=\"bilinear\", align_corners=False)\n    tensor = tensor.squeeze(0).repeat(3, 1, 1)\n    tensor = (tensor - IMAGENET_MEAN) / IMAGENET_STD\n    return tensor\n\n\ndef sample_series_slices(series_uid, n_slices=N_SLICES_PER_SERIES):\n    fpaths = SERIES_MANIFEST.get(series_uid, [])\n    if not fpaths:\n        return None\n    n = min(n_slices, len(fpaths))\n    idxs = sorted(set(np.linspace(0, len(fpaths) - 1, n).round().astype(int).tolist()))\n    tensors = []\n    for idx in idxs:\n        try:\n            tensors.append(load_and_preprocess_slice(fpaths[idx]))\n        except Exception:\n            continue\n    if not tensors:\n        return None\n    return torch.stack(tensors)\n"},{"cell_type":"markdown","id":"924f6ac0","metadata":{},"source":"## 4. Dataset\n\nOne item is one study: a ragged list of per-series slice tensors (series\ncount varies 3-14 per the EDA) plus the plane/fluid-bit for each series\nand the soft-label target. Batch size is fixed at 1 study, with gradient\naccumulation at the training-loop level standing in for a larger\neffective batch - avoids padding/masking logic for the ragged bag\nstructure.\n"},{"cell_type":"code","execution_count":null,"id":"6c1b846e","metadata":{},"outputs":[],"source":"from torch.utils.data import Dataset\n\n\nclass StudyDataset(Dataset):\n    def __init__(self, study_uids):\n        self.study_uids = [uid for uid in study_uids if uid in series_by_study]\n\n    def __len__(self):\n        return len(self.study_uids)\n\n    def __getitem__(self, idx):\n        study_uid = self.study_uids[idx]\n        series_tensors, plane_ids, fluid_ids = [], [], []\n        for series_uid, plane, fluid_bit in series_by_study.get(study_uid, []):\n            slices = sample_series_slices(series_uid)\n            if slices is None:\n                continue\n            series_tensors.append(slices)\n            plane_ids.append(PLANE_VOCAB.get(plane, PLANE_VOCAB[\"<unk>\"]))\n            fluid_ids.append(fluid_bit)\n\n        target = torch.tensor(silver_labels.loc[study_uid, LABEL_COLS].values.astype(\"float32\"))\n        return {\n            \"series_tensors\": series_tensors,\n            \"plane_ids\": torch.tensor(plane_ids, dtype=torch.long),\n            \"fluid_ids\": torch.tensor(fluid_ids, dtype=torch.long),\n            \"target\": target,\n            \"study_uid\": study_uid,\n        }\n\n\ndef collate_single(batch):\n    return batch[0]\n"},{"cell_type":"markdown","id":"55537aad","metadata":{},"source":"## 5. Model — two-level gated-attention MIL\n\nGated attention pooling (Ilse et al. 2018) at both levels: slices within\na series, then series within a study. Plane and fluid-sensitivity are\nconcatenated into the series embedding as small learned embeddings\nrather than architected as separate branches.\n"},{"cell_type":"code","execution_count":null,"id":"0d09cf6a","metadata":{},"outputs":[],"source":"import torch.nn as nn\nimport torchvision\n\n\nclass GatedAttentionPool(nn.Module):\n    def __init__(self, dim, hidden=128):\n        super().__init__()\n        self.V = nn.Linear(dim, hidden)\n        self.U = nn.Linear(dim, hidden)\n        self.w = nn.Linear(hidden, 1)\n\n    def forward(self, x):\n        a = torch.tanh(self.V(x)) * torch.sigmoid(self.U(x))\n        weights = torch.softmax(self.w(a), dim=0)\n        return (weights * x).sum(dim=0)\n\n\nclass TwoLevelMIL(nn.Module):\n    def __init__(self, n_planes, embed_dim=256, plane_dim=16, fluid_dim=8,\n                 pool_hidden=128, head_hidden=128, n_labels=12):\n        super().__init__()\n        backbone = torchvision.models.efficientnet_b0(\n            weights=torchvision.models.EfficientNet_B0_Weights.IMAGENET1K_V1\n        )\n        backbone.classifier = nn.Identity()\n        self.backbone = backbone\n        self.proj = nn.Linear(1280, embed_dim)\n        self.slice_pool = GatedAttentionPool(embed_dim, pool_hidden)\n        self.plane_embed = nn.Embedding(n_planes, plane_dim)\n        self.fluid_embed = nn.Embedding(2, fluid_dim)\n        series_dim = embed_dim + plane_dim + fluid_dim\n        self.study_pool = GatedAttentionPool(series_dim, pool_hidden)\n        self.head = nn.Sequential(\n            nn.Linear(series_dim, head_hidden),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(head_hidden, n_labels),\n        )\n\n    def forward(self, series_tensors, plane_ids, fluid_ids):\n        series_embeds = []\n        for i, slices in enumerate(series_tensors):\n            slice_feats = self.proj(self.backbone(slices))\n            series_embed = self.slice_pool(slice_feats)\n            p_emb = self.plane_embed(plane_ids[i])\n            f_emb = self.fluid_embed(fluid_ids[i])\n            series_embeds.append(torch.cat([series_embed, p_emb, f_emb]))\n        series_embeds = torch.stack(series_embeds)\n        study_embed = self.study_pool(series_embeds)\n        return self.head(study_embed)\n"},{"cell_type":"markdown","id":"5d4759ed","metadata":{},"source":"## 6. Training utilities\n"},{"cell_type":"code","execution_count":null,"id":"105d19c6","metadata":{},"outputs":[],"source":"from sklearn.metrics import roc_auc_score\n\nGRAD_ACCUM_STEPS = 8\nBACKBONE_LR = 1e-5\nHEAD_LR = 5e-4\nMAX_GRAD_NORM = 1.0\nMEMORY_CLEANUP_EVERY = 50\n\n\ndef release_cuda_memory():\n    gc.collect()\n    torch.cuda.empty_cache()\n\n\ndef build_model_and_optimizer():\n    model = TwoLevelMIL(n_planes=N_PLANES).to(DEVICE)\n    backbone_params = list(model.backbone.parameters())\n    backbone_ids = {id(p) for p in backbone_params}\n    other_params = [p for p in model.parameters() if id(p) not in backbone_ids]\n    optimizer = torch.optim.AdamW([\n        {\"params\": backbone_params, \"lr\": BACKBONE_LR},\n        {\"params\": other_params, \"lr\": HEAD_LR},\n    ])\n    scaler = torch.amp.GradScaler(\"cuda\")\n    return model, optimizer, scaler\n\n\ndef _clip_and_step(model, optimizer, scaler):\n    scaler.unscale_(optimizer)\n    torch.nn.utils.clip_grad_norm_(model.parameters(), MAX_GRAD_NORM)\n    scaler.step(optimizer)\n    scaler.update()\n    optimizer.zero_grad()\n\n\ndef train_one_epoch(model, loader, optimizer, scaler, criterion, accum_steps=GRAD_ACCUM_STEPS):\n    model.train()\n    optimizer.zero_grad()\n    total_loss, n_batches, n_skipped, i = 0.0, 0, 0, -1\n    for i, item in enumerate(loader):\n        if len(item[\"series_tensors\"]) == 0:\n            continue\n        series_tensors = [s.to(DEVICE, non_blocking=True) for s in item[\"series_tensors\"]]\n        plane_ids = item[\"plane_ids\"].to(DEVICE)\n        fluid_ids = item[\"fluid_ids\"].to(DEVICE)\n        target = item[\"target\"].to(DEVICE)\n\n        with torch.amp.autocast(\"cuda\", dtype=torch.float16):\n            logits = model(series_tensors, plane_ids, fluid_ids)\n            loss = criterion(logits, target) / accum_steps\n\n        if not torch.isfinite(loss):\n            # A single degenerate study (extreme/corrupted pixel data etc.)\n            # skips its contribution rather than poisoning the whole run's\n            # gradients - this is defense-in-depth, gradient clipping below\n            # is the primary fix for the training-diverges-over-time case.\n            n_skipped += 1\n            optimizer.zero_grad()\n            continue\n\n        scaler.scale(loss).backward()\n        if (i + 1) % accum_steps == 0:\n            _clip_and_step(model, optimizer, scaler)\n\n        total_loss += loss.item() * accum_steps\n        n_batches += 1\n        if (i + 1) % MEMORY_CLEANUP_EVERY == 0:\n            release_cuda_memory()\n\n    if (i + 1) % accum_steps != 0:\n        _clip_and_step(model, optimizer, scaler)\n\n    if n_skipped:\n        print(f\"  (skipped {n_skipped} non-finite-loss studies this epoch)\")\n    return total_loss / max(n_batches, 1)\n\n\n@torch.no_grad()\ndef evaluate(model, loader):\n    model.eval()\n    all_probs, all_targets, all_uids = [], [], []\n    for item in loader:\n        if len(item[\"series_tensors\"]) == 0:\n            continue\n        series_tensors = [s.to(DEVICE, non_blocking=True) for s in item[\"series_tensors\"]]\n        plane_ids = item[\"plane_ids\"].to(DEVICE)\n        fluid_ids = item[\"fluid_ids\"].to(DEVICE)\n        with torch.amp.autocast(\"cuda\", dtype=torch.float16):\n            logits = model(series_tensors, plane_ids, fluid_ids)\n        all_probs.append(torch.sigmoid(logits).float().cpu().numpy())\n        all_targets.append(item[\"target\"].numpy())\n        all_uids.append(item[\"study_uid\"])\n\n    preds = pd.DataFrame(all_probs, columns=LABEL_COLS, index=all_uids)\n    trues = pd.DataFrame(all_targets, columns=LABEL_COLS, index=all_uids)\n    n_nan = int(preds.isna().any(axis=1).sum())\n    if n_nan:\n        print(f\"  WARNING: {n_nan} studies produced non-finite predictions - replacing with 0.5\")\n        preds = preds.fillna(0.5)\n    rows = []\n    for label in LABEL_COLS:\n        y_true = (trues[label] >= 0.5).astype(int)\n        auroc = roc_auc_score(y_true, preds[label]) if y_true.nunique() > 1 else float(\"nan\")\n        rows.append({\"label\": label, \"auroc\": auroc})\n    return pd.DataFrame(rows), preds\n"},{"cell_type":"markdown","id":"4700160d","metadata":{},"source":"## 7. Smoke test\n\nOne forward+backward pass on a few studies - confirms the data pipeline\nand model are wired correctly and gives a rough per-study timing/memory\nestimate before committing to a real training run. Scoped to a manifest\nof just these studies' own series (seconds, not the ~24k-series full\nbuild) so a broken pipeline is cheap to discover and re-push against.\n"},{"cell_type":"code","execution_count":null,"id":"b9e806a5","metadata":{},"outputs":[],"source":"from torch.utils.data import DataLoader\n\nsmoke_uids = (gt_study_uids[:3] or silver_only_uids[:3])\nSERIES_MANIFEST = build_series_manifest(study_uids=smoke_uids)\n\nsmoke_loader = DataLoader(StudyDataset(smoke_uids), batch_size=1, shuffle=False, collate_fn=collate_single)\n\nmodel, optimizer, scaler = build_model_and_optimizer()\ncriterion = nn.BCEWithLogitsLoss()\n\nt0 = time.time()\navg_loss = train_one_epoch(model, smoke_loader, optimizer, scaler, criterion, accum_steps=1)\nelapsed = time.time() - t0\nprint(f\"smoke test: {len(smoke_uids)} studies, avg_loss={avg_loss:.4f}, {elapsed:.1f}s total\")\nif torch.cuda.is_available():\n    print(f\"peak GPU memory: {torch.cuda.max_memory_allocated() / 1e9:.2f} GB\")\nrelease_cuda_memory()\n"},{"cell_type":"markdown","id":"d77360fa","metadata":{},"source":"## 8. Full DICOM manifest\n\nNow that the pipeline is confirmed working end-to-end on a few studies,\nget the full manifest covering every series - needed from here on since\nthe dry run, CV, and final training all draw from the whole dataset.\nLoaded directly from the pre-built manifest Dataset if one is attached\n(a one-time ~28min build otherwise, parallelized but apparently\nthroughput- rather than thread-count-bound - increasing worker count\nalone did not meaningfully speed up a from-scratch build).\n"},{"cell_type":"code","execution_count":null,"id":"feb3695b","metadata":{},"outputs":[],"source":"if PREBUILT_MANIFEST_DIR is not None:\n    t0 = time.time()\n    with open(os.path.join(PREBUILT_MANIFEST_DIR, \"series_manifest.pkl\"), \"rb\") as f:\n        SERIES_MANIFEST = pickle.load(f)\n    print(f\"loaded pre-built manifest for {len(SERIES_MANIFEST)} series in {time.time() - t0:.0f}s\")\nelse:\n    SERIES_MANIFEST = build_series_manifest(cache_path=MANIFEST_PATH)\n\nn_empty = sum(1 for v in SERIES_MANIFEST.values() if not v)\nprint(f\"{n_empty} series have no matched DICOM files (skipped at training time)\")\n"},{"cell_type":"markdown","id":"396c430f","metadata":{},"source":"## 9. Dry run — single train/val split (gate)\n\nA cheap go/no-go checkpoint before the full 5-fold run: one stratified\nsplit of the 58 ground-truth studies, a handful of epochs, and a look at\nwhether validation AUROC clears chance. All silver-labeled studies are\nalways in the training split - only the ground-truth studies have real\nlabels to hold out against.\n"},{"cell_type":"code","execution_count":null,"id":"a116b0d2","metadata":{},"outputs":[],"source":"from iterstrat.ml_stratifiers import MultilabelStratifiedKFold\n\nN_FOLDS = 5\nN_DRY_EPOCHS = 6  # matches N_CV_EPOCHS below - a shorter dry run wouldn't\n                  # have caught the mid-training instability that a real\n                  # fold's full epoch count needs to be tested against\n\ngt_targets_binary = (silver_labels.loc[gt_study_uids, LABEL_COLS] >= 0.5).astype(int).values\nmskf = MultilabelStratifiedKFold(n_splits=N_FOLDS, shuffle=True, random_state=0)\nfolds = list(mskf.split(np.array(gt_study_uids), gt_targets_binary))\n\ndry_train_idx, dry_val_idx = folds[0]\ndry_train_uids = silver_only_uids + [gt_study_uids[i] for i in dry_train_idx]\ndry_val_uids = [gt_study_uids[i] for i in dry_val_idx]\nprint(f\"dry run: {len(dry_train_uids)} train studies, {len(dry_val_uids)} val studies\")\n\nRUN_DRY_RUN = False  # already validated (6 healthy epochs, no NaN) - re-enable\n                      # if the pipeline changes again and needs re-checking\n\nif RUN_DRY_RUN:\n    dry_train_loader = DataLoader(StudyDataset(dry_train_uids), batch_size=1, shuffle=True, collate_fn=collate_single, num_workers=2)\n    dry_val_loader = DataLoader(StudyDataset(dry_val_uids), batch_size=1, shuffle=False, collate_fn=collate_single, num_workers=2)\n\n    model, optimizer, scaler = build_model_and_optimizer()\n    criterion = nn.BCEWithLogitsLoss()\n\n    for epoch in range(N_DRY_EPOCHS):\n        t0 = time.time()\n        train_loss = train_one_epoch(model, dry_train_loader, optimizer, scaler, criterion)\n        metrics_df, _ = evaluate(model, dry_val_loader)\n        print(f\"epoch {epoch + 1}/{N_DRY_EPOCHS}: train_loss={train_loss:.4f}, \"\n              f\"val macro AUROC={metrics_df['auroc'].mean():.3f}, {time.time() - t0:.0f}s\")\n        release_cuda_memory()\n\n    print()\n    print(metrics_df.sort_values(\"auroc\", ascending=False).to_string(index=False))\nelse:\n    print(\"RUN_DRY_RUN is False - skipping (already validated in an earlier run).\")\n"},{"cell_type":"markdown","id":"8c168047","metadata":{},"source":"## 10. Full 5-fold cross-validation (gated)\n\nTrains 5 separate models end-to-end. `RUN_FULL_CV=False` stops here so\nthis can be re-checked cheaply if something looks off. Epoch count is\ndeliberately conservative (3, not the dry run's 6) - real per-epoch\ntiming observed on Kaggle (~23-25min) means 5 folds at 6 epochs each\ndoes not fit in a single session (confirmed: an earlier run was\nauto-cancelled by Kaggle mid-fold-4 at the ~12h session limit). 3\nepochs/fold keeps the full 5-fold run comfortably under that ceiling,\ntrading some per-fold undertraining for actually finishing - the CV\nestimate is a slightly pessimistic read on generalization, not the\nfinal model itself (that comes from a separate, later full-epoch run\non all data with nothing held out).\n"},{"cell_type":"code","execution_count":null,"id":"95a43988","metadata":{},"outputs":[],"source":"RUN_FULL_CV = False  # completed successfully - macro AUROC 0.614 (range\n                      # 0.57-0.65 across 5 folds). Re-enable if the model/\n                      # data pipeline changes again and needs re-validating.\nN_CV_EPOCHS = 3\n\ncv_metrics = []\nif RUN_FULL_CV:\n    for fold_idx, (train_idx, val_idx) in enumerate(folds):\n        train_uids = silver_only_uids + [gt_study_uids[i] for i in train_idx]\n        val_uids = [gt_study_uids[i] for i in val_idx]\n\n        train_loader = DataLoader(StudyDataset(train_uids), batch_size=1, shuffle=True, collate_fn=collate_single, num_workers=2)\n        val_loader = DataLoader(StudyDataset(val_uids), batch_size=1, shuffle=False, collate_fn=collate_single, num_workers=2)\n\n        model, optimizer, scaler = build_model_and_optimizer()\n        criterion = nn.BCEWithLogitsLoss()\n\n        t0 = time.time()\n        for epoch in range(N_CV_EPOCHS):\n            train_loss = train_one_epoch(model, train_loader, optimizer, scaler, criterion)\n            release_cuda_memory()\n            # Per-epoch checkpoint (overwritten each epoch) so a mid-fold\n            # cutoff - which is exactly what happened at fold 4 last run -\n            # doesn't lose the whole fold's progress.\n            torch.save(model.state_dict(), os.path.join(WORK_DIR, f\"model_fold{fold_idx}_inprogress.pt\"))\n        metrics_df, _ = evaluate(model, val_loader)\n        macro_auroc = metrics_df[\"auroc\"].mean()\n        print(f\"fold {fold_idx + 1}/{N_FOLDS}: macro AUROC={macro_auroc:.3f}, {time.time() - t0:.0f}s\")\n\n        metrics_df[\"fold\"] = fold_idx\n        cv_metrics.append(metrics_df)\n        torch.save(model.state_dict(), os.path.join(WORK_DIR, f\"model_fold{fold_idx}.pt\"))\n        release_cuda_memory()\n\n    cv_metrics_df = pd.concat(cv_metrics, ignore_index=True)\n    cv_metrics_df.to_csv(os.path.join(WORK_DIR, \"cv_metrics.csv\"), index=False)\n    summary = cv_metrics_df.groupby(\"label\")[\"auroc\"].agg([\"mean\", \"std\"]).sort_values(\"mean\", ascending=False)\n    print()\n    print(summary.to_string())\n    print(f\"\\noverall macro AUROC: {cv_metrics_df['auroc'].mean():.3f}\")\nelse:\n    print(\"RUN_FULL_CV is False - skipping the 5-fold run.\")\n"},{"cell_type":"markdown","id":"039f5be8","metadata":{},"source":"## 11. Final production training (gated)\n\nOnce the CV numbers above look acceptable, retrain once on all 58\nground-truth studies plus all silver-labeled studies (no held-out split)\n- these are the weights that actually ship to the submission notebook.\nRun as its own separate push once CV has completed - stacking this on\ntop of the same session that also does CV is what overran the session\nlimit last time.\n"},{"cell_type":"code","execution_count":null,"id":"0c0e929e","metadata":{},"outputs":[],"source":"RUN_FINAL_TRAIN = True  # CV validated the setup (see above) - now the real run\nN_FINAL_EPOCHS = 8\n\nif RUN_FINAL_TRAIN:\n    final_uids = silver_only_uids + gt_study_uids\n    final_loader = DataLoader(StudyDataset(final_uids), batch_size=1, shuffle=True, collate_fn=collate_single, num_workers=2)\n\n    model, optimizer, scaler = build_model_and_optimizer()\n    criterion = nn.BCEWithLogitsLoss()\n\n    for epoch in range(N_FINAL_EPOCHS):\n        t0 = time.time()\n        train_loss = train_one_epoch(model, final_loader, optimizer, scaler, criterion)\n        print(f\"epoch {epoch + 1}/{N_FINAL_EPOCHS}: train_loss={train_loss:.4f}, {time.time() - t0:.0f}s\")\n        release_cuda_memory()\n        torch.save(model.state_dict(), os.path.join(WORK_DIR, \"model_final.pt\"))\n\n    print(f\"\\nfinal model trained on {len(final_uids)} studies, saved to model_final.pt\")\nelse:\n    print(\"RUN_FINAL_TRAIN is False - skipping the final production training run.\")\n"},{"cell_type":"markdown","id":"88a027b1","metadata":{},"source":"## Summary\n\n`model_final.pt` (written to `/kaggle/working`) is the weight file the\nsubmission notebook loads for image-only inference. `cv_metrics.csv`\ngives the 5-fold held-out AUROC per finding - the honest read on how\nwell this generalizes, since the final model itself was trained on all\navailable data with nothing held out.\n"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.13"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":154281},{"sourceType":"datasetVersion","sourceId":"aruneembhowmick/knee-mri-silver-labels"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat":4,"nbformat_minor":5}