{"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":"a2c9dea0-f0c9-43f1-8820-7155def01117","cell_type":"markdown","source":"# RSNA Knee Abnormality Detection\n## Multi-label MRI Classification\n\nPipeline:\n- Load competition data\n- Read CSV metadata\n- Multi-label stratified train/validation split\n- DICOM study loader\n- Sample multiple slices per study\n- ConvNeXt-Tiny encoder\n- Attention pooling across slices\n- Weighted BCE loss\n- Macro ROC-AUC evaluation\n- Test inference\n- `submission.csv`\n\nThis notebook is designed for Kaggle GPU notebooks.","metadata":{}},{"id":"f8f95a26-07f8-41cb-b097-65ff979717eb","cell_type":"code","source":"# Cell 1: Install dependencies\n\n!pip install -q timm pydicom pylibjpeg pylibjpeg-libjpeg pylibjpeg-openjpeg iterative-stratification\n\nprint(\"✅ Cell 1 complete: dependencies installed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:29:00.317189Z","iopub.execute_input":"2026-08-09T12:29:00.318051Z","iopub.status.idle":"2026-08-09T12:29:04.183534Z","shell.execute_reply.started":"2026-08-09T12:29:00.31802Z","shell.execute_reply":"2026-08-09T12:29:04.182331Z"}},"outputs":[],"execution_count":null},{"id":"8fdf3218-2a2f-4c21-9a80-27acab5b51fe","cell_type":"code","source":"# Cell 2: Imports + config\n\nimport os\nimport gc\nimport random\nimport warnings\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport pydicom\nimport cv2\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\n\nimport timm\n\nfrom sklearn.metrics import roc_auc_score\nfrom iterstrat.ml_stratifiers import MultilabelStratifiedShuffleSplit\n\nwarnings.filterwarnings(\"ignore\")\n\nSEED = 42\nIMAGE_SIZE = 224\nN_SLICES = 16\nBATCH_SIZE = 2\nEPOCHS = 10\nLR = 1e-4\nWEIGHT_DECAY = 1e-4\nNUM_WORKERS = 2\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\ndef seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n\nseed_everything(SEED)\n\nprint(\"✅ Cell 2 complete\")\nprint(\"Device       :\", device)\nprint(\"Torch        :\", torch.__version__)\nprint(\"Image size   :\", IMAGE_SIZE)\nprint(\"Slices/study :\", N_SLICES)\nprint(\"Batch size   :\", BATCH_SIZE)\nprint(\"Epochs       :\", EPOCHS)\nprint(\"LR           :\", LR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:29:07.363142Z","iopub.execute_input":"2026-08-09T12:29:07.364189Z","iopub.status.idle":"2026-08-09T12:29:07.374661Z","shell.execute_reply.started":"2026-08-09T12:29:07.364158Z","shell.execute_reply":"2026-08-09T12:29:07.373858Z"}},"outputs":[],"execution_count":null},{"id":"cd6c8684-d4f8-445b-9921-4c7fc223bfd2","cell_type":"code","source":"# Cell 3: Setup DATA_DIR\n\nCOMPETITION = \"rsna-knee-abnormality-detection\"\nKAGGLE_INPUT = Path(f\"/kaggle/input/{COMPETITION}\")\n\nif KAGGLE_INPUT.exists():\n    DATA_DIR = KAGGLE_INPUT\n    source = \"Kaggle Input\"\nelse:\n    print(\"📦 Dataset not found in /kaggle/input\")\n    print(\"Trying kagglehub download...\")\n\n    !pip install -q -U kagglehub\n    import kagglehub\n\n    try:\n        DATA_DIR = Path(kagglehub.competition_download(COMPETITION))\n    except Exception:\n        print(\"🔐 Kaggle authentication required\")\n        kagglehub.login()\n        DATA_DIR = Path(kagglehub.competition_download(COMPETITION))\n\n    source = \"kagglehub\"\n\nprint(\"✅ Cell 3 complete\")\nprint(\"Dataset source:\", source)\nprint(\"DATA_DIR      :\", DATA_DIR)\nprint(\"Exists        :\", DATA_DIR.exists())\n\nitems = sorted(DATA_DIR.iterdir())\nprint(\"Top-level item count:\", len(items))\nfor p in items[:20]:\n    print(\" -\", p.name)\n\nif len(items) > 20:\n    print(f\" ... and {len(items) - 20} more\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:29:10.037744Z","iopub.execute_input":"2026-08-09T12:29:10.038179Z","iopub.status.idle":"2026-08-09T12:29:15.98601Z","shell.execute_reply.started":"2026-08-09T12:29:10.038151Z","shell.execute_reply":"2026-08-09T12:29:15.98484Z"}},"outputs":[],"execution_count":null},{"id":"2f932eea-f604-40e1-a7b6-b6f6f5f89f90","cell_type":"code","source":"# Cell 4: Load CSVs + split labeled/unlabeled\n\nTARGETS = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\"\n]\n\ncsv_names = [\n    \"train.csv\",\n    \"train_series.csv\",\n    \"test.csv\",\n    \"test_series.csv\",\n    \"sample_submission.csv\"\n]\n\nframes = {}\n\nfor name in tqdm(csv_names, desc=\"Loading CSV\", unit=\"file\"):\n    file_path = DATA_DIR / name\n\n    if not file_path.exists():\n        raise FileNotFoundError(f\"Missing required file: {file_path}\")\n\n    frames[name] = pd.read_csv(file_path)\n\ntrain_df = frames[\"train.csv\"]\ntrain_series_df = frames[\"train_series.csv\"]\ntest_df = frames[\"test.csv\"]\ntest_series_df = frames[\"test_series.csv\"]\nsample_submission = frames[\"sample_submission.csv\"]\n\nmissing_targets = [c for c in TARGETS if c not in train_df.columns]\n\nif missing_targets:\n    raise ValueError(f\"Missing target columns: {missing_targets}\")\n\nlabeled_mask = train_df[TARGETS].notna().all(axis=1)\n\nlabeled_df = train_df.loc[labeled_mask].copy().reset_index(drop=True)\nunlabeled_df = train_df.loc[~labeled_mask].copy().reset_index(drop=True)\n\nprint(\"✅ Cell 4 complete\")\nprint(\"train.csv             :\", train_df.shape)\nprint(\"train_series.csv      :\", train_series_df.shape)\nprint(\"test.csv              :\", test_df.shape)\nprint(\"test_series.csv       :\", test_series_df.shape)\nprint(\"sample_submission.csv :\", sample_submission.shape)\nprint()\nprint(f\"All train studies : {len(train_df):,}\")\nprint(f\"Labeled studies   : {len(labeled_df):,}\")\nprint(f\"Unlabeled studies : {len(unlabeled_df):,}\")\nprint(f\"Targets           : {len(TARGETS)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:29:19.197038Z","iopub.execute_input":"2026-08-09T12:29:19.197551Z","iopub.status.idle":"2026-08-09T12:29:19.524328Z","shell.execute_reply.started":"2026-08-09T12:29:19.197515Z","shell.execute_reply":"2026-08-09T12:29:19.523339Z"}},"outputs":[],"execution_count":null},{"id":"a7626abc-b166-4ecd-b073-1d64685c8db1","cell_type":"code","source":"# Cell 5: Label distribution\n\nstats = []\n\nfor target in TARGETS:\n    pos = int(labeled_df[target].sum())\n    neg = int((labeled_df[target] == 0).sum())\n    total = pos + neg\n    prevalence = 100.0 * pos / max(total, 1)\n\n    stats.append({\n        \"Target\": target,\n        \"Positive\": pos,\n        \"Negative\": neg,\n        \"Prevalence_%\": prevalence\n    })\n\nlabel_stats = pd.DataFrame(stats)\n\nprint(\"✅ Cell 5 complete: label distribution\")\ndisplay(label_stats)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:29:23.4563Z","iopub.execute_input":"2026-08-09T12:29:23.45703Z","iopub.status.idle":"2026-08-09T12:29:23.489835Z","shell.execute_reply.started":"2026-08-09T12:29:23.456997Z","shell.execute_reply":"2026-08-09T12:29:23.489021Z"}},"outputs":[],"execution_count":null},{"id":"c5e4df2e-c103-4af3-b560-96f65eb813f7","cell_type":"code","source":"# Cell 6: Multilabel stratified train/validation split\n\nX = labeled_df[\"StudyInstanceUID\"].values\ny = labeled_df[TARGETS].values.astype(int)\n\nsplitter = MultilabelStratifiedShuffleSplit(\n    n_splits=1,\n    test_size=0.20,\n    random_state=SEED\n)\n\ntrain_idx, val_idx = next(splitter.split(X, y))\n\ntrain_split = labeled_df.iloc[train_idx].reset_index(drop=True)\nval_split = labeled_df.iloc[val_idx].reset_index(drop=True)\n\nval_check = []\n\nfor t in TARGETS:\n    pos = int(val_split[t].sum())\n    neg = int((val_split[t] == 0).sum())\n    auc_ok = (pos > 0 and neg > 0)\n\n    val_check.append({\n        \"Target\": t,\n        \"Val_Positive\": pos,\n        \"Val_Negative\": neg,\n        \"AUC_Valid\": auc_ok\n    })\n\nval_check_df = pd.DataFrame(val_check)\nvalid_auc_labels = int(val_check_df[\"AUC_Valid\"].sum())\n\nprint(\"✅ Cell 6 complete\")\nprint(\"Train studies :\", len(train_split))\nprint(\"Val studies   :\", len(val_split))\nprint(f\"Valid AUC labels: {valid_auc_labels}/12\")\ndisplay(val_check_df)\n\nif valid_auc_labels < 12:\n    print(\"⚠️ Some labels still cannot produce ROC-AUC on this split.\")\n    print(\"   Metric code will safely ignore invalid labels instead of crashing.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:29:30.275404Z","iopub.execute_input":"2026-08-09T12:29:30.276241Z","iopub.status.idle":"2026-08-09T12:29:30.302166Z","shell.execute_reply.started":"2026-08-09T12:29:30.276208Z","shell.execute_reply":"2026-08-09T12:29:30.301156Z"}},"outputs":[],"execution_count":null},{"id":"3c783349-9e43-4513-9e3a-534357da0997","cell_type":"code","source":"# Cell 7: Build study -> series metadata lookup\n\ndef make_series_lookup(series_df):\n    lookup = {}\n\n    for study_uid, group in tqdm(\n        series_df.groupby(\"StudyInstanceUID\"),\n        total=series_df[\"StudyInstanceUID\"].nunique(),\n        desc=\"Building series lookup\",\n        leave=False\n    ):\n        records = []\n\n        for _, row in group.iterrows():\n            records.append({\n                \"SeriesInstanceUID\": str(row[\"SeriesInstanceUID\"]),\n                \"Anatomical_Plane\": row.get(\"Anatomical_Plane\", None),\n                \"Fluid_Sensitive\": row.get(\"Fluid_Sensitive\", None),\n                \"Fat_Suppression\": row.get(\"Fat_Suppression\", None),\n            })\n\n        lookup[str(study_uid)] = records\n\n    return lookup\n\ntrain_series_lookup = make_series_lookup(train_series_df)\ntest_series_lookup = make_series_lookup(test_series_df)\n\nprint(\"✅ Cell 7 complete\")\nprint(\"Train study lookup:\", len(train_series_lookup))\nprint(\"Test study lookup :\", len(test_series_lookup))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:29:34.633019Z","iopub.execute_input":"2026-08-09T12:29:34.633504Z","iopub.status.idle":"2026-08-09T12:29:36.313232Z","shell.execute_reply.started":"2026-08-09T12:29:34.633423Z","shell.execute_reply":"2026-08-09T12:29:36.312155Z"}},"outputs":[],"execution_count":null},{"id":"c4c2635b-b0ff-4c87-aa77-89fb3f743206","cell_type":"code","source":"# Cell 8: DICOM helpers\n\ndef dicom_sort_key(path):\n    try:\n        dcm = pydicom.dcmread(\n            str(path),\n            stop_before_pixels=True,\n            force=True\n        )\n\n        if hasattr(dcm, \"InstanceNumber\"):\n            return float(dcm.InstanceNumber)\n\n        if hasattr(dcm, \"ImagePositionPatient\"):\n            return float(dcm.ImagePositionPatient[-1])\n\n    except Exception:\n        pass\n\n    return path.name\n\n\ndef read_dicom(path, image_size=224):\n    dcm = pydicom.dcmread(str(path), force=True)\n    img = dcm.pixel_array.astype(np.float32)\n\n    slope = float(getattr(dcm, \"RescaleSlope\", 1.0))\n    intercept = float(getattr(dcm, \"RescaleIntercept\", 0.0))\n    img = img * slope + intercept\n\n    lo, hi = np.percentile(img, [1, 99])\n\n    if hi > lo:\n        img = np.clip(img, lo, hi)\n        img = (img - lo) / (hi - lo)\n    else:\n        img = np.zeros_like(img, dtype=np.float32)\n\n    img = cv2.resize(\n        img,\n        (image_size, image_size),\n        interpolation=cv2.INTER_AREA\n    )\n\n    return img.astype(np.float32)\n\n\ndef evenly_sample(items, n):\n    if len(items) == 0:\n        return []\n\n    idx = np.linspace(\n        0,\n        len(items) - 1,\n        n\n    ).round().astype(int)\n\n    return [items[i] for i in idx]\n\nprint(\"✅ Cell 8 complete: DICOM helper functions ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:29:43.554336Z","iopub.execute_input":"2026-08-09T12:29:43.554781Z","iopub.status.idle":"2026-08-09T12:29:43.564071Z","shell.execute_reply.started":"2026-08-09T12:29:43.554751Z","shell.execute_reply":"2026-08-09T12:29:43.563195Z"}},"outputs":[],"execution_count":null},{"id":"32e56f6b-e9d7-4f5b-b05f-87834e60f9c2","cell_type":"code","source":"# Cell 9: Study-level slice selection\n\nPLANE_ORDER = [\"Sagittal\", \"Coronal\", \"Axial\"]\n\ndef get_study_slice_paths(\n    study_uid,\n    split=\"train\",\n    total_slices=N_SLICES\n):\n    if split == \"train\":\n        root = DATA_DIR / \"train_series\"\n        lookup = train_series_lookup\n    else:\n        root = DATA_DIR / \"test_series\"\n        lookup = test_series_lookup\n\n    study_uid = str(study_uid)\n    series_records = lookup.get(study_uid, [])\n\n    if not series_records:\n        return []\n\n    # Prefer fluid-sensitive and fat-suppressed series\n    scored = []\n\n    for rec in series_records:\n        score = 0\n\n        if rec.get(\"Fluid_Sensitive\") == 1:\n            score += 2\n\n        if rec.get(\"Fat_Suppression\") == 1:\n            score += 1\n\n        plane = rec.get(\"Anatomical_Plane\")\n\n        scored.append((score, plane, rec))\n\n    # Best series per plane\n    selected_series = []\n\n    for plane in PLANE_ORDER:\n        candidates = [\n            x for x in scored\n            if str(x[1]) == plane\n        ]\n\n        if candidates:\n            candidates.sort(key=lambda x: x[0], reverse=True)\n            selected_series.append(candidates[0][2])\n\n    # Fallback if plane metadata is missing\n    if len(selected_series) == 0:\n        scored.sort(key=lambda x: x[0], reverse=True)\n        selected_series = [x[2] for x in scored[:3]]\n\n    per_series = max(1, total_slices // max(len(selected_series), 1))\n    paths = []\n\n    for rec in selected_series:\n        series_uid = rec[\"SeriesInstanceUID\"]\n        series_dir = root / study_uid / series_uid\n\n        if not series_dir.exists():\n            continue\n\n        files = sorted(\n            series_dir.glob(\"*.dcm\"),\n            key=dicom_sort_key\n        )\n\n        paths.extend(\n            evenly_sample(files, per_series)\n        )\n\n    # Fill/trim to fixed number\n    if len(paths) > total_slices:\n        paths = evenly_sample(paths, total_slices)\n\n    elif len(paths) > 0 and len(paths) < total_slices:\n        paths = evenly_sample(paths, total_slices)\n\n    return paths\n\nprint(\"✅ Cell 9 complete: study slice selector ready\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:29:46.137807Z","iopub.execute_input":"2026-08-09T12:29:46.13809Z","iopub.status.idle":"2026-08-09T12:29:46.149882Z","shell.execute_reply.started":"2026-08-09T12:29:46.138069Z","shell.execute_reply":"2026-08-09T12:29:46.14916Z"}},"outputs":[],"execution_count":null},{"id":"34f9cc0c-6e74-4900-af8c-503536ebd07c","cell_type":"code","source":"# Cell 10: Dataset class\n\nclass KneeStudyDataset(Dataset):\n\n    def __init__(\n        self,\n        df,\n        split=\"train\",\n        targets=None,\n        n_slices=16,\n        image_size=224\n    ):\n        self.df = df.reset_index(drop=True)\n        self.split = split\n        self.targets = targets\n        self.n_slices = n_slices\n        self.image_size = image_size\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        study_uid = str(row[\"StudyInstanceUID\"])\n\n        paths = get_study_slice_paths(\n            study_uid,\n            split=self.split,\n            total_slices=self.n_slices\n        )\n\n        images = []\n\n        for p in paths:\n            try:\n                img = read_dicom(\n                    p,\n                    image_size=self.image_size\n                )\n            except Exception:\n                img = np.zeros(\n                    (self.image_size, self.image_size),\n                    dtype=np.float32\n                )\n\n            images.append(img)\n\n        if len(images) == 0:\n            images = [\n                np.zeros(\n                    (self.image_size, self.image_size),\n                    dtype=np.float32\n                )\n                for _ in range(self.n_slices)\n            ]\n\n        while len(images) < self.n_slices:\n            images.append(images[-1].copy())\n\n        images = images[:self.n_slices]\n\n        x = np.stack(images, axis=0)\n        x = torch.from_numpy(x).unsqueeze(1)  # [S, 1, H, W]\n\n        if self.targets is None:\n            return x, study_uid\n\n        y = row[self.targets].values.astype(np.float32)\n        y = torch.tensor(y, dtype=torch.float32)\n\n        return x, y\n\nprint(\"✅ Cell 10 complete: Dataset class ready\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:29:50.360108Z","iopub.execute_input":"2026-08-09T12:29:50.360848Z","iopub.status.idle":"2026-08-09T12:29:50.36993Z","shell.execute_reply.started":"2026-08-09T12:29:50.360821Z","shell.execute_reply":"2026-08-09T12:29:50.369057Z"}},"outputs":[],"execution_count":null},{"id":"6c46fcbd-b161-49e2-9c6d-070948080a39","cell_type":"code","source":"# Cell 11: Quick DICOM sanity check\n\nsample_uid = train_split.iloc[0][\"StudyInstanceUID\"]\n\npaths = get_study_slice_paths(\n    sample_uid,\n    split=\"train\",\n    total_slices=N_SLICES\n)\n\nprint(\"Sample UID      :\", sample_uid)\nprint(\"Selected slices:\", len(paths))\n\nif len(paths) == 0:\n    raise RuntimeError(\n        \"No DICOM slices found for sample study. \"\n        \"Check DATA_DIR and train_series structure.\"\n    )\n\nimg = read_dicom(\n    paths[len(paths) // 2],\n    IMAGE_SIZE\n)\n\nprint(\"Image shape    :\", img.shape)\nprint(\"Intensity min  :\", float(img.min()))\nprint(\"Intensity max  :\", float(img.max()))\nprint(\"Intensity mean :\", float(img.mean()))\nprint(\"✅ Cell 11 complete: DICOM read test passed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:29:52.826115Z","iopub.execute_input":"2026-08-09T12:29:52.8268Z","iopub.status.idle":"2026-08-09T12:29:53.849703Z","shell.execute_reply.started":"2026-08-09T12:29:52.826771Z","shell.execute_reply":"2026-08-09T12:29:53.848615Z"}},"outputs":[],"execution_count":null},{"id":"f2e43284-4612-4f1c-8b6a-eefde28cedb3","cell_type":"code","source":"# Cell 12: Dataloaders\n\ntrain_ds = KneeStudyDataset(\n    train_split,\n    split=\"train\",\n    targets=TARGETS,\n    n_slices=N_SLICES,\n    image_size=IMAGE_SIZE\n)\n\nval_ds = KneeStudyDataset(\n    val_split,\n    split=\"train\",\n    targets=TARGETS,\n    n_slices=N_SLICES,\n    image_size=IMAGE_SIZE\n)\n\ntrain_loader = DataLoader(\n    train_ds,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=NUM_WORKERS,\n    pin_memory=torch.cuda.is_available(),\n    persistent_workers=(NUM_WORKERS > 0)\n)\n\nval_loader = DataLoader(\n    val_ds,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=torch.cuda.is_available(),\n    persistent_workers=(NUM_WORKERS > 0)\n)\n\nxb, yb = next(iter(train_loader))\n\nprint(\"✅ Cell 12 complete\")\nprint(\"Train studies :\", len(train_ds))\nprint(\"Val studies   :\", len(val_ds))\nprint(\"Train batches :\", len(train_loader))\nprint(\"Val batches   :\", len(val_loader))\nprint(\"Batch X shape :\", tuple(xb.shape))\nprint(\"Batch Y shape :\", tuple(yb.shape))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:29:55.830182Z","iopub.execute_input":"2026-08-09T12:29:55.831561Z","iopub.status.idle":"2026-08-09T12:29:57.649906Z","shell.execute_reply.started":"2026-08-09T12:29:55.831529Z","shell.execute_reply":"2026-08-09T12:29:57.648653Z"}},"outputs":[],"execution_count":null},{"id":"a56ac63a-ad40-4ad6-a0d6-36bd21f70635","cell_type":"code","source":"# Cell 13: ConvNeXt + slice attention model\n\nclass KneeAttentionModel(nn.Module):\n\n    def __init__(\n        self,\n        model_name=\"convnext_tiny\",\n        num_classes=12,\n        pretrained=True\n    ):\n        super().__init__()\n\n        self.encoder = timm.create_model(\n            model_name,\n            pretrained=pretrained,\n            in_chans=1,\n            num_classes=0\n        )\n\n        feat_dim = self.encoder.num_features\n\n        self.attention = nn.Sequential(\n            nn.Linear(feat_dim, 256),\n            nn.Tanh(),\n            nn.Linear(256, 1)\n        )\n\n        self.head = nn.Sequential(\n            nn.LayerNorm(feat_dim),\n            nn.Dropout(0.30),\n            nn.Linear(feat_dim, num_classes)\n        )\n\n    def forward(self, x):\n        B, S, C, H, W = x.shape\n\n        x = x.reshape(B * S, C, H, W)\n        feat = self.encoder(x)\n        feat = feat.reshape(B, S, -1)\n\n        attn = torch.softmax(\n            self.attention(feat),\n            dim=1\n        )\n\n        study_feat = (feat * attn).sum(dim=1)\n        return self.head(study_feat)\n\n\nmodel = KneeAttentionModel(\n    num_classes=len(TARGETS)\n).to(device)\n\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(\n    p.numel() for p in model.parameters()\n    if p.requires_grad\n)\n\nwith torch.no_grad():\n    test_logits = model(\n        xb[:1].to(device)\n    )\n\nprint(\"✅ Cell 13 complete\")\nprint(\"Total params    :\", f\"{total_params:,}\")\nprint(\"Trainable params:\", f\"{trainable_params:,}\")\nprint(\"Test output     :\", tuple(test_logits.shape))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:30:00.12142Z","iopub.execute_input":"2026-08-09T12:30:00.122482Z","iopub.status.idle":"2026-08-09T12:30:07.811131Z","shell.execute_reply.started":"2026-08-09T12:30:00.122383Z","shell.execute_reply":"2026-08-09T12:30:07.810095Z"}},"outputs":[],"execution_count":null},{"id":"82c5155a-e038-4437-abc1-3c566698dbba","cell_type":"code","source":"# Cell 14: Stable class-balanced BCE loss\n\ny_train = train_split[TARGETS].values.astype(np.float32)\n\npos = y_train.sum(axis=0)\nneg = len(y_train) - pos\n\n# Raw inverse prevalence ratio\nraw_pos_weight = neg / np.maximum(pos, 1.0)\n\n# Stabilize rare classes:\n# sqrt compresses extreme ratios and clipping prevents a single class\n# from dominating gradients.\nstable_pos_weight = np.sqrt(raw_pos_weight)\nstable_pos_weight = np.clip(\n    stable_pos_weight,\n    1.0,\n    4.0\n)\n\npos_weight = torch.tensor(\n    stable_pos_weight,\n    dtype=torch.float32,\n    device=device\n)\n\nweight_table = pd.DataFrame({\n    \"Target\": TARGETS,\n    \"Positive\": pos.astype(int),\n    \"Negative\": neg.astype(int),\n    \"Raw_weight\": raw_pos_weight,\n    \"Stable_weight\": stable_pos_weight\n})\n\ncriterion = nn.BCEWithLogitsLoss(\n    pos_weight=pos_weight\n)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=LR,\n    weight_decay=WEIGHT_DECAY\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=EPOCHS,\n    eta_min=LR * 0.05\n)\n\nprint(\"✅ Cell 14 complete: stable weighted BCE configured\")\ndisplay(weight_table.round(4))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:30:08.123596Z","iopub.execute_input":"2026-08-09T12:30:08.124173Z","iopub.status.idle":"2026-08-09T12:30:08.140388Z","shell.execute_reply.started":"2026-08-09T12:30:08.124144Z","shell.execute_reply":"2026-08-09T12:30:08.139673Z"}},"outputs":[],"execution_count":null},{"id":"60c869a7-d878-47dd-9f85-2a1ff48b58ca","cell_type":"code","source":"# Cell 15: ROC-AUC helper\n\ndef calculate_auc(y_true, y_pred):\n    scores = {}\n    valid_scores = []\n\n    for i, target in enumerate(TARGETS):\n        yt = y_true[:, i]\n        yp = y_pred[:, i]\n\n        if len(np.unique(yt)) < 2:\n            score = np.nan\n        else:\n            score = roc_auc_score(yt, yp)\n            valid_scores.append(score)\n\n        scores[target] = score\n\n    macro_auc = (\n        float(np.mean(valid_scores))\n        if valid_scores\n        else np.nan\n    )\n\n    return macro_auc, scores, len(valid_scores)\n\nprint(\"✅ Cell 15 complete: ROC-AUC helper ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:30:11.413888Z","iopub.execute_input":"2026-08-09T12:30:11.414325Z","iopub.status.idle":"2026-08-09T12:30:11.421247Z","shell.execute_reply.started":"2026-08-09T12:30:11.414297Z","shell.execute_reply":"2026-08-09T12:30:11.42046Z"}},"outputs":[],"execution_count":null},{"id":"6dd40d24-88c9-4d2e-ae2a-b98dd55e579e","cell_type":"code","source":"# Cell 16: Train one epoch\n\ndef train_one_epoch(model, loader):\n    model.train()\n\n    running_loss = 0.0\n\n    pbar = tqdm(\n        loader,\n        desc=\"Training\",\n        leave=False\n    )\n\n    for x, y in pbar:\n        x = x.to(device, non_blocking=True)\n        y = y.to(device, non_blocking=True)\n\n        optimizer.zero_grad(set_to_none=True)\n\n        logits = model(x)\n\n        loss = criterion(logits, y)\n\n        loss.backward()\n\n        torch.nn.utils.clip_grad_norm_(\n            model.parameters(),\n            max_norm=1.0\n        )\n\n        optimizer.step()\n\n        running_loss += loss.item()\n\n        pbar.set_postfix(\n            loss=f\"{loss.item():.4f}\"\n        )\n\n    return running_loss / max(len(loader), 1)\n\nprint(\"✅ Cell 16 complete: training function ready\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:30:15.372702Z","iopub.execute_input":"2026-08-09T12:30:15.37313Z","iopub.status.idle":"2026-08-09T12:30:15.380098Z","shell.execute_reply.started":"2026-08-09T12:30:15.373103Z","shell.execute_reply":"2026-08-09T12:30:15.37922Z"}},"outputs":[],"execution_count":null},{"id":"81f67b69-106c-43a0-bdbd-368f310d4e1e","cell_type":"code","source":"# Cell 17: Validation\n\n@torch.no_grad()\ndef validate(model, loader):\n    model.eval()\n\n    all_targets = []\n    all_probs = []\n\n    pbar = tqdm(\n        loader,\n        desc=\"Validation\",\n        leave=False\n    )\n\n    for x, y in pbar:\n        x = x.to(device, non_blocking=True)\n\n        logits = model(x)\n        probs = torch.sigmoid(logits)\n\n        all_targets.append(y.numpy())\n        all_probs.append(probs.cpu().numpy())\n\n    y_true = np.concatenate(all_targets, axis=0)\n    y_pred = np.concatenate(all_probs, axis=0)\n\n    macro_auc, scores, n_valid = calculate_auc(\n        y_true,\n        y_pred\n    )\n\n    return macro_auc, scores, n_valid, y_true, y_pred\n\nprint(\"✅ Cell 17 complete: validation function ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:30:20.099921Z","iopub.execute_input":"2026-08-09T12:30:20.100291Z","iopub.status.idle":"2026-08-09T12:30:20.10713Z","shell.execute_reply.started":"2026-08-09T12:30:20.100264Z","shell.execute_reply":"2026-08-09T12:30:20.106152Z"}},"outputs":[],"execution_count":null},{"id":"a4413492-4bd5-4f91-aa4b-988076c1f548","cell_type":"code","source":"# Cell 18: Training loop\n\nWORK_DIR = Path(\"/kaggle/working\") if Path(\"/kaggle/working\").exists() else Path(\"./working\")\nWORK_DIR.mkdir(parents=True, exist_ok=True)\n\nBEST_MODEL_PATH = WORK_DIR / \"best_rsna_knee_model.pt\"\nHISTORY_PATH = WORK_DIR / \"training_history.csv\"\n\nbest_auc = -np.inf\nhistory = []\n\nprint(\"Starting training...\")\nprint(\"=\" * 70)\n\nfor epoch in range(1, EPOCHS + 1):\n\n    train_loss = train_one_epoch(\n        model,\n        train_loader\n    )\n\n    macro_auc, scores, n_valid, y_true, y_pred = validate(\n        model,\n        val_loader\n    )\n\n    current_lr = optimizer.param_groups[0][\"lr\"]\n    scheduler.step()\n\n    history.append({\n        \"epoch\": epoch,\n        \"train_loss\": train_loss,\n        \"macro_auc\": macro_auc,\n        \"valid_auc_labels\": n_valid,\n        \"lr\": current_lr\n    })\n\n    print(\n        f\"Epoch {epoch:02d}/{EPOCHS} | \"\n        f\"loss={train_loss:.4f} | \"\n        f\"macro AUC={macro_auc:.4f} | \"\n        f\"valid={n_valid}/12 | \"\n        f\"lr={current_lr:.2e}\"\n    )\n\n    for target in TARGETS:\n        score = scores[target]\n\n        if np.isnan(score):\n            print(f\"  {target:<20} nan\")\n        else:\n            print(f\"  {target:<20} {score:.4f}\")\n\n    if not np.isnan(macro_auc) and macro_auc > best_auc:\n        best_auc = macro_auc\n\n        torch.save(\n            model.state_dict(),\n            BEST_MODEL_PATH\n        )\n\n        print(\"  ✅ Saved new best model\")\n\n    print(\"-\" * 70)\n\n    pd.DataFrame(history).to_csv(\n        HISTORY_PATH,\n        index=False\n    )\n\n    gc.collect()\n\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\nprint(\"✅ Cell 18 complete: training finished\")\nprint(\"Best macro AUC:\", best_auc)\nprint(\"Best model    :\", BEST_MODEL_PATH)\nprint(\"History       :\", HISTORY_PATH)\n\ndisplay(pd.DataFrame(history))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:30:22.808552Z","iopub.execute_input":"2026-08-09T12:30:22.809531Z","iopub.status.idle":"2026-08-09T12:34:29.925189Z","shell.execute_reply.started":"2026-08-09T12:30:22.809501Z","shell.execute_reply":"2026-08-09T12:34:29.924312Z"}},"outputs":[],"execution_count":null},{"id":"af1cd063-a489-4e2d-a2d7-d46331800d59","cell_type":"code","source":"# Cell 19: Load best model\n\nif not BEST_MODEL_PATH.exists():\n    raise FileNotFoundError(\n        f\"Best model not found: {BEST_MODEL_PATH}\"\n    )\n\nmodel.load_state_dict(\n    torch.load(\n        BEST_MODEL_PATH,\n        map_location=device\n    )\n)\n\nmodel.eval()\n\nfinal_auc, final_scores, final_n_valid, _, _ = validate(\n    model,\n    val_loader\n)\n\nprint(\"✅ Cell 19 complete: best model loaded\")\nprint(f\"Best validation macro AUC: {final_auc:.4f}\")\nprint(f\"Valid AUC labels          : {final_n_valid}/12\")\n\nfor t in TARGETS:\n    s = final_scores[t]\n    print(\n        f\"{t:<20} \"\n        + (\"nan\" if np.isnan(s) else f\"{s:.4f}\")\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:36:01.058071Z","iopub.execute_input":"2026-08-09T12:36:01.058379Z","iopub.status.idle":"2026-08-09T12:36:05.159344Z","shell.execute_reply.started":"2026-08-09T12:36:01.058347Z","shell.execute_reply":"2026-08-09T12:36:05.158754Z"}},"outputs":[],"execution_count":null},{"id":"7448a02e-cf17-4838-87b9-4513a73e8b28","cell_type":"code","source":"# Cell 20: Test dataset + dataloader\n\ntest_ds = KneeStudyDataset(\n    test_df,\n    split=\"test\",\n    targets=None,\n    n_slices=N_SLICES,\n    image_size=IMAGE_SIZE\n)\n\ntest_loader = DataLoader(\n    test_ds,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=torch.cuda.is_available(),\n    persistent_workers=(NUM_WORKERS > 0)\n)\n\nprint(\"✅ Cell 20 complete\")\nprint(\"Test studies:\", len(test_ds))\nprint(\"Test batches:\", len(test_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:36:12.837181Z","iopub.execute_input":"2026-08-09T12:36:12.837728Z","iopub.status.idle":"2026-08-09T12:36:12.843573Z","shell.execute_reply.started":"2026-08-09T12:36:12.837699Z","shell.execute_reply":"2026-08-09T12:36:12.842749Z"}},"outputs":[],"execution_count":null},{"id":"b38a1caa-4e82-4ca3-b80e-7385ad2edf96","cell_type":"code","source":"# Cell 21: Test inference\n\nall_uids = []\nall_preds = []\n\nmodel.eval()\n\nwith torch.no_grad():\n    for x, uids in tqdm(\n        test_loader,\n        desc=\"Test inference\"\n    ):\n        x = x.to(\n            device,\n            non_blocking=True\n        )\n\n        logits = model(x)\n        probs = torch.sigmoid(logits)\n\n        all_preds.append(\n            probs.cpu().numpy()\n        )\n\n        all_uids.extend(\n            list(uids)\n        )\n\nif len(all_preds) == 0:\n    raise RuntimeError(\"No test predictions were generated\")\n\npreds = np.concatenate(\n    all_preds,\n    axis=0\n)\n\nprint(\"✅ Cell 21 complete\")\nprint(\"Prediction shape:\", preds.shape)\nprint(\"Prediction min  :\", float(preds.min()))\nprint(\"Prediction max  :\", float(preds.max()))\nprint(\"Prediction mean :\", float(preds.mean()))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:36:15.54575Z","iopub.execute_input":"2026-08-09T12:36:15.546242Z","iopub.status.idle":"2026-08-09T12:36:20.526523Z","shell.execute_reply.started":"2026-08-09T12:36:15.546214Z","shell.execute_reply":"2026-08-09T12:36:20.525111Z"}},"outputs":[],"execution_count":null},{"id":"5e1d047c-2df9-4ec6-8423-7661046a369c","cell_type":"code","source":"# Cell 22: Build submission.csv\n\nsubmission = pd.DataFrame({\n    \"StudyInstanceUID\": all_uids\n})\n\nfor i, target in enumerate(TARGETS):\n    submission[target] = preds[:, i]\n\nsubmission = submission[\n    sample_submission.columns\n]\n\nOUTPUT_PATH = WORK_DIR / \"submission.csv\"\n\nsubmission.to_csv(\n    OUTPUT_PATH,\n    index=False\n)\n\nif submission.isna().any().any():\n    raise ValueError(\"submission.csv contains NaN values\")\n\nif len(submission) != len(test_df):\n    raise ValueError(\n        f\"Submission row count {len(submission)} \"\n        f\"does not match test row count {len(test_df)}\"\n    )\n\nprint(\"✅ Cell 22 complete: submission ready\")\nprint(\"Saved to:\", OUTPUT_PATH)\nprint(\"Shape   :\", submission.shape)\nprint(\"NaNs    :\", int(submission.isna().sum().sum()))\ndisplay(submission.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:36:25.502495Z","iopub.execute_input":"2026-08-09T12:36:25.502858Z","iopub.status.idle":"2026-08-09T12:36:25.529053Z","shell.execute_reply.started":"2026-08-09T12:36:25.502824Z","shell.execute_reply":"2026-08-09T12:36:25.528404Z"}},"outputs":[],"execution_count":null},{"id":"f8ab89e7-8736-4310-99a3-a436e1d607b7","cell_type":"markdown","source":"## Optional next stage: self-supervised / semi-supervised training\n\nThe current notebook is the strongest stable baseline for this dataset structure.\n\nA next-stage competition improvement can use:\n1. All labeled + unlabeled MRI for self-supervised encoder pretraining.\n2. Reports to derive weak/pseudo labels for unlabeled studies.\n3. 3–5 fold multilabel stratified cross-validation.\n4. Ensemble predictions across folds.\n5. Separate encoders per anatomical plane.","metadata":{}},{"id":"dec8bad3-655c-470a-9e84-75febe201ef3","cell_type":"code","source":"# Cell 23: Final notebook diagnostics\n\nprint(\"=\" * 70)\nprint(\"FINAL RUN SUMMARY\")\nprint(\"=\" * 70)\n\nprint(\"Device               :\", device)\nprint(\"Train studies        :\", len(train_split))\nprint(\"Validation studies   :\", len(val_split))\nprint(\"Test studies         :\", len(test_df))\nprint(\"Slices per study     :\", N_SLICES)\nprint(\"Best macro AUC       :\", best_auc)\nprint(\"Best model path      :\", BEST_MODEL_PATH)\nprint(\"Submission path      :\", OUTPUT_PATH)\n\nprint(\"\\n✅ Cell 23 complete\")\nprint(\"✅ Notebook pipeline finished successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-09T12:36:40.066748Z","iopub.execute_input":"2026-08-09T12:36:40.06704Z","iopub.status.idle":"2026-08-09T12:36:40.074271Z","shell.execute_reply.started":"2026-08-09T12:36:40.067017Z","shell.execute_reply":"2026-08-09T12:36:40.073295Z"}},"outputs":[],"execution_count":null}]}