{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":[{"id":"54718e0a-19aa-4a98-8c59-510bcda4c7cf","cell_type":"markdown","source":"# 🦴 RSNA Knee Abnormality Detection\n## Multimodal Training Notebook (Kaggle GPU)\n\n**Competition**: [RSNA Knee Abnormality Detection](https://www.kaggle.com/competitions/rsna-knee-abnormality-detection)  \n**Architecture**: 2.5D ConvNeXt-Base (vision) + XLM-RoBERTa (text) → 12-label sigmoid  \n**Metric**: Macro-averaged AUC-ROC  \n\n### 🚀 How to use this notebook:\n1. Go to **Edit** → **Notebook Settings** → set **Accelerator = GPU T4 x2** (or P100)\n2. Make sure **Internet** is turned ON (needed for model downloads)\n3. Click **Run All**\n4. Submit `submission.csv` from the output","metadata":{}},{"id":"8a044008-8c15-4cd1-8cbf-d2f9345985c2","cell_type":"markdown","source":"## ⚙️ 1. Install Dependencies","metadata":{}},{"id":"43f50f89-4ca8-42c6-85df-2ccd1b8e0670","cell_type":"code","source":"%%capture\n!pip install timm>=1.0.3 albumentations>=1.4.6 transformers>=4.40.0 sentencepiece pydicom SimpleITK torchio wandb -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T18:28:08.883199Z","iopub.execute_input":"2026-08-24T18:28:08.884204Z","iopub.status.idle":"2026-08-24T18:28:12.55079Z","shell.execute_reply.started":"2026-08-24T18:28:08.884149Z","shell.execute_reply":"2026-08-24T18:28:12.549965Z"}},"outputs":[],"execution_count":null},{"id":"fc3b1ea3-3804-4c12-8e7c-192aae0c16ae","cell_type":"markdown","source":"## 📦 2. Imports & Seed","metadata":{}},{"id":"f715dc38-730f-4f96-a6c1-1819300526e4","cell_type":"code","source":"import os, gc, math, random, warnings, time\nfrom pathlib import Path\nfrom typing import Dict, List, Optional, Tuple\n\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport pydicom\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.cuda.amp import GradScaler, autocast\nfrom torch.utils.data import Dataset, DataLoader\n\nimport timm\nimport albumentations as A\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score\nfrom tqdm.notebook import tqdm\nfrom transformers import AutoModel, AutoTokenizer\n\nwarnings.filterwarnings('ignore')\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n\nset_seed(42)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'Device: {device}')\nif torch.cuda.is_available():\n    print(f'GPU: {torch.cuda.get_device_name(0)}')\n    print(f'VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T18:28:12.552792Z","iopub.execute_input":"2026-08-24T18:28:12.553062Z","iopub.status.idle":"2026-08-24T18:28:12.56442Z","shell.execute_reply.started":"2026-08-24T18:28:12.553035Z","shell.execute_reply":"2026-08-24T18:28:12.563674Z"}},"outputs":[],"execution_count":null},{"id":"7d9e437e-3e2c-459f-9a22-1c0a003beadd","cell_type":"markdown","source":"## ⚙️ 3. Configuration","metadata":{}},{"id":"5a1fb66e-d5e3-4f33-a416-bc8092d36997","cell_type":"code","source":"# ─── Paths ───────────────────────────────────────────────────────\n\nCOMP_DIR = Path('/kaggle/input/competitions/rsna-knee-abnormality-detection')\n\nOUTPUT_DIR  = Path('/kaggle/working')\nCKPT_DIR    = OUTPUT_DIR / 'checkpoints'\nCKPT_DIR.mkdir(parents=True, exist_ok=True)\n\n# ─── Labels ──────────────────────────────────────────────────────\nLABEL_NAMES = [\n    'ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', \n    'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', \n    'Synovitis', \"Baker's\", 'Contusion', 'Fracture'\n]\nNUM_LABELS = len(LABEL_NAMES)  # 12\n\n# ─── Data ────────────────────────────────────────────────────────\nNUM_SLICES   = 16     # slices per MRI volume (2.5D)\nIMG_SIZE     = 256    # spatial resolution\nN_FOLDS      = 5\nSEED         = 42\n\n# ─── Model ───────────────────────────────────────────────────────\nVISION_BACKBONE    = 'convnext_base'     # timm model\nTEXT_BACKBONE      = 'xlm-roberta-base'  # HuggingFace model\nVISION_OUT_DIM     = 512\nTEXT_OUT_DIM       = 256\nTEXT_MAX_LEN       = 256\nFREEZE_TEXT_LAYERS = 8\nFUSION_DIM         = 512\nUSE_TEXT           = True   # Set False for vision-only baseline\n\n# ─── Training ────────────────────────────────────────────────────\nEPOCHS          = 15\nBATCH_SIZE      = 4    # per GPU; increase if you have A100\nGRAD_ACCUM      = 8    # effective batch = 32\nLR_VISION       = 2e-5\nLR_TEXT         = 5e-6\nLR_HEAD         = 1e-4\nWEIGHT_DECAY    = 1e-2\nWARMUP_EPOCHS   = 2\nLABEL_SMOOTHING = 0.05\nMAX_GRAD_NORM   = 1.0\nAMP             = True\nNUM_WORKERS     = 2\n\n# ─── Which fold to train (set to None for all folds) ─────────────\nTRAIN_FOLD = None   # None = all 5 folds; 1 = fold 1 only (faster)\n\nprint('Config loaded ✓')\nprint(f'  Backbone: {VISION_BACKBONE}  |  Text: {USE_TEXT}  |  Slices: {NUM_SLICES}  |  IMG: {IMG_SIZE}')\nprint(f'  Epochs: {EPOCHS}  |  Batch: {BATCH_SIZE}×{GRAD_ACCUM}={BATCH_SIZE*GRAD_ACCUM} effective  |  Folds: {N_FOLDS}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T18:28:12.565398Z","iopub.execute_input":"2026-08-24T18:28:12.565853Z","iopub.status.idle":"2026-08-24T18:28:12.584206Z","shell.execute_reply.started":"2026-08-24T18:28:12.56583Z","shell.execute_reply":"2026-08-24T18:28:12.583676Z"}},"outputs":[],"execution_count":null},{"id":"d896a13b-8c12-4cf7-946a-6809c934faad","cell_type":"markdown","source":"## 📂 4. Load & Inspect Data","metadata":{}},{"id":"102c1889-f1e7-4b46-8015-f26c87b7ee88","cell_type":"code","source":"# List available files\nprint('Competition files:')\nfor f in sorted(COMP_DIR.iterdir()):\n    size = f.stat().st_size / 1e9 if f.is_file() else 0\n    print(f'  {f.name}  ({size:.2f} GB)' if f.is_file() else f'  {f.name}/')\n\n# Load CSVs\n# In Section 4 (Load & Inspect Data)\ntrain_df = pd.read_csv(COMP_DIR / 'train.csv')\n\n# Add this line to fill NaN labels with 0\ntrain_df[LABEL_NAMES] = train_df[LABEL_NAMES].fillna(0)\n\nprint(f'\\nTrain: {len(train_df):,} studies | Test: {len(test_df):,} studies')\nprint(f'Columns: {list(train_df.columns)}')\ntrain_df.head(3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T18:28:12.585098Z","iopub.execute_input":"2026-08-24T18:28:12.585265Z","iopub.status.idle":"2026-08-24T18:28:12.704174Z","shell.execute_reply.started":"2026-08-24T18:28:12.585249Z","shell.execute_reply":"2026-08-24T18:28:12.703643Z"}},"outputs":[],"execution_count":null},{"id":"fb98dcc0-55e0-4299-b6d2-4d17334a0ae7","cell_type":"code","source":"# Label distribution\nfig, axes = plt.subplots(1, 2, figsize=(16, 5))\npos_rates = train_df[LABEL_NAMES].mean().sort_values(ascending=True)\ncolors = plt.cm.plasma(np.linspace(0.2, 0.9, len(LABEL_NAMES)))\naxes[0].barh(pos_rates.index, pos_rates.values, color=colors)\naxes[0].set_title('Positive Rate per Label', fontweight='bold')\naxes[0].set_xlabel('Rate')\nfor i, v in enumerate(pos_rates.values):\n    axes[0].text(v + 0.005, i, f'{v:.1%}', va='center', fontsize=8)\n\nlabel_counts = train_df[LABEL_NAMES].sum(axis=1)\naxes[1].hist(label_counts, bins=range(14), color='#7B2FBE', edgecolor='white', alpha=0.85)\naxes[1].set_title('Labels per Study', fontweight='bold')\naxes[1].set_xlabel('# Positive Labels')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T18:28:12.706098Z","iopub.execute_input":"2026-08-24T18:28:12.706632Z","iopub.status.idle":"2026-08-24T18:28:13.019684Z","shell.execute_reply.started":"2026-08-24T18:28:12.706609Z","shell.execute_reply":"2026-08-24T18:28:13.018925Z"}},"outputs":[],"execution_count":null},{"id":"6b16387c-e23a-46fc-9179-90a0494f97ef","cell_type":"markdown","source":"## 🔬 5. DICOM Utilities","metadata":{}},{"id":"41f44b5c-748b-4665-a0aa-0090ca16de24","cell_type":"code","source":"def load_volume(study_id: str, img_dir: Path, num_slices: int = 16, img_size: int = 256) -> np.ndarray:\n    \"\"\"\n    Load a knee MRI study as a (num_slices, H, W) uint8 numpy array.\n    Handles nested directory structure: img_dir/study_id/**/series/*.dcm\n    Picks the longest series (most slices) per study.\n    \"\"\"\n    study_path = img_dir / str(study_id)\n    all_dcm = list(study_path.rglob('*.dcm'))\n\n    if not all_dcm:\n        return np.zeros((num_slices, img_size, img_size), dtype=np.uint8)\n\n    # Group by parent directory (series)\n    from collections import defaultdict\n    series_map = defaultdict(list)\n    for f in all_dcm:\n        series_map[f.parent].append(f)\n\n    # Pick largest series\n    best_series = max(series_map.values(), key=len)\n\n    # Load & sort by InstanceNumber\n    slices = []\n    for f in best_series:\n        try:\n            ds = pydicom.dcmread(str(f))\n            img = ds.pixel_array.astype(np.float32)\n            slope = float(getattr(ds, 'RescaleSlope', 1.0))\n            intercept = float(getattr(ds, 'RescaleIntercept', 0.0))\n            img = img * slope + intercept\n            inst = int(getattr(ds, 'InstanceNumber', 0))\n            slices.append((inst, img))\n        except Exception:\n            continue\n\n    if not slices:\n        return np.zeros((num_slices, img_size, img_size), dtype=np.uint8)\n\n    slices.sort(key=lambda x: x[0])\n    volume = np.stack([s for _, s in slices], axis=0)  # (N, H, W)\n\n    # Percentile normalize\n    lo, hi = np.percentile(volume, 1), np.percentile(volume, 99)\n    volume = np.clip((volume - lo) / (hi - lo + 1e-7), 0, 1) * 255\n    volume = volume.astype(np.uint8)\n\n    # Sample num_slices evenly\n    n = volume.shape[0]\n    if n >= num_slices:\n        idx = np.linspace(0, n - 1, num_slices, dtype=int)\n    else:\n        idx = np.arange(num_slices) % n\n    volume = volume[idx]  # (num_slices, H, W)\n\n    # Resize each slice\n    out = np.zeros((num_slices, img_size, img_size), dtype=np.uint8)\n    for i in range(num_slices):\n        out[i] = cv2.resize(volume[i], (img_size, img_size), interpolation=cv2.INTER_LINEAR)\n\n    return out  # (num_slices, H, W) uint8\n\n\n# Quick sanity check\n# Change 'study_id' -> 'StudyInstanceUID'\nsample_id = train_df['StudyInstanceUID'].iloc[0]\n# Update 'train_images' -> 'train_series'\nimg_dir = COMP_DIR / 'train_series'\nsample_vol = load_volume(sample_id, img_dir, NUM_SLICES, IMG_SIZE)\nprint(f'Volume shape: {sample_vol.shape}  dtype: {sample_vol.dtype}  range: [{sample_vol.min()}, {sample_vol.max()}]')\n\n# Visualize\nfig, axes = plt.subplots(2, 8, figsize=(20, 6))\nfor i, ax in enumerate(axes.flatten()):\n    ax.imshow(sample_vol[i], cmap='bone')\n    ax.set_title(f'Slice {i}', fontsize=8)\n    ax.axis('off')\nplt.suptitle(f'Study: {sample_id}', fontsize=12)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T18:28:13.020454Z","iopub.execute_input":"2026-08-24T18:28:13.020742Z","iopub.status.idle":"2026-08-24T18:28:14.33776Z","shell.execute_reply.started":"2026-08-24T18:28:13.020722Z","shell.execute_reply":"2026-08-24T18:28:14.336944Z"}},"outputs":[],"execution_count":null},{"id":"33fd3deb-1b4c-46a6-b582-8205b254da47","cell_type":"markdown","source":"## 🔄 6. Transforms","metadata":{}},{"id":"f0276a95-0a42-48d7-bf71-17df406921e2","cell_type":"code","source":"def get_train_transforms(img_size=256):\n    return A.Compose([\n        A.RandomResizedCrop(\n            size=(img_size, img_size),  # Updated format for Albumentations 1.4+\n            scale=(0.8, 1.0), \n            p=1.0\n        ),\n        A.HorizontalFlip(p=0.5),\n        A.Rotate(limit=15, border_mode=cv2.BORDER_REFLECT, p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1, rotate_limit=10, border_mode=cv2.BORDER_REFLECT, p=0.3),\n        A.ElasticTransform(alpha=80, sigma=6, p=0.2),\n        A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.5),\n        A.GaussNoise(var_limit=(5.0, 25.0), p=0.3),\n        A.GaussianBlur(blur_limit=(3, 5), p=0.2),\n        A.CLAHE(clip_limit=3.0, p=0.2),\n        A.CoarseDropout(max_holes=6, max_height=img_size // 16, max_width=img_size // 16, p=0.2),\n        A.Normalize(mean=0.0, std=1.0, max_pixel_value=255.0),\n    ])\n\n\ndef get_val_transforms(img_size=256):\n    return A.Compose([\n        A.Resize(height=img_size, width=img_size),\n        A.Normalize(mean=0.0, std=1.0, max_pixel_value=255.0),\n    ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T18:28:14.338678Z","iopub.execute_input":"2026-08-24T18:28:14.339589Z","iopub.status.idle":"2026-08-24T18:28:14.348442Z","shell.execute_reply.started":"2026-08-24T18:28:14.339551Z","shell.execute_reply":"2026-08-24T18:28:14.347596Z"}},"outputs":[],"execution_count":null},{"id":"c9404654-cf82-40f9-9084-5fa8e9bc5a7d","cell_type":"markdown","source":"## 📊 7. Dataset","metadata":{}},{"id":"9364ff07-993a-42a4-b834-8fb92c65d6c2","cell_type":"code","source":"class KneeDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None, tokenizer=None,\n                 num_slices=16, img_size=256, text_max_len=256, mode='train'):\n        self.df = df.reset_index(drop=True)\n        self.img_dir = Path(img_dir)\n        self.transform = transform\n        self.tokenizer = tokenizer\n        self.num_slices = num_slices\n        self.img_size = img_size\n        self.text_max_len = text_max_len\n        self.mode = mode\n        self.has_labels = all(l in df.columns for l in LABEL_NAMES)\n\n    def __len__(self):\n        return len(self.df)\n\n    def _apply_transform(self, volume):\n        \"\"\"Apply 2D transform to each slice. volume: (N, H, W) uint8\"\"\"\n        out = []\n        for i in range(volume.shape[0]):\n            result = self.transform(image=volume[i])\n            out.append(result['image'])\n        return np.stack(out, axis=0)  # (N, H, W) float\n\n    def _get_text(self, report):\n        if self.tokenizer is None or report is None or str(report).strip() == '':\n            return {\n                'input_ids': torch.zeros(self.text_max_len, dtype=torch.long),\n                'attention_mask': torch.zeros(self.text_max_len, dtype=torch.long),\n            }\n        enc = self.tokenizer(\n            str(report), max_length=self.text_max_len,\n            padding='max_length', truncation=True, return_tensors='pt'\n        )\n        return {\n            'input_ids': enc['input_ids'].squeeze(0),\n            'attention_mask': enc['attention_mask'].squeeze(0),\n        }\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        \n        # 1. Update 'study_id' -> 'StudyInstanceUID'\n        study_id = str(row['StudyInstanceUID'])\n        \n        volume = load_volume(study_id, self.img_dir, self.num_slices, self.img_size)\n        \n        if self.transform:\n            volume = self._apply_transform(volume)\n        else:\n            norm = get_val_transforms(self.img_size)\n            volume = self._apply_transform_fixed(volume, norm)\n            \n        image = torch.from_numpy(volume).float()\n        \n        # 2. Update 'report' -> 'Report'\n        report = row.get('Report', None)\n        text = self._get_text(report)\n        \n        if self.has_labels:\n            labels = torch.tensor([float(row[l]) for l in LABEL_NAMES], dtype=torch.float32)\n        else:\n            labels = torch.zeros(NUM_LABELS, dtype=torch.float32)\n            \n        return {\n            'study_id': study_id,\n            'image': image,\n            'input_ids': text['input_ids'],\n            'attention_mask': text['attention_mask'],\n            'labels': labels\n        }\n\n    def _apply_transform_fixed(self, volume, transform):\n        out = []\n        for i in range(volume.shape[0]):\n            result = transform(image=volume[i])\n            out.append(result['image'])\n        return np.stack(out, axis=0)\n\nprint('Dataset class defined ✓')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T18:28:14.349637Z","iopub.execute_input":"2026-08-24T18:28:14.350005Z","iopub.status.idle":"2026-08-24T18:28:14.36617Z","shell.execute_reply.started":"2026-08-24T18:28:14.349985Z","shell.execute_reply":"2026-08-24T18:28:14.365181Z"}},"outputs":[],"execution_count":null},{"id":"a78c40a7-9090-414b-9908-906db12be369","cell_type":"markdown","source":"## 🤖 8. Model Architecture","metadata":{}},{"id":"df095759-868a-4132-9491-a19a0b30f432","cell_type":"code","source":"# ── Vision Encoder ────────────────────────────────────────────────\nclass VisionEncoder(nn.Module):\n    def __init__(self, backbone=VISION_BACKBONE, num_slices=NUM_SLICES,\n                 out_dim=VISION_OUT_DIM, pretrained=True, drop_rate=0.3):\n        super().__init__()\n        self.backbone = timm.create_model(\n            backbone, pretrained=pretrained, num_classes=0,\n            drop_rate=drop_rate, drop_path_rate=0.2\n        )\n        self._adapt_first_conv(num_slices)\n        feat_dim = self.backbone.num_features\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.proj = nn.Sequential(\n            nn.Flatten(), nn.LayerNorm(feat_dim),\n            nn.Linear(feat_dim, out_dim), nn.GELU(), nn.Dropout(drop_rate)\n        )\n\n    def _adapt_first_conv(self, num_slices):\n        first_conv, first_name = None, None\n        for name, m in self.backbone.named_modules():\n            if isinstance(m, nn.Conv2d):\n                first_conv, first_name = m, name\n                break\n        if first_conv is None or first_conv.in_channels == num_slices:\n            return\n        new_conv = nn.Conv2d(num_slices, first_conv.out_channels,\n                             first_conv.kernel_size, first_conv.stride,\n                             first_conv.padding, bias=first_conv.bias is not None)\n        with torch.no_grad():\n            avg_w = first_conv.weight.mean(dim=1, keepdim=True)\n            new_conv.weight.data = avg_w.repeat(1, num_slices, 1, 1) / num_slices\n            if first_conv.bias is not None:\n                new_conv.bias.data = first_conv.bias.data.clone()\n        parts = first_name.split('.')\n        parent = self.backbone\n        for p in parts[:-1]: parent = getattr(parent, p)\n        setattr(parent, parts[-1], new_conv)\n\n    def forward(self, x):\n        feats = self.backbone.forward_features(x)\n        if feats.dim() == 4: feats = self.pool(feats)\n        return self.proj(feats)\n\n\n# ── Text Encoder ─────────────────────────────────────────────────\nclass TextEncoder(nn.Module):\n    def __init__(self, model_name=TEXT_BACKBONE, out_dim=TEXT_OUT_DIM,\n                 freeze_layers=FREEZE_TEXT_LAYERS):\n        super().__init__()\n        self.backbone = AutoModel.from_pretrained(model_name)\n        hidden = self.backbone.config.hidden_size\n        for p in self.backbone.embeddings.parameters(): p.requires_grad = False\n        for i, layer in enumerate(self.backbone.encoder.layer):\n            if i < freeze_layers:\n                for p in layer.parameters(): p.requires_grad = False\n        self.proj = nn.Sequential(\n            nn.LayerNorm(hidden), nn.Linear(hidden, out_dim), nn.GELU(), nn.Dropout(0.1)\n        )\n\n    def forward(self, input_ids, attention_mask):\n        out = self.backbone(input_ids=input_ids, attention_mask=attention_mask)\n        tok = out.last_hidden_state  # (B, L, H)\n        mask = attention_mask.unsqueeze(-1).float()\n        pooled = (tok * mask).sum(1) / mask.sum(1).clamp(min=1e-9)\n        return self.proj(pooled)\n\n\n# ── Fusion Model ─────────────────────────────────────────────────\nclass MultimodalKneeModel(nn.Module):\n    def __init__(self, use_text=USE_TEXT):\n        super().__init__()\n        self.use_text = use_text\n        self.vision = VisionEncoder()\n        if use_text:\n            self.text = TextEncoder()\n            in_dim = VISION_OUT_DIM + TEXT_OUT_DIM\n        else:\n            self.text = None\n            in_dim = VISION_OUT_DIM\n        self.head = nn.Sequential(\n            nn.LayerNorm(in_dim),\n            nn.Linear(in_dim, FUSION_DIM), nn.GELU(), nn.Dropout(0.4),\n            nn.Linear(FUSION_DIM, FUSION_DIM // 2), nn.GELU(), nn.Dropout(0.2),\n            nn.Linear(FUSION_DIM // 2, NUM_LABELS)\n        )\n\n    def forward(self, image, input_ids=None, attention_mask=None):\n        v = self.vision(image)\n        if self.use_text and self.text is not None and input_ids is not None:\n            has_text = (attention_mask.sum(-1) > 0).float().unsqueeze(-1)\n            t = self.text(input_ids, attention_mask) * has_text\n            feat = torch.cat([v, t], dim=-1)\n        else:\n            feat = v\n        logits = self.head(feat)\n        return {'logits': logits, 'probs': torch.sigmoid(logits)}\n\n    def param_groups(self):\n        return [\n            {'params': list(self.vision.backbone.parameters()), 'lr': LR_VISION},\n            {'params': list(self.vision.proj.parameters()), 'lr': LR_HEAD},\n            {'params': list(self.head.parameters()), 'lr': LR_HEAD},\n        ] + ([\n            {'params': list(self.text.backbone.parameters()), 'lr': LR_TEXT},\n            {'params': list(self.text.proj.parameters()), 'lr': LR_HEAD},\n        ] if self.use_text and self.text else [])\n\n\n# Quick model test\nmodel = MultimodalKneeModel().to(device)\ntotal = sum(p.numel() for p in model.parameters())\ntrainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f'Model: {total/1e6:.1f}M total | {trainable/1e6:.1f}M trainable')\n\ndummy_img = torch.randn(2, NUM_SLICES, IMG_SIZE, IMG_SIZE).to(device)\ndummy_ids = torch.ones(2, 64, dtype=torch.long).to(device)\ndummy_mask = torch.ones(2, 64, dtype=torch.long).to(device)\nwith torch.no_grad():\n    out = model(dummy_img, dummy_ids, dummy_mask)\nprint(f'Forward pass OK! logits: {out[\"logits\"].shape}, probs: {out[\"probs\"].shape}')\ndel model, dummy_img, dummy_ids, dummy_mask; gc.collect(); torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T18:28:14.367311Z","iopub.execute_input":"2026-08-24T18:28:14.367614Z","iopub.status.idle":"2026-08-24T18:28:27.824556Z","shell.execute_reply.started":"2026-08-24T18:28:14.367594Z","shell.execute_reply":"2026-08-24T18:28:27.823705Z"}},"outputs":[],"execution_count":null},{"id":"df06542a-a3a1-4ff1-b3c8-6bc964df763f","cell_type":"markdown","source":"## 📐 9. Loss, Optimizer & Scheduler","metadata":{}},{"id":"ef45c8a1-c8d0-4bc5-be2b-9bcacc11abaf","cell_type":"code","source":"class LabelSmoothBCE(nn.Module):\n    def __init__(self, smoothing=0.05):\n        super().__init__()\n        self.smoothing = smoothing\n    def forward(self, logits, targets):\n        targets = targets * (1 - self.smoothing) + 0.5 * self.smoothing\n        return F.binary_cross_entropy_with_logits(logits, targets)\n\ndef get_scheduler(optimizer, num_warmup_steps, num_training_steps):\n    from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR\n    warmup = LinearLR(optimizer, start_factor=0.01, end_factor=1.0, total_iters=num_warmup_steps)\n    cosine = CosineAnnealingLR(optimizer, T_max=num_training_steps - num_warmup_steps, eta_min=1e-7)\n    return SequentialLR(optimizer, schedulers=[warmup, cosine], milestones=[num_warmup_steps])\n\ndef compute_auc(y_true, y_pred):\n    per_label = {}\n    aucs = []\n    for i, name in enumerate(LABEL_NAMES):\n        col_true = y_true[:, i]\n        col_pred = y_pred[:, i]\n        if len(np.unique(col_true)) < 2:\n            per_label[name] = 0.5; aucs.append(0.5)\n        else:\n            auc = roc_auc_score(col_true, col_pred)\n            per_label[name] = float(auc); aucs.append(auc)\n    return float(np.mean(aucs)), per_label\n\nprint('Loss & metrics defined ✓')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T18:28:27.825551Z","iopub.execute_input":"2026-08-24T18:28:27.825761Z","iopub.status.idle":"2026-08-24T18:28:27.833815Z","shell.execute_reply.started":"2026-08-24T18:28:27.825742Z","shell.execute_reply":"2026-08-24T18:28:27.832906Z"}},"outputs":[],"execution_count":null},{"id":"ec5122dd-1403-499d-ae47-51d7360fc0bf","cell_type":"markdown","source":"## 🏋️ 10. Training Loop","metadata":{}},{"id":"d24a5a78-b46d-4329-b6f8-e279e85fbe85","cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, scheduler, criterion, scaler, epoch):\n    model.train()\n    losses = []\n    optimizer.zero_grad()\n    pbar = tqdm(enumerate(loader), total=len(loader), desc=f'Train E{epoch}')\n    for step, batch in pbar:\n        imgs  = batch['image'].to(device, non_blocking=True)\n        ids   = batch['input_ids'].to(device, non_blocking=True)\n        mask  = batch['attention_mask'].to(device, non_blocking=True)\n        lbls  = batch['labels'].to(device, non_blocking=True)\n        with autocast(enabled=AMP):\n            out  = model(imgs, ids, mask)\n            loss = criterion(out['logits'], lbls) / GRAD_ACCUM\n        scaler.scale(loss).backward()\n        if (step + 1) % GRAD_ACCUM == 0:\n            scaler.unscale_(optimizer)\n            nn.utils.clip_grad_norm_(model.parameters(), MAX_GRAD_NORM)\n            scaler.step(optimizer); scaler.update()\n            scheduler.step(); optimizer.zero_grad()\n        losses.append(loss.item() * GRAD_ACCUM)\n        pbar.set_postfix({'loss': f'{np.mean(losses[-20:]):.4f}',\n                          'lr': f'{optimizer.param_groups[0][\"lr\"]:.2e}'})\n    return np.mean(losses)\n\n@torch.no_grad()\ndef validate(model, loader, criterion):\n    model.eval()\n    losses, all_labels, all_probs = [], [], []\n    for batch in tqdm(loader, desc='Val', leave=False):\n        imgs = batch['image'].to(device, non_blocking=True)\n        ids  = batch['input_ids'].to(device, non_blocking=True)\n        mask = batch['attention_mask'].to(device, non_blocking=True)\n        lbls = batch['labels'].to(device, non_blocking=True)\n        with autocast(enabled=AMP):\n            out  = model(imgs, ids, mask)\n            loss = criterion(out['logits'], lbls)\n        losses.append(loss.item())\n        all_labels.append(lbls.cpu().numpy())\n        all_probs.append(out['probs'].cpu().numpy())\n    y_true = np.concatenate(all_labels)\n    y_pred = np.concatenate(all_probs)\n    macro, per = compute_auc(y_true, y_pred)\n    return np.mean(losses), macro, per\n\nprint('Training loop defined ✓')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T18:28:27.834886Z","iopub.execute_input":"2026-08-24T18:28:27.83556Z","iopub.status.idle":"2026-08-24T18:28:27.857225Z","shell.execute_reply.started":"2026-08-24T18:28:27.835532Z","shell.execute_reply":"2026-08-24T18:28:27.856693Z"}},"outputs":[],"execution_count":null},{"id":"920abd04-1648-4e14-84d4-53f1d01f7833","cell_type":"markdown","source":"## 🚀 11. Cross-Validation Training","metadata":{}},{"id":"64c96d2b-248b-4caa-89bb-8f087fb74a15","cell_type":"code","source":"# CV split\ntrain_df['_strat'] = train_df[LABEL_NAMES].sum(axis=1).astype(int).clip(0, 4)\nkfold = StratifiedKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\nsplits = list(kfold.split(train_df, train_df['_strat']))\n\n# Tokenizer\ntokenizer = AutoTokenizer.from_pretrained(TEXT_BACKBONE) if USE_TEXT else None\n\nfold_aucs = []\ntrain_img_dir = COMP_DIR / 'train_images'\n\nfolds_to_run = [TRAIN_FOLD - 1] if TRAIN_FOLD else range(N_FOLDS)\n\nfor fold_idx in folds_to_run:\n    fold = fold_idx + 1\n    print(f'\\n{\"=\"*60}')\n    print(f'  FOLD {fold}/{N_FOLDS}')\n    print(f'{\"=\"*60}')\n\n    train_idx, val_idx = splits[fold_idx]\n    fold_train = train_df.iloc[train_idx].drop(columns=['_strat'])\n    fold_val   = train_df.iloc[val_idx].drop(columns=['_strat'])\n\n    # DataLoaders\n    train_ds = KneeDataset(fold_train, train_img_dir,\n                           transform=get_train_transforms(IMG_SIZE),\n                           tokenizer=tokenizer, num_slices=NUM_SLICES,\n                           img_size=IMG_SIZE, text_max_len=TEXT_MAX_LEN, mode='train')\n    val_ds   = KneeDataset(fold_val, train_img_dir,\n                           transform=get_val_transforms(IMG_SIZE),\n                           tokenizer=tokenizer, num_slices=NUM_SLICES,\n                           img_size=IMG_SIZE, text_max_len=TEXT_MAX_LEN, mode='val')\n\n    train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True,\n                              num_workers=NUM_WORKERS, pin_memory=True, drop_last=True)\n    val_loader   = DataLoader(val_ds, batch_size=BATCH_SIZE*2, shuffle=False,\n                              num_workers=NUM_WORKERS, pin_memory=True)\n\n    # Model\n    model = MultimodalKneeModel(use_text=USE_TEXT).to(device)\n\n    # Optimizer\n    optimizer = torch.optim.AdamW(model.param_groups(), weight_decay=WEIGHT_DECAY)\n    criterion = LabelSmoothBCE(LABEL_SMOOTHING)\n    scaler    = GradScaler(enabled=AMP)\n\n    total_steps  = EPOCHS * (len(train_loader) // GRAD_ACCUM)\n    warmup_steps = WARMUP_EPOCHS * (len(train_loader) // GRAD_ACCUM)\n    scheduler    = get_scheduler(optimizer, warmup_steps, total_steps)\n\n    best_auc, best_ckpt = 0.0, None\n\n    for epoch in range(1, EPOCHS + 1):\n        t0 = time.time()\n        train_loss = train_one_epoch(model, train_loader, optimizer, scheduler, criterion, scaler, epoch)\n        val_loss, macro_auc, per_label = validate(model, val_loader, criterion)\n        elapsed = time.time() - t0\n\n        print(f'\\nEpoch {epoch:2d}/{EPOCHS} ({elapsed/60:.1f} min) | '\n              f'train_loss={train_loss:.4f} | val_loss={val_loss:.4f} | '\n              f'macro_AUC={macro_auc:.4f} {\"★ NEW BEST\" if macro_auc > best_auc else \"\"}')\n        for lbl, auc in per_label.items():\n            bar = '█' * int(auc * 15) + '░' * (15 - int(auc * 15))\n            print(f'  {lbl:<20} [{bar}] {auc:.4f}')\n\n        if macro_auc > best_auc:\n            best_auc = macro_auc\n            best_ckpt = CKPT_DIR / f'fold{fold}_best.pt'\n            torch.save(model.state_dict(), best_ckpt)\n            print(f'  → Saved checkpoint: {best_ckpt}')\n\n    fold_aucs.append(best_auc)\n    print(f'\\nFold {fold} Best AUC: {best_auc:.4f}')\n    del model, optimizer, scheduler, scaler, train_ds, val_ds\n    gc.collect(); torch.cuda.empty_cache()\n\nprint(f'\\n{\"=\"*60}')\nprint(f'  CV Results: {[f\"{a:.4f}\" for a in fold_aucs]}')\nprint(f'  Mean CV AUC: {np.mean(fold_aucs):.4f} ± {np.std(fold_aucs):.4f}')\nprint(f'{\"=\"*60}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T18:28:27.858021Z","iopub.execute_input":"2026-08-24T18:28:27.85846Z","iopub.status.idle":"2026-08-24T21:10:53.257643Z","shell.execute_reply.started":"2026-08-24T18:28:27.85844Z","shell.execute_reply":"2026-08-24T21:10:53.256916Z"}},"outputs":[],"execution_count":null},{"id":"4e383ec0-bdab-4c55-baae-11eea7466cbc","cell_type":"markdown","source":"## 📝 12. Inference & Submission","metadata":{}},{"id":"4a3815de-84f8-4933-8c46-c38f53ec93d4","cell_type":"code","source":"test_img_dir = COMP_DIR / 'test_images'\n\ntest_ds = KneeDataset(test_df, test_img_dir,\n                      transform=get_val_transforms(IMG_SIZE),\n                      tokenizer=tokenizer, num_slices=NUM_SLICES,\n                      img_size=IMG_SIZE, text_max_len=TEXT_MAX_LEN, mode='test')\ntest_loader = DataLoader(test_ds, batch_size=BATCH_SIZE*2, shuffle=False,\n                         num_workers=NUM_WORKERS, pin_memory=True)\n\nall_fold_probs = []\n\nfor ckpt_path in sorted(CKPT_DIR.glob('fold*_best.pt')):\n    print(f'Loading: {ckpt_path.name}')\n    model = MultimodalKneeModel(use_text=USE_TEXT).to(device)\n    model.load_state_dict(torch.load(str(ckpt_path), map_location=device))\n    model.eval()\n\n    fold_probs = []\n    # Standard pass\n    with torch.no_grad():\n        for batch in tqdm(test_loader, desc='Inference', leave=False):\n            imgs = batch['image'].to(device)\n            ids  = batch['input_ids'].to(device)\n            mask = batch['attention_mask'].to(device)\n            with autocast(enabled=AMP):\n                out = model(imgs, ids, mask)\n            fold_probs.append(out['probs'].cpu().numpy())\n\n    # TTA: horizontal flip\n    tta_probs = []\n    tta_transform = get_val_transforms(IMG_SIZE)  # already defined\n    with torch.no_grad():\n        for batch in tqdm(test_loader, desc='TTA flip', leave=False):\n            imgs = batch['image'].to(device)\n            imgs_flip = torch.flip(imgs, dims=[-1])  # horizontal flip\n            ids  = batch['input_ids'].to(device)\n            mask = batch['attention_mask'].to(device)\n            with autocast(enabled=AMP):\n                out = model(imgs_flip, ids, mask)\n            tta_probs.append(out['probs'].cpu().numpy())\n\n    fold_preds = (np.concatenate(fold_probs) + np.concatenate(tta_probs)) / 2\n    all_fold_probs.append(fold_preds)\n\n    del model; gc.collect(); torch.cuda.empty_cache()\n\n# Ensemble average\nensemble_preds = np.mean(all_fold_probs, axis=0)\nprint(f'\\nEnsemble shape: {ensemble_preds.shape}')\n\n# Build submission\nsubmission = pd.DataFrame({'StudyInstanceUID': test_df['StudyInstanceUID']})\nfor i, lbl in enumerate(LABEL_NAMES):\n    submission[lbl] = ensemble_preds[:, i]\n\nsub_path = OUTPUT_DIR / 'submission.csv'\nsubmission.to_csv(sub_path, index=False)\nprint(f'\\n✅ Submission saved: {sub_path}')\nprint(f'Shape: {submission.shape}')\nsubmission.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T21:14:53.051939Z","iopub.execute_input":"2026-08-24T21:14:53.052564Z","iopub.status.idle":"2026-08-24T21:15:16.020668Z","shell.execute_reply.started":"2026-08-24T21:14:53.052525Z","shell.execute_reply":"2026-08-24T21:15:16.019967Z"}},"outputs":[],"execution_count":null},{"id":"80c33e83-755a-401e-b456-4ba79e9aacb0","cell_type":"code","source":"# Final stats\nprint('Prediction summary:')\nprint(submission[LABEL_NAMES].describe().T[['mean','std','min','max']].to_string())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-24T21:22:45.246247Z","iopub.execute_input":"2026-08-24T21:22:45.246696Z","iopub.status.idle":"2026-08-24T21:22:45.274646Z","shell.execute_reply.started":"2026-08-24T21:22:45.246666Z","shell.execute_reply":"2026-08-24T21:22:45.27382Z"}},"outputs":[],"execution_count":null},{"id":"9a422b71-70ec-479a-bd6d-d5a6ed8670f0","cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}