{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"cell-000","cell_type":"markdown","source":"# Knee MRI Abnormality Classification — ViT (vit-base-mri) Fine-tuning\n\nGoal: 12-class **multi-label** classification of knee MRI studies\n(`rsna-knee-abnormality-detection`), using labels extracted by an LLM from\nradiology reports (`rsna-knee-llm-report-labels/llm_labels_full.csv`).\n\nPipeline overview:\n1. Load labels, auto-detect the 12 target columns.\n2. For each study, locate up to 3 series (Axial / Sagittal / Coronal) directly\n   from DICOM headers (no dependency on an unverified series-metadata CSV).\n3. Sample a handful of slices per plane, preprocess them (percentile\n   normalize, square-pad, resize) the same way as `EDA.ipynb`.\n4. Feed all slices of a study through a shared **ViT-base** backbone\n   (pretrained on MRI data: [`raedinkhaled/vit-base-mri`](https://huggingface.co/raedinkhaled/vit-base-mri)),\n   then combine the per-slice embeddings with **attention-based multiple\n   instance learning (MIL)** into one study-level embedding.\n5. A linear head maps the study embedding to 12 sigmoid outputs (multi-label).\n6. Train with masked BCE (robust to missing labels), track per-class ROC-AUC.\n7. Run inference on `test.csv` and write `submission.csv`.\n\nRuntime notes:\n- Requires a GPU (Settings → Accelerator → GPU) and Internet **on** (to\n  download the HF backbone weights the first time).\n- `raedinkhaled/vit-base-mri` is a ViT-base backbone that was fine-tuned on a\n  different (cardiac, binary cad/healthy) MRI dataset — it is **not**\n  knee-specific. We only reuse its MRI-domain-adapted weights and attach a\n  brand-new 12-way head on top, then fine-tune.\n- Section 2 (`CONFIG`) is the single place to tweak paths/hyperparameters.\n  **Run the \"Inspect raw data\" cell first and check that the printed paths /\n  column names actually match your Kaggle input tab** — dataset file layouts\n  occasionally differ from what's assumed here.","metadata":{}},{"id":"cell-001","cell_type":"markdown","source":"## 1. Setup\n\nInstalls (Kaggle images usually already ship `torch`/`pydicom`/`opencv`; the\n`-q` installs are harmless no-ops if already present) and imports.","metadata":{}},{"id":"cell-002","cell_type":"code","source":"!pip install -q transformers pydicom opencv-python-headless scikit-learn","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-003","cell_type":"code","source":"import os\nimport glob\nimport random\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport cv2\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import roc_auc_score\n\nfrom transformers import ViTModel, get_cosine_schedule_with_warmup\n\nseed = 0\n\ndef seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nseed_everything(seed)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'device: {device}')\nif device.type != 'cuda':\n    print('WARNING: no GPU attached to this session (check Settings -> Accelerator, '\n          'and that this Kaggle account is phone-verified with GPU quota remaining). '\n          'A CPU run of this notebook is far too slow to finish a full training pass, '\n          'so Section 11 automatically caps the training set size when device is CPU.')","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-004","cell_type":"markdown","source":"## 2. Config\n\nCentral place for every path and hyperparameter. `SLICES_PER_PLANE=4` with\n3 planes gives 12 images/study through the ViT backbone per forward pass —\nsized to fit a single Kaggle T4/P100 with AMP. Raise it if you have headroom,\nlower it if you hit OOM.","metadata":{}},{"id":"cell-005","cell_type":"code","source":"CONFIG = {\n    # --- paths (verify against the \"Inspect raw data\" cell output) ---\n    'TRAIN_STUDY_GLOB': '/kaggle/input/competitions/rsna-knee-abnormality-detection/train_series/*',\n    'TEST_STUDY_GLOB':  '/kaggle/input/competitions/rsna-knee-abnormality-detection/test_series/*',\n    'TEST_CSV':         '/kaggle/input/competitions/rsna-knee-abnormality-detection/test.csv',\n    'SAMPLE_SUB_CSV':   '/kaggle/input/competitions/rsna-knee-abnormality-detection/sample_submission.csv',\n    'LABELS_CSV':       '/kaggle/input/datasets/stevenleehans/rsna-knee-llm-report-labels/llm_labels_full.csv',\n    'OUTPUT_DIR':       '/kaggle/working',\n\n    # --- model ---\n    'HF_BACKBONE': 'raedinkhaled/vit-base-mri',\n    'IMG_SIZE': 224,\n    'N_PLANES': 3,\n    'SLICES_PER_PLANE': 4,\n    'FREEZE_ENCODER_LAYERS': 8,     # of 12 ViT-base encoder layers; rest fine-tune\n    'ATTN_HIDDEN_DIM': 256,\n\n    # --- training ---\n    'BATCH_SIZE': 2,                # studies per step\n    'GRAD_ACCUM_STEPS': 4,          # effective batch size 8 studies\n    'EPOCHS': 8,\n    'LR_BACKBONE': 2e-5,\n    'LR_HEAD': 1e-3,\n    'WEIGHT_DECAY': 0.01,\n    'WARMUP_RATIO': 0.1,\n    'VAL_SPLIT': 0.15,\n    'NUM_WORKERS': 2,\n\n    # labels are soft LLM-confidence scores in [0, 1], not strictly binary;\n    # this only thresholds them for the ROC-AUC metric, not for the BCE loss target\n    'LABEL_THRESHOLD': 0.5,\n\n    # if no GPU is attached, cap the training set so a debug run finishes in\n    # minutes instead of hours (see Section 11) -- raise/remove once on GPU\n    'CPU_SAFETY_SUBSET_STUDIES': 40,\n\n    # fallback target columns if auto-detection from LABELS_CSV is ambiguous\n    'FALLBACK_TARGET_COLS': [\n        'ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus',\n        'Medial OA', 'Lateral OA', 'PF OA', 'Effusion',\n        'Synovitis', \"Baker's\", 'Contusion', 'Fracture',\n    ],\n    'NON_LABEL_COLS': [\n        'StudyInstanceUID', 'SeriesInstanceUID', 'PatientID',\n        'report', 'Report', 'text', 'Text', 'notes', 'Notes',\n    ],\n}\n\nos.makedirs(CONFIG['OUTPUT_DIR'], exist_ok=True)\nCONFIG","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-006","cell_type":"markdown","source":"## 3. Inspect raw data\n\nRun this before anything else. It prints what's actually on disk so the\n`CONFIG` paths above (and the column-detection logic in Section 4) can be\ncorrected if the real layout differs.","metadata":{}},{"id":"cell-007","cell_type":"code","source":"train_study_dirs = sorted(glob.glob(CONFIG['TRAIN_STUDY_GLOB']))\nprint(f\"train studies found: {len(train_study_dirs)}\")\nif train_study_dirs:\n    print('example study dir:', train_study_dirs[0])\n    print('example series subfolders:', glob.glob(os.path.join(train_study_dirs[0], '*'))[:5])\n\nlabels_preview = pd.read_csv(CONFIG['LABELS_CSV'])\nprint('\\nlabels csv shape:', labels_preview.shape)\nprint('labels csv columns:', list(labels_preview.columns))\nlabels_preview.head()","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-008","cell_type":"code","source":"test_df_preview = pd.read_csv(CONFIG['TEST_CSV'])\nprint('test.csv shape:', test_df_preview.shape)\nprint('test.csv columns:', list(test_df_preview.columns))\ntest_df_preview.head()","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-009","cell_type":"markdown","source":"## 4. Labels: load + auto-detect the 12 target columns\n\nRather than hard-coding column names that may not match the LLM-generated\nCSV exactly, we detect them: drop known ID/text columns, keep the rest, and\nsanity-check the count is 12 (falling back to `FALLBACK_TARGET_COLS`\nintersected with what's available if detection is ambiguous).","metadata":{}},{"id":"cell-010","cell_type":"code","source":"def detect_target_columns(df, config):\n    candidate_cols = [c for c in df.columns if c not in config['NON_LABEL_COLS']]\n\n    if len(candidate_cols) == 12:\n        target_cols = candidate_cols\n    else:\n        fallback = [c for c in config['FALLBACK_TARGET_COLS'] if c in df.columns]\n        print(f'auto-detected {len(candidate_cols)} candidate columns (expected 12); '\n              f'falling back to {len(fallback)} known target columns.')\n        target_cols = fallback\n\n    assert len(target_cols) == 12, (\n        f'Could not resolve exactly 12 target columns, got {len(target_cols)}: {target_cols}. '\n        f'Inspect labels_df.columns and update CONFIG[\"FALLBACK_TARGET_COLS\"] / '\n        f'CONFIG[\"NON_LABEL_COLS\"] accordingly.'\n    )\n    return target_cols\n\n\nlabels_df = pd.read_csv(CONFIG['LABELS_CSV'])\nTARGET_COLS = detect_target_columns(labels_df, CONFIG)\nNUM_CLASSES = len(TARGET_COLS)\nprint('target columns:', TARGET_COLS)\n\nprint('\\nlabel prevalence (fraction positive, ignoring NaN):')\nprint((labels_df[TARGET_COLS].mean(numeric_only=True)).round(3))","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-011","cell_type":"markdown","source":"## 5. Series selection from DICOM headers\n\nInstead of assuming an unverified `train_series.csv` metadata file exists,\nthe anatomical plane is derived directly from each series' first DICOM's\n`ImageOrientationPatient` (standard direction-cosine trick: the axis the\nslice-normal is most aligned with tells you Axial/Sagittal/Coronal). Ties\nwithin a plane are broken with a small keyword heuristic on\n`SeriesDescription` (fluid-sensitive / fat-sat sequences), then by slice\ncount — mirroring the preference logic in `EDA.ipynb`'s `get_series_df`.","metadata":{}},{"id":"cell-012","cell_type":"code","source":"PLANE_BY_AXIS = {0: 'Sagittal', 1: 'Coronal', 2: 'Axial'}\n\nFLUID_SENSITIVE_KEYWORDS = ['PD', 'PROTON', 'STIR', 'T2', 'FSE']\nFAT_SAT_KEYWORDS = ['FS', 'FATSAT', 'FAT SAT', 'SPAIR', 'SPIR']\n\n\ndef infer_anatomical_plane(dcm):\n    iop = getattr(dcm, 'ImageOrientationPatient', None)\n    if iop is None or len(iop) != 6:\n        return None\n    row_cosine = np.array(iop[0:3], dtype=float)\n    col_cosine = np.array(iop[3:6], dtype=float)\n    normal = np.cross(row_cosine, col_cosine)\n    axis = int(np.argmax(np.abs(normal)))\n    return PLANE_BY_AXIS[axis]\n\n\ndef series_quality_score(series_description):\n    if series_description is None:\n        return 0\n    desc = str(series_description).upper()\n    score = 0\n    for kw in FLUID_SENSITIVE_KEYWORDS + FAT_SAT_KEYWORDS:\n        if kw in desc:\n            score += 1\n    return score\n\n\ndef select_series_per_plane(study_dir):\n    \"\"\"Returns {plane_name: series_dir} picking the best series per plane.\"\"\"\n    best = {}\n    for series_dir in glob.glob(os.path.join(study_dir, '*')):\n        dcm_files = glob.glob(os.path.join(series_dir, '*.dcm'))\n        if not dcm_files:\n            continue\n        dcm = pydicom.dcmread(dcm_files[0], stop_before_pixels=True)\n        plane = infer_anatomical_plane(dcm)\n        if plane is None:\n            continue\n        score = series_quality_score(getattr(dcm, 'SeriesDescription', None))\n        n_slices = len(dcm_files)\n        if plane not in best or (score, n_slices) > (best[plane][0], best[plane][1]):\n            best[plane] = (score, n_slices, series_dir)\n    return {plane: v[2] for plane, v in best.items()}","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-013","cell_type":"code","source":"_sample_study = train_study_dirs[0]\n_planes = select_series_per_plane(_sample_study)\nprint('study:', os.path.basename(_sample_study))\nfor plane, series_dir in _planes.items():\n    n_slices = len(glob.glob(os.path.join(series_dir, '*.dcm')))\n    print(f'  {plane:10s} -> {os.path.basename(series_dir)}  ({n_slices} slices)')","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-014","cell_type":"markdown","source":"## 6. Slice sampling + preprocessing\n\n`sample_slice_paths` sorts a series by `InstanceNumber` and evenly samples\n`SLICES_PER_PLANE` of them (same `np.linspace` trick as `EDA.ipynb`'s\n`get_dcm_df`, just returning fewer slices to keep the MIL forward pass\ncheap). `preprocess_slice` is copied from `EDA.ipynb` unchanged (MONOCHROME1\ninversion, 1st–99th percentile normalization, symmetric square padding,\nresize) — only `target_size` now defaults to 224 to match ViT's input.","metadata":{}},{"id":"cell-015","cell_type":"code","source":"def sample_slice_paths(series_dir, n_slices):\n    dcm_paths = glob.glob(os.path.join(series_dir, '*.dcm'))\n    rows = []\n    for p in dcm_paths:\n        dcm = pydicom.dcmread(p, stop_before_pixels=True)\n        rows.append((p, int(dcm.InstanceNumber)))\n    rows.sort(key=lambda r: r[1])\n    paths_sorted = [r[0] for r in rows]\n    if not paths_sorted:\n        return []\n    idx = np.linspace(0, len(paths_sorted) - 1, min(n_slices, len(paths_sorted)))\n    idx = sorted(set(idx.round().astype(int).tolist()))\n    return [paths_sorted[i] for i in idx]\n\n\ndef preprocess_slice(dcm_path, target_size=(224, 224)):\n    # some DICOM files in this dataset have truncated/corrupt pixel data\n    # (e.g. \"number of bytes of pixel data is less than expected\"); skip\n    # those rather than crashing the whole training run\n    try:\n        dcm = pydicom.dcmread(dcm_path)\n        img = dcm.pixel_array.astype(np.float32)\n    except Exception as e:\n        print(f'skipping unreadable DICOM {dcm_path}: {e}')\n        return None\n\n    if getattr(dcm, 'PhotometricInterpretation', 'MONOCHROME2') == 'MONOCHROME1':\n        img = img.max() - img\n\n    p1, p99 = np.percentile(img, (1, 99))\n    if p99 > p1:\n        img = np.clip(img, p1, p99)\n        img = (img - p1) / (p99 - p1)\n    else:\n        img = np.zeros_like(img)\n\n    h, w = img.shape\n    if h > w:\n        pad_left = (h - w) // 2\n        pad_right = (h - w) - pad_left\n        img = np.pad(img, ((0, 0), (pad_left, pad_right)), mode='constant', constant_values=0)\n    elif w > h:\n        pad_top = (w - h) // 2\n        pad_bottom = (w - h) - pad_top\n        img = np.pad(img, ((pad_top, pad_bottom), (0, 0)), mode='constant', constant_values=0)\n\n    img = cv2.resize(img, target_size, interpolation=cv2.INTER_AREA)\n    return img","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-016","cell_type":"markdown","source":"## 7. Dataset\n\nBuilds a fixed-size `(N_MAX, 3, 224, 224)` tensor per study, where\n`N_MAX = N_PLANES * SLICES_PER_PLANE`. Grayscale slices are replicated to 3\nchannels (ViT was pretrained on natural RGB images) and normalized with\n`mean=std=0.5` — the exact stats in `vit-base-mri`'s `preprocessor_config.json`.\nStudies with fewer than 3 planes found, or a plane with fewer available\nslices, are **padded with zeros and a boolean `mask`** rather than dropped —\nthe attention-MIL pooling in Section 9 respects this mask so padding never\ninfluences the study embedding.","metadata":{}},{"id":"cell-017","cell_type":"code","source":"class RSNAViTDataset(Dataset):\n    def __init__(self, study_dirs, config, labels_df=None, target_cols=None):\n        self.study_dirs = study_dirs\n        self.config = config\n        self.labels_df = labels_df\n        self.target_cols = target_cols\n        self.img_size = config['IMG_SIZE']\n        self.slices_per_plane = config['SLICES_PER_PLANE']\n        self.n_max = config['N_PLANES'] * config['SLICES_PER_PLANE']\n\n    def __len__(self):\n        return len(self.study_dirs)\n\n    def _load_study_volume(self, study_dir):\n        planes = select_series_per_plane(study_dir)\n        imgs = []\n        for plane in ('Axial', 'Sagittal', 'Coronal'):\n            series_dir = planes.get(plane)\n            if series_dir is None:\n                continue\n            for slice_path in sample_slice_paths(series_dir, self.slices_per_plane):\n                sl = preprocess_slice(slice_path, (self.img_size, self.img_size))\n                if sl is not None:\n                    imgs.append(sl)\n        if not imgs:\n            return np.zeros((0, 3, self.img_size, self.img_size), dtype=np.float32)\n        vol = np.stack(imgs).astype(np.float32)               # (n_found, H, W), in [0, 1]\n        vol = (vol - 0.5) / 0.5                                 # ViT normalization (mean=std=0.5)\n        vol = np.repeat(vol[:, None, :, :], 3, axis=1)          # grayscale -> pseudo-RGB\n        return vol\n\n    def __getitem__(self, idx):\n        study_dir = self.study_dirs[idx]\n        study_uid = os.path.basename(study_dir)\n        vol = self._load_study_volume(study_dir)\n\n        images = np.zeros((self.n_max, 3, self.img_size, self.img_size), dtype=np.float32)\n        mask = np.zeros((self.n_max,), dtype=np.float32)\n        n_use = min(vol.shape[0], self.n_max)\n        images[:n_use] = vol[:n_use]\n        mask[:n_use] = 1.0\n\n        images = torch.from_numpy(images)\n        mask = torch.from_numpy(mask)\n\n        if self.labels_df is None:\n            return images, mask, study_uid\n\n        row = self.labels_df[self.labels_df['StudyInstanceUID'] == study_uid]\n        if len(row) == 0:\n            label = np.full(len(self.target_cols), np.nan, dtype=np.float32)\n        else:\n            label = row.iloc[0][self.target_cols].to_numpy(dtype=np.float32)\n        label = torch.from_numpy(label)\n        return images, mask, label, study_uid","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-018","cell_type":"markdown","source":"## 8. Sanity check\n\nLoad one study end-to-end and visualize its sampled slices before spending\nany GPU time on training.","metadata":{}},{"id":"cell-019","cell_type":"code","source":"_check_ds = RSNAViTDataset(train_study_dirs[:1], CONFIG, labels_df, TARGET_COLS)\n_images, _mask, _label, _uid = _check_ds[0]\nprint('images:', _images.shape, 'mask sum (real slices):', int(_mask.sum().item()))\nprint('study:', _uid)\nprint('label:', dict(zip(TARGET_COLS, _label.tolist())))\n\n_n_show = int(_mask.sum().item())\nfig, axes = plt.subplots(1, max(_n_show, 1), figsize=(2.5 * max(_n_show, 1), 3))\naxes = np.atleast_1d(axes)\nfor i in range(_n_show):\n    img = (_images[i, 0].numpy() * 0.5) + 0.5   # undo normalization for display\n    axes[i].imshow(img, cmap='gray')\n    axes[i].axis('off')\nplt.tight_layout()\nplt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-020","cell_type":"markdown","source":"## 9. Model: ViT backbone + attention-MIL head\n\nEvery sampled slice of a study is embedded independently by the shared ViT\nbackbone (its `[CLS]` token). An attention module (Ilse et al.-style) scores\neach slice's embedding and produces a softmax-weighted sum — this is the MIL\naggregation step that turns a variable number of 2D slice embeddings into\none fixed-size study embedding. Padded slots are masked out of the softmax\nwith `-inf` so they get exactly zero weight. A linear layer then maps the\npooled 768-d study embedding to the 12 class logits.\n\nThe first `FREEZE_ENCODER_LAYERS` of the 12 ViT-base encoder layers (plus\nthe patch/position embeddings) are frozen — full fine-tuning is expensive\nhere because each study forwards `N_PLANES * SLICES_PER_PLANE` images\nthrough the backbone, effectively multiplying the batch size.","metadata":{}},{"id":"cell-021","cell_type":"code","source":"class AttentionMIL(nn.Module):\n    def __init__(self, embed_dim, hidden_dim):\n        super().__init__()\n        self.attn = nn.Sequential(\n            nn.Linear(embed_dim, hidden_dim),\n            nn.Tanh(),\n            nn.Linear(hidden_dim, 1),\n        )\n\n    def forward(self, embeddings, mask):\n        logits = self.attn(embeddings).squeeze(-1)                  # (B, N)\n        logits = logits.masked_fill(mask == 0, float('-inf'))\n        weights = torch.softmax(logits, dim=1)                       # (B, N)\n        pooled = torch.sum(weights.unsqueeze(-1) * embeddings, dim=1)  # (B, D)\n        return pooled, weights\n\n\nclass ViTKneeMIL(nn.Module):\n    def __init__(self, config, num_classes):\n        super().__init__()\n        self.backbone = ViTModel.from_pretrained(config['HF_BACKBONE'], add_pooling_layer=False)\n        embed_dim = self.backbone.config.hidden_size\n\n        for p in self.backbone.embeddings.parameters():\n            p.requires_grad = False\n        for layer in self.backbone.encoder.layer[:config['FREEZE_ENCODER_LAYERS']]:\n            for p in layer.parameters():\n                p.requires_grad = False\n\n        self.pool = AttentionMIL(embed_dim, config['ATTN_HIDDEN_DIM'])\n        self.classifier = nn.Linear(embed_dim, num_classes)\n\n    def forward(self, images, mask):\n        B, N, C, H, W = images.shape\n        flat = images.view(B * N, C, H, W)\n        cls_tokens = self.backbone(pixel_values=flat).last_hidden_state[:, 0]  # (B*N, D)\n        embeddings = cls_tokens.view(B, N, -1)\n        pooled, attn_weights = self.pool(embeddings, mask)\n        logits = self.classifier(pooled)\n        return logits, attn_weights\n\n\nmodel = ViTKneeMIL(CONFIG, NUM_CLASSES).to(device)\nn_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\nn_total = sum(p.numel() for p in model.parameters())\nprint(f'trainable params: {n_trainable:,} / {n_total:,}')","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-022","cell_type":"markdown","source":"## 10. Loss (NaN-masked multi-label BCE)\n\nThe LLM-derived labels may be missing/uncertain for some study-class pairs.\n`masked_bce_loss` computes BCE only over non-NaN entries so those don't\ncontribute gradient or skew the loss.","metadata":{}},{"id":"cell-023","cell_type":"code","source":"def masked_bce_loss(logits, targets, pos_weight=None):\n    valid = ~torch.isnan(targets)\n    targets_filled = torch.nan_to_num(targets, nan=0.0)\n    loss = F.binary_cross_entropy_with_logits(\n        logits, targets_filled, pos_weight=pos_weight, reduction='none'\n    )\n    loss = loss * valid.float()\n    denom = valid.float().sum().clamp(min=1.0)\n    return loss.sum() / denom","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-024","cell_type":"markdown","source":"## 11. Train / validation split\n\nA plain random split (multi-label stratification isn't worth the extra\ndependency at this dataset size). Only studies with at least one non-NaN\nlabel are kept, since an all-NaN row contributes nothing to the loss.","metadata":{}},{"id":"cell-025","cell_type":"code","source":"labeled_mask = labels_df[TARGET_COLS].notna().any(axis=1)\nlabeled_uids = set(labels_df.loc[labeled_mask, 'StudyInstanceUID'])\n\nusable_study_dirs = [d for d in train_study_dirs if os.path.basename(d) in labeled_uids]\nprint(f'studies with a directory AND at least one label: {len(usable_study_dirs)} / {len(train_study_dirs)}')\n\nif device.type != 'cuda':\n    rng = random.Random(seed)\n    rng.shuffle(usable_study_dirs)\n    usable_study_dirs = usable_study_dirs[:CONFIG['CPU_SAFETY_SUBSET_STUDIES']]\n    print(f'no GPU attached -> capped to {len(usable_study_dirs)} studies for a CPU debug run '\n          f'(raise/remove CONFIG[\"CPU_SAFETY_SUBSET_STUDIES\"] once GPU is available)')\n\ntrain_dirs, val_dirs = train_test_split(\n    usable_study_dirs, test_size=CONFIG['VAL_SPLIT'], random_state=seed\n)\nprint(f'train: {len(train_dirs)}  val: {len(val_dirs)}')\n\ntrain_ds = RSNAViTDataset(train_dirs, CONFIG, labels_df, TARGET_COLS)\nval_ds = RSNAViTDataset(val_dirs, CONFIG, labels_df, TARGET_COLS)\n\ntrain_loader = DataLoader(train_ds, batch_size=CONFIG['BATCH_SIZE'], shuffle=True,\n                           num_workers=CONFIG['NUM_WORKERS'], drop_last=True)\nval_loader = DataLoader(val_ds, batch_size=CONFIG['BATCH_SIZE'], shuffle=False,\n                         num_workers=CONFIG['NUM_WORKERS'])","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-026","cell_type":"markdown","source":"## 12. Optimizer, scheduler, class-imbalance weighting\n\nA differential learning rate is used: a low LR for the (partially frozen)\npretrained backbone, a higher LR for the freshly-initialized attention/head.\n`pos_weight` is derived from training-set label prevalence so rare\nabnormalities (e.g. Fracture) aren't drowned out by common negatives.","metadata":{}},{"id":"cell-027","cell_type":"code","source":"prevalence = labels_df.loc[labels_df['StudyInstanceUID'].isin({os.path.basename(d) for d in train_dirs}), TARGET_COLS].mean(numeric_only=True)\npos_weight = ((1 - prevalence) / prevalence.clip(lower=0.01)).clip(upper=20.0)\npos_weight = torch.tensor(pos_weight.to_numpy(dtype=np.float32), device=device)\nprint('pos_weight per class:')\nprint(pd.Series(pos_weight.cpu().numpy(), index=TARGET_COLS).round(2))\n\nbackbone_params = [p for p in model.backbone.parameters() if p.requires_grad]\nhead_params = list(model.pool.parameters()) + list(model.classifier.parameters())\n\noptimizer = torch.optim.AdamW([\n    {'params': backbone_params, 'lr': CONFIG['LR_BACKBONE']},\n    {'params': head_params, 'lr': CONFIG['LR_HEAD']},\n], weight_decay=CONFIG['WEIGHT_DECAY'])\n\nsteps_per_epoch = len(train_loader) // CONFIG['GRAD_ACCUM_STEPS']\ntotal_steps = max(steps_per_epoch * CONFIG['EPOCHS'], 1)\nscheduler = get_cosine_schedule_with_warmup(\n    optimizer,\n    num_warmup_steps=int(total_steps * CONFIG['WARMUP_RATIO']),\n    num_training_steps=total_steps,\n)\nscaler = torch.cuda.amp.GradScaler(enabled=(device.type == 'cuda'))","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-028","cell_type":"markdown","source":"## 13. Train / eval loops\n\nMixed precision (`autocast` + `GradScaler`) and gradient accumulation keep\nmemory in check given each study expands to `N_PLANES * SLICES_PER_PLANE`\nimages. Validation reports macro ROC-AUC, skipping any class that has only\none label value present in the current val split (AUC is undefined there).","metadata":{}},{"id":"cell-029","cell_type":"code","source":"def run_epoch(model, loader, optimizer, scheduler, scaler, pos_weight, train):\n    model.train(mode=train)\n    total_loss, n_batches = 0.0, 0\n    all_logits, all_targets = [], []\n\n    if train:\n        optimizer.zero_grad()\n\n    for step, (images, mask, targets, _uid) in enumerate(loader):\n        images, mask, targets = images.to(device), mask.to(device), targets.to(device)\n\n        with torch.set_grad_enabled(train), torch.cuda.amp.autocast(enabled=(device.type == 'cuda')):\n            logits, _ = model(images, mask)\n            loss = masked_bce_loss(logits, targets, pos_weight=pos_weight)\n\n        if train:\n            scaler.scale(loss / CONFIG['GRAD_ACCUM_STEPS']).backward()\n            if (step + 1) % CONFIG['GRAD_ACCUM_STEPS'] == 0:\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n                scheduler.step()\n\n        total_loss += loss.item()\n        n_batches += 1\n        all_logits.append(logits.detach().float().cpu())\n        all_targets.append(targets.detach().float().cpu())\n\n    logits_cat = torch.cat(all_logits).numpy()\n    targets_cat = torch.cat(all_targets).numpy()\n    probs_cat = 1 / (1 + np.exp(-logits_cat))\n\n    # labels are soft LLM-confidence scores, not strictly {0, 1}; roc_auc_score\n    # needs a binary ground truth, so threshold only for this metric\n    targets_binary = (targets_cat >= CONFIG['LABEL_THRESHOLD']).astype(np.float32)\n    targets_binary[np.isnan(targets_cat)] = np.nan\n\n    aucs = {}\n    for i, col in enumerate(TARGET_COLS):\n        valid = ~np.isnan(targets_binary[:, i])\n        y_true = targets_binary[valid, i]\n        if valid.sum() > 0 and len(np.unique(y_true)) > 1:\n            aucs[col] = roc_auc_score(y_true, probs_cat[valid, i])\n\n    macro_auc = float(np.mean(list(aucs.values()))) if aucs else float('nan')\n    return total_loss / max(n_batches, 1), macro_auc, aucs","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-030","cell_type":"markdown","source":"## 14. Run training\n\nSaves the checkpoint with the best validation macro-AUC to\n`CONFIG['OUTPUT_DIR']/vit_knee_mil_best.pt`.","metadata":{}},{"id":"cell-031","cell_type":"code","source":"history = {'train_loss': [], 'val_loss': [], 'val_auc': []}\nbest_val_auc = -1.0\nbest_ckpt_path = os.path.join(CONFIG['OUTPUT_DIR'], 'vit_knee_mil_best.pt')\n\nfor epoch in range(CONFIG['EPOCHS']):\n    train_loss, _, _ = run_epoch(model, train_loader, optimizer, scheduler, scaler, pos_weight, train=True)\n    val_loss, val_auc, val_aucs = run_epoch(model, val_loader, optimizer, scheduler, scaler, pos_weight, train=False)\n\n    history['train_loss'].append(train_loss)\n    history['val_loss'].append(val_loss)\n    history['val_auc'].append(val_auc)\n\n    print(f\"epoch {epoch + 1}/{CONFIG['EPOCHS']}  \"\n          f\"train_loss={train_loss:.4f}  val_loss={val_loss:.4f}  val_macro_auc={val_auc:.4f}\")\n\n    if val_auc > best_val_auc:\n        best_val_auc = val_auc\n        torch.save({'model_state_dict': model.state_dict(),\n                    'target_cols': TARGET_COLS,\n                    'config': CONFIG,\n                    'val_auc': val_auc}, best_ckpt_path)\n        print(f'  -> new best (macro AUC {val_auc:.4f}), saved to {best_ckpt_path}')\n\nprint('\\nper-class AUC (final epoch):')\nprint(pd.Series(val_aucs).round(3))","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-032","cell_type":"markdown","source":"## 15. Training curves","metadata":{}},{"id":"cell-033","cell_type":"code","source":"fig, axes = plt.subplots(1, 2, figsize=(12, 4))\naxes[0].plot(history['train_loss'], label='train')\naxes[0].plot(history['val_loss'], label='val')\naxes[0].set_title('masked BCE loss')\naxes[0].set_xlabel('epoch')\naxes[0].legend()\n\naxes[1].plot(history['val_auc'], color='tab:green')\naxes[1].set_title('val macro ROC-AUC')\naxes[1].set_xlabel('epoch')\n\nplt.tight_layout()\nplt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-034","cell_type":"markdown","source":"## 16. Reload best checkpoint\n\nReloads the best validation-AUC weights before running inference, so a\nlater epoch that overfit doesn't get used for predictions.","metadata":{}},{"id":"cell-035","cell_type":"code","source":"checkpoint = torch.load(best_ckpt_path, map_location=device)\nmodel.load_state_dict(checkpoint['model_state_dict'])\nmodel.eval()\nprint(f\"loaded checkpoint with val macro AUC = {checkpoint['val_auc']:.4f}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-036","cell_type":"markdown","source":"## 17. Inference on `test.csv`\n\nBuilds the test set the same way as train (no labels), runs the model with\nno gradient tracking, and sigmoids the logits into per-class probabilities.","metadata":{}},{"id":"cell-037","cell_type":"code","source":"test_df = pd.read_csv(CONFIG['TEST_CSV'])\ntest_study_dirs_all = sorted(glob.glob(CONFIG['TEST_STUDY_GLOB']))\ntest_dir_by_uid = {os.path.basename(d): d for d in test_study_dirs_all}\n\ntest_study_uids = test_df['StudyInstanceUID'].tolist()\nmissing = [u for u in test_study_uids if u not in test_dir_by_uid]\nif missing:\n    print(f'warning: {len(missing)} test studies have no matching directory under TEST_STUDY_GLOB '\n          f'(check CONFIG[\"TEST_STUDY_GLOB\"]); they will get 0.5 default probabilities.')\n\ntest_dirs_ordered = [test_dir_by_uid[u] for u in test_study_uids if u in test_dir_by_uid]\ntest_ds = RSNAViTDataset(test_dirs_ordered, CONFIG, labels_df=None, target_cols=None)\ntest_loader = DataLoader(test_ds, batch_size=CONFIG['BATCH_SIZE'], shuffle=False,\n                          num_workers=CONFIG['NUM_WORKERS'])\n\npred_rows = {}\nwith torch.no_grad():\n    for images, mask, uids in test_loader:\n        images, mask = images.to(device), mask.to(device)\n        with torch.cuda.amp.autocast(enabled=(device.type == 'cuda')):\n            logits, _ = model(images, mask)\n        probs = torch.sigmoid(logits).float().cpu().numpy()\n        for uid, row in zip(uids, probs):\n            pred_rows[uid] = row\n\nprint(f'predicted {len(pred_rows)} / {len(test_study_uids)} test studies')","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-038","cell_type":"markdown","source":"## 18. Build `submission.csv`\n\nAny test study whose directory wasn't found falls back to a neutral 0.5\nprobability per class (flagged by the warning above) rather than crashing.","metadata":{}},{"id":"cell-039","cell_type":"code","source":"default_probs = np.full(NUM_CLASSES, 0.5, dtype=np.float32)\nsubmission_rows = []\nfor uid in test_study_uids:\n    probs = pred_rows.get(uid, default_probs)\n    submission_rows.append([uid] + probs.tolist())\n\nsubmission_df = pd.DataFrame(submission_rows, columns=['StudyInstanceUID'] + TARGET_COLS)\n\nsubmission_path = os.path.join(CONFIG['OUTPUT_DIR'], 'submission.csv')\nsubmission_df.to_csv(submission_path, index=False)\nprint(f'wrote {submission_path}  shape={submission_df.shape}')\nsubmission_df.head()","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-040","cell_type":"markdown","source":"## 19. Caveats / what to tune\n\n- **Backbone domain mismatch**: `vit-base-mri` was fine-tuned on a cardiac\n  (cad/healthy) MRI dataset, not knee MRI. It's used purely as an\n  MRI-adapted starting point for the ViT-base weights — the 12-class head is\n  trained from scratch here, and unfreezing more encoder layers\n  (`CONFIG['FREEZE_ENCODER_LAYERS']`) may help once you have a training-time\n  budget to spend.\n- **Series/plane selection is heuristic**: `select_series_per_plane` infers\n  Axial/Sagittal/Coronal from `ImageOrientationPatient` and scores series by\n  keywords in `SeriesDescription`. If a study has multiple series per plane\n  with no distinguishing keywords, the tie-break falls to slice count —\n  double check a few studies visually (Section 8) if results look off.\n- **`TEST_STUDY_GLOB` / `TRAIN_STUDY_GLOB` paths are assumptions** — the\n  train path was given directly, but the test series directory name\n  (`test_series`) was inferred by analogy. Confirm via Section 3's directory\n  listing and fix `CONFIG` if it differs.\n- **Label sparsity**: `masked_bce_loss` ignores NaN labels per-class per-study\n  rather than dropping the whole row, which matters if the LLM-derived CSV\n  has partial coverage.\n- **Compute budget**: `SLICES_PER_PLANE`, `BATCH_SIZE`, and `EPOCHS` are set\n  conservatively for a single Kaggle GPU session. Scale up if you have more\n  GPU-hours available.","metadata":{}}]}