{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session\n\n# Use the kagglehub client library to attach Kaggle resources like competitions, datasets, and models to your session\n# Learn more about kagglehub: https://github.com/Kaggle/kagglehub/blob/main/README.md\n\nimport kagglehub\n# kagglehub.dataset_download('<owner>/<dataset-slug>')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-09-14T08:46:47.831983Z","iopub.execute_input":"2026-09-14T08:46:47.832325Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# RSNA KNEE ABNORMALITY DETECTION\n# FINAL SWIN TRANSFORMER PIPELINE\n#\n# Architecture:\n#\n# DICOM MRI\n#    |\n#    +--> multiple MRI series\n#    |\n#    +--> multiple anatomical slice locations\n#    |\n#    +--> 2.5D 3-channel images\n#    |\n#    +--> Swin Transformer\n#    |\n#    +--> view attention\n#    |\n#    +--> 12-label classifier\n#\n# Training:\n#   - 5-fold study-level CV\n#   - gold labels only for validation\n#   - optional report-derived pseudo labels\n#   - mixed precision\n#   - gradient accumulation\n#   - cosine LR\n#\n# Output:\n#   /kaggle/working/submission.csv\n#\n# ============================================================\n\n\n# ============================================================\n# 1. IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport glob\nimport math\nimport random\nimport warnings\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\n\nwarnings.filterwarnings(\"ignore\")\n\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import roc_auc_score\n\nimport pydicom\n\nfrom PIL import Image\n\nimport torchvision\nfrom torchvision.models import (\n    swin_t,\n    Swin_T_Weights\n)\n\n\n# ============================================================\n# 2. CONFIGURATION\n# ============================================================\n\nSEED = 2026\n\nIMG_SIZE = 224\n\n# Number of study-level views.\n# Each view is a 3-slice 2.5D image.\nNUM_VIEWS = 12\n\n# Number of series used per study.\nMAX_SERIES = 3\n\nBATCH_SIZE = 2\n\n# Gradient accumulation lets us simulate larger batches.\nACCUM_STEPS = 4\n\nEPOCHS = 12\n\nLR = 1.5e-4\n\nWEIGHT_DECAY = 1e-4\n\nNUM_WORKERS = 2\n\nN_FOLDS = 5\n\nUSE_AMP = True\n\n# Use report-derived pseudo labels?\n#\n# Default False because this keeps the baseline fully\n# ground-truth supervised.\n#\n# You can enable it after validating the report parser.\nUSE_PSEUDO_LABELS = False\n\nPSEUDO_WEIGHT = 0.20\n\nDEVICE = torch.device(\n    \"cuda\" if torch.cuda.is_available() else \"cpu\"\n)\n\nprint(\"PyTorch:\", torch.__version__)\nprint(\"Torchvision:\", torchvision.__version__)\nprint(\"Device:\", DEVICE)\n\n\n# ============================================================\n# 3. REPRODUCIBILITY\n# ============================================================\n\ndef seed_everything(seed):\n\n    random.seed(seed)\n    np.random.seed(seed)\n\n    torch.manual_seed(seed)\n\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n\n    torch.backends.cudnn.benchmark = True\n    torch.backends.cudnn.deterministic = False\n\n\nseed_everything(SEED)\n\n\n# ============================================================\n# 4. FIND COMPETITION DATA\n# ============================================================\n\ndef find_competition_root():\n\n    direct_candidates = [\n        \"/kaggle/input/rsna-knee-abnormality-detection\",\n        \"/kaggle/input/rsna-knee-abnormality-detection-2026\",\n    ]\n\n    for root in direct_candidates:\n\n        if os.path.exists(\n            os.path.join(root, \"train.csv\")\n        ):\n\n            return root\n\n    # Search all mounted Kaggle inputs.\n    for root, dirs, files in os.walk(\n        \"/kaggle/input\"\n    ):\n\n        if (\n            \"train.csv\" in files\n            and \"train_series.csv\" in files\n        ):\n\n            return root\n\n    raise FileNotFoundError(\n        \"RSNA competition dataset was not found \"\n        \"under /kaggle/input.\"\n    )\n\n\nDATA_ROOT = find_competition_root()\n\nprint(\"\\nDATA ROOT:\")\nprint(DATA_ROOT)\n\n\nTRAIN_CSV = os.path.join(\n    DATA_ROOT,\n    \"train.csv\"\n)\n\nTRAIN_SERIES_CSV = os.path.join(\n    DATA_ROOT,\n    \"train_series.csv\"\n)\n\nTEST_CSV = os.path.join(\n    DATA_ROOT,\n    \"test.csv\"\n)\n\nTEST_SERIES_CSV = os.path.join(\n    DATA_ROOT,\n    \"test_series.csv\"\n)\n\nTRAIN_SERIES_DIR = os.path.join(\n    DATA_ROOT,\n    \"train_series\"\n)\n\nTEST_SERIES_DIR = os.path.join(\n    DATA_ROOT,\n    \"test_series\"\n)\n\n\n# ============================================================\n# 5. LOAD CSV FILES\n# ============================================================\n\ntrain = pd.read_csv(TRAIN_CSV)\n\ntrain_series = pd.read_csv(\n    TRAIN_SERIES_CSV\n)\n\ntest = pd.read_csv(TEST_CSV)\n\ntest_series = pd.read_csv(\n    TEST_SERIES_CSV\n)\n\n\nprint(\"\\nShapes\")\nprint(\"train:\", train.shape)\nprint(\"train_series:\", train_series.shape)\nprint(\"test:\", test.shape)\nprint(\"test_series:\", test_series.shape)\n\n\n# ============================================================\n# 6. TARGETS\n# ============================================================\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\nNUM_CLASSES = len(TARGETS)\n\n\n# ============================================================\n# 7. GOLD LABEL DATA\n# ============================================================\n\ngold_mask = (\n    train[TARGETS]\n    .notna()\n    .all(axis=1)\n)\n\ngold = train.loc[\n    gold_mask\n].copy().reset_index(drop=True)\n\nprint(\n    \"\\nGold labelled studies:\",\n    len(gold)\n)\n\nif len(gold) != 58:\n\n    print(\n        \"WARNING: expected approximately 58 \"\n        \"gold-labelled studies.\"\n    )\n\n\n# ============================================================\n# 8. LABEL DISTRIBUTION\n# ============================================================\n\nprint(\"\\nLabel prevalence\")\n\nfor target in TARGETS:\n\n    n = int(\n        gold[target].sum()\n    )\n\n    p = float(\n        gold[target].mean()\n    )\n\n    print(\n        f\"{target:20s} \"\n        f\"{n:3d}/{len(gold):3d} \"\n        f\"({p:.3f})\"\n    )\n\n\n# ============================================================\n# 9. SERIES LOOKUP\n# ============================================================\n\ntrain_series_groups = {\n    str(k): v.copy()\n    for k, v in train_series.groupby(\n        \"StudyInstanceUID\"\n    )\n}\n\ntest_series_groups = {\n    str(k): v.copy()\n    for k, v in test_series.groupby(\n        \"StudyInstanceUID\"\n    )\n}\n\n\n# ============================================================\n# 10. DICOM SORTING\n# ============================================================\n\ndef dicom_sort_key(ds):\n\n    # Prefer ImagePositionPatient.\n    if hasattr(\n        ds,\n        \"ImagePositionPatient\"\n    ):\n\n        try:\n            return float(\n                ds.ImagePositionPatient[-1]\n            )\n\n        except Exception:\n            pass\n\n    # Fall back to InstanceNumber.\n    if hasattr(\n        ds,\n        \"InstanceNumber\"\n    ):\n\n        try:\n            return float(\n                ds.InstanceNumber\n            )\n\n        except Exception:\n            pass\n\n    return 0.0\n\n\n# ============================================================\n# 11. LOAD DICOM SERIES\n# ============================================================\n\ndef load_series(\n    series_path\n):\n\n    files = glob.glob(\n        os.path.join(\n            series_path,\n            \"*.dcm\"\n        )\n    )\n\n    if not files:\n        return None\n\n    slices = []\n\n    for path in files:\n\n        try:\n\n            ds = pydicom.dcmread(\n                path,\n                force=True\n            )\n\n            if not hasattr(\n                ds,\n                \"PixelData\"\n            ):\n                continue\n\n            slices.append(ds)\n\n        except Exception:\n            continue\n\n    if not slices:\n        return None\n\n    slices.sort(\n        key=dicom_sort_key\n    )\n\n    volume = []\n\n    for ds in slices:\n\n        try:\n\n            arr = ds.pixel_array.astype(\n                np.float32\n            )\n\n        except Exception:\n            continue\n\n        # DICOM rescaling.\n        slope = float(\n            getattr(\n                ds,\n                \"RescaleSlope\",\n                1.0\n            )\n        )\n\n        intercept = float(\n            getattr(\n                ds,\n                \"RescaleIntercept\",\n                0.0\n            )\n        )\n\n        arr = (\n            arr * slope\n            + intercept\n        )\n\n        # Robust MRI normalization.\n        p01 = np.percentile(\n            arr,\n            1\n        )\n\n        p99 = np.percentile(\n            arr,\n            99\n        )\n\n        if p99 > p01:\n\n            arr = np.clip(\n                arr,\n                p01,\n                p99\n            )\n\n            arr = (\n                arr - p01\n            ) / (\n                p99 - p01\n            )\n\n        else:\n\n            arr = np.zeros_like(\n                arr\n            )\n\n        volume.append(\n            arr\n        )\n\n    if not volume:\n        return None\n\n    # Make dimensions consistent.\n    h = min(\n        x.shape[0]\n        for x in volume\n    )\n\n    w = min(\n        x.shape[1]\n        for x in volume\n    )\n\n    volume = [\n        x[:h, :w]\n        for x in volume\n    ]\n\n    return np.stack(\n        volume\n    ).astype(\n        np.float32\n    )\n\n\n# ============================================================\n# 12. SERIES RANKING\n# ============================================================\n\ndef rank_series(\n    study_id,\n    series_df\n):\n\n    rows = series_df.get(\n        str(study_id),\n        None\n    )\n\n    if rows is None:\n        return []\n\n    rows = rows.copy()\n\n    rows[\"score\"] = 0.0\n\n    # Fluid sensitive is extremely useful for many\n    # ligament, meniscus, effusion and bone findings.\n    if \"Fluid_Sensitive\" in rows.columns:\n\n        rows[\"score\"] += (\n            rows[\n                \"Fluid_Sensitive\"\n            ]\n            .fillna(0)\n            .astype(float)\n            * 4.0\n        )\n\n    # Fat suppression.\n    if \"Fat_Suppression\" in rows.columns:\n\n        rows[\"score\"] += (\n            rows[\n                \"Fat_Suppression\"\n            ]\n            .fillna(0)\n            .astype(float)\n            * 2.0\n        )\n\n    # Anatomical plane.\n    if \"Anatomical_Plane\" in rows.columns:\n\n        plane = (\n            rows[\n                \"Anatomical_Plane\"\n            ]\n            .fillna(\"\")\n            .astype(str)\n            .str.lower()\n        )\n\n        rows.loc[\n            plane.eq(\"sagittal\"),\n            \"score\"\n        ] += 2.0\n\n        rows.loc[\n            plane.eq(\"coronal\"),\n            \"score\"\n        ] += 1.5\n\n        rows.loc[\n            plane.eq(\"axial\"),\n            \"score\"\n        ] += 1.0\n\n    rows = rows.sort_values(\n        \"score\",\n        ascending=False\n    )\n\n    return rows[\n        \"SeriesInstanceUID\"\n    ].astype(str).tolist()[\n        :MAX_SERIES\n    ]\n\n\n# ============================================================\n# 13. RESIZE\n# ============================================================\n\ndef resize_slice(\n    arr,\n    size=IMG_SIZE\n):\n\n    arr = np.clip(\n        arr,\n        0,\n        1\n    )\n\n    image = Image.fromarray(\n        (\n            arr * 255\n        ).astype(\n            np.uint8\n        )\n    )\n\n    image = image.resize(\n        (size, size),\n        Image.Resampling.BILINEAR\n    )\n\n    return (\n        np.asarray(\n            image\n        ).astype(\n            np.float32\n        )\n        / 255.0\n    )\n\n\n# ============================================================\n# 14. 2.5D VIEW GENERATION\n# ============================================================\n\ndef make_views(\n    volume,\n    num_views=NUM_VIEWS\n):\n\n    if volume is None:\n\n        return np.zeros(\n            (\n                num_views,\n                3,\n                IMG_SIZE,\n                IMG_SIZE\n            ),\n            dtype=np.float32\n        )\n\n    depth = volume.shape[0]\n\n    if depth < 3:\n\n        base = resize_slice(\n            volume[0]\n        )\n\n        return np.repeat(\n            base[None, None, :, :],\n            num_views * 3,\n            axis=0\n        ).reshape(\n            num_views,\n            3,\n            IMG_SIZE,\n            IMG_SIZE\n        )\n\n    # Don't use the extreme slices.\n    lo = max(\n        1,\n        int(depth * 0.10)\n    )\n\n    hi = min(\n        depth - 2,\n        int(depth * 0.90)\n    )\n\n    centers = np.linspace(\n        lo,\n        hi,\n        num_views\n    ).astype(int)\n\n    views = []\n\n    for center in centers:\n\n        prev_idx = max(\n            0,\n            center - 1\n        )\n\n        next_idx = min(\n            depth - 1,\n            center + 1\n        )\n\n        # Three neighbouring slices.\n        c0 = resize_slice(\n            volume[prev_idx]\n        )\n\n        c1 = resize_slice(\n            volume[center]\n        )\n\n        c2 = resize_slice(\n            volume[next_idx]\n        )\n\n        view = np.stack(\n            [\n                c0,\n                c1,\n                c2\n            ],\n            axis=0\n        )\n\n        views.append(\n            view\n        )\n\n    return np.stack(\n        views\n    ).astype(\n        np.float32\n    )\n\n\n# ============================================================\n# 15. CACHE\n# ============================================================\n\nCACHE_ROOT = (\n    \"/kaggle/working/\"\n    \"rsna_swin_cache\"\n)\n\nos.makedirs(\n    CACHE_ROOT,\n    exist_ok=True\n)\n\n\ndef study_cache_path(\n    study_id\n):\n\n    return os.path.join(\n        CACHE_ROOT,\n        f\"{study_id}.npy\"\n    )\n\n\n# ============================================================\n# 16. LOAD STUDY VIEWS\n# ============================================================\n\ndef get_study_views(\n    study_id,\n    series_lookup,\n    series_root\n):\n\n    cache_path = study_cache_path(\n        study_id\n    )\n\n    if os.path.exists(\n        cache_path\n    ):\n\n        try:\n\n            return np.load(\n                cache_path\n            )\n\n        except Exception:\n            pass\n\n    series_ids = rank_series(\n        study_id,\n        series_lookup\n    )\n\n    all_views = []\n\n    for series_id in series_ids:\n\n        path = os.path.join(\n            series_root,\n            str(study_id),\n            str(series_id)\n        )\n\n        if not os.path.isdir(\n            path\n        ):\n            continue\n\n        volume = load_series(\n            path\n        )\n\n        if volume is None:\n            continue\n\n        views = make_views(\n            volume\n        )\n\n        all_views.append(\n            views\n        )\n\n    if not all_views:\n\n        output = np.zeros(\n            (\n                MAX_SERIES,\n                NUM_VIEWS,\n                3,\n                IMG_SIZE,\n                IMG_SIZE\n            ),\n            dtype=np.float32\n        )\n\n    else:\n\n        # Pad series dimension.\n        output = np.zeros(\n            (\n                MAX_SERIES,\n                NUM_VIEWS,\n                3,\n                IMG_SIZE,\n                IMG_SIZE\n            ),\n            dtype=np.float32\n        )\n\n        for i, views in enumerate(\n            all_views[\n                :MAX_SERIES\n            ]\n        ):\n\n            output[i] = views\n\n    np.save(\n        cache_path,\n        output.astype(\n            np.float32\n        )\n    )\n\n    return output\n\n\n# ============================================================\n# 17. TEST ONE STUDY\n# ============================================================\n\nexample_id = str(\n    gold.iloc[0][\n        \"StudyInstanceUID\"\n    ]\n)\n\nprint(\n    \"\\nTesting study:\",\n    example_id\n)\n\nexample_views = get_study_views(\n    example_id,\n    train_series_groups,\n    TRAIN_SERIES_DIR\n)\n\nprint(\n    \"Study tensor:\",\n    example_views.shape\n)\n\n# Expected:\n# [MAX_SERIES, NUM_VIEWS, 3, 224, 224]\n\n\n# ============================================================\n# 18. DATASET\n# ============================================================\n\nclass KneeSwinDataset(\n    Dataset\n):\n\n    def __init__(\n        self,\n        dataframe,\n        series_lookup,\n        series_root,\n        training=True\n    ):\n\n        self.df = (\n            dataframe\n            .reset_index(drop=True)\n        )\n\n        self.series_lookup = (\n            series_lookup\n        )\n\n        self.series_root = (\n            series_root\n        )\n\n        self.training = training\n\n    def __len__(self):\n\n        return len(\n            self.df\n        )\n\n    def __getitem__(\n        self,\n        idx\n    ):\n\n        row = self.df.iloc[idx]\n\n        study_id = str(\n            row[\n                \"StudyInstanceUID\"\n            ]\n        )\n\n        x = get_study_views(\n            study_id,\n            self.series_lookup,\n            self.series_root\n        )\n\n        # Flatten series + views.\n        #\n        # [series, views, C, H, W]\n        #\n        # -> [series*views, C, H, W]\n\n        x = x.reshape(\n            MAX_SERIES * NUM_VIEWS,\n            3,\n            IMG_SIZE,\n            IMG_SIZE\n        )\n\n        # Randomly choose views during training.\n        if self.training:\n\n            valid = (\n                np.where(\n                    x.reshape(\n                        x.shape[0],\n                        -1\n                    ).sum(axis=1) > 0\n                )[0]\n            )\n\n            if len(valid) > NUM_VIEWS:\n\n                selected = np.random.choice(\n                    valid,\n                    size=NUM_VIEWS,\n                    replace=False\n                )\n\n                x = x[\n                    selected\n                ]\n\n            else:\n\n                selected = np.arange(\n                    min(\n                        NUM_VIEWS,\n                        len(x)\n                    )\n                )\n\n                x = x[\n                    selected\n                ]\n\n        else:\n\n            # Deterministic inference.\n            valid = (\n                np.where(\n                    x.reshape(\n                        x.shape[0],\n                        -1\n                    ).sum(axis=1) > 0\n                )[0]\n            )\n\n            if len(valid) >= NUM_VIEWS:\n\n                selected = np.linspace(\n                    0,\n                    len(valid) - 1,\n                    NUM_VIEWS\n                ).astype(int)\n\n                x = x[\n                    valid[selected]\n                ]\n\n        # If fewer views than required,\n        # pad by repeating.\n        if x.shape[0] < NUM_VIEWS:\n\n            repeat_count = (\n                NUM_VIEWS\n                - x.shape[0]\n            )\n\n            if x.shape[0] > 0:\n\n                extra = x[\n                    :1\n                ].repeat(\n                    repeat_count,\n                    axis=0\n                )\n\n                x = np.concatenate(\n                    [x, extra],\n                    axis=0\n                )\n\n            else:\n\n                x = np.zeros(\n                    (\n                        NUM_VIEWS,\n                        3,\n                        IMG_SIZE,\n                        IMG_SIZE\n                    ),\n                    dtype=np.float32\n                )\n\n        image = torch.from_numpy(\n            x.astype(\n                np.float32\n            )\n        )\n\n        if self.training:\n\n            y = torch.tensor(\n                row[\n                    TARGETS\n                ].values.astype(\n                    np.float32\n                )\n            )\n\n            return image, y\n\n        return (\n            image,\n            study_id\n        )\n\n\n# ============================================================\n# 19. SWIN MODEL\n# ============================================================\n\nclass SwinKneeModel(\n    nn.Module\n):\n\n    def __init__(\n        self,\n        num_classes=12,\n        pretrained=True\n    ):\n\n        super().__init__()\n\n        # ----------------------------------------------------\n        # Load ImageNet-pretrained Swin-Tiny.\n        # ----------------------------------------------------\n\n        weights = None\n\n        if pretrained:\n\n            try:\n\n                weights = (\n                    Swin_T_Weights\n                    .IMAGENET1K_V1\n                )\n\n                print(\n                    \"Using ImageNet \"\n                    \"pretrained Swin-T.\"\n                )\n\n            except Exception:\n\n                weights = None\n\n        self.backbone = swin_t(\n            weights=weights\n        )\n\n        feature_dim = (\n            self.backbone.head.in_features\n        )\n\n        # Remove original ImageNet head.\n        self.backbone.head = (\n            nn.Identity()\n        )\n\n        # ----------------------------------------------------\n        # View attention.\n        # ----------------------------------------------------\n\n        self.attention = nn.Sequential(\n\n            nn.Linear(\n                feature_dim,\n                256\n            ),\n\n            nn.LayerNorm(\n                256\n            ),\n\n            nn.GELU(),\n\n            nn.Linear(\n                256,\n                1\n            )\n        )\n\n        # ----------------------------------------------------\n        # Classification head.\n        # ----------------------------------------------------\n\n        self.dropout = nn.Dropout(\n            0.25\n        )\n\n        self.classifier = nn.Sequential(\n\n            nn.LayerNorm(\n                feature_dim\n            ),\n\n            nn.Dropout(\n                0.25\n            ),\n\n            nn.Linear(\n                feature_dim,\n                512\n            ),\n\n            nn.GELU(),\n\n            nn.Dropout(\n                0.20\n            ),\n\n            nn.Linear(\n                512,\n                num_classes\n            )\n        )\n\n    def forward(\n        self,\n        x\n    ):\n\n        # x:\n        # [B, V, 3, H, W]\n\n        B, V, C, H, W = x.shape\n\n        x = x.reshape(\n            B * V,\n            C,\n            H,\n            W\n        )\n\n        # Swin features.\n        features = (\n            self.backbone(\n                x\n            )\n        )\n\n        # [B*V, D]\n        features = features.reshape(\n            B,\n            V,\n            -1\n        )\n\n        # Attention score per view.\n        scores = self.attention(\n            features\n        )\n\n        # [B,V,1]\n        weights = torch.softmax(\n            scores,\n            dim=1\n        )\n\n        # Study representation.\n        pooled = (\n            features\n            * weights\n        ).sum(\n            dim=1\n        )\n\n        logits = self.classifier(\n            pooled\n        )\n\n        return logits\n\n\n# ============================================================\n# 20. MODEL CHECK\n# ============================================================\n\nmodel = SwinKneeModel(\n    num_classes=NUM_CLASSES,\n    pretrained=True\n)\n\nmodel = model.to(\n    DEVICE\n)\n\ndummy = torch.randn(\n    1,\n    NUM_VIEWS,\n    3,\n    IMG_SIZE,\n    IMG_SIZE,\n    device=DEVICE\n)\n\nwith torch.no_grad():\n\n    dummy_output = model(\n        dummy\n    )\n\nprint(\n    \"\\nModel output:\",\n    dummy_output.shape\n)\n\ndel model\ndel dummy\ndel dummy_output\n\ngc.collect()\n\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n\n\n# ============================================================\n# 21. METRIC\n# ============================================================\n\ndef calculate_auc(\n    y_true,\n    y_pred\n):\n\n    scores = []\n\n    for i, target in enumerate(\n        TARGETS\n    ):\n\n        try:\n\n            score = roc_auc_score(\n                y_true[:, i],\n                y_pred[:, i]\n            )\n\n        except ValueError:\n\n            score = np.nan\n\n        scores.append(\n            score\n        )\n\n    return (\n        float(\n            np.nanmean(\n                scores\n            )\n        ),\n        scores\n    )\n\n\n# ============================================================\n# 22. TRAIN ONE FOLD\n# ============================================================\n\ndef train_fold(\n    train_df,\n    valid_df,\n    fold\n):\n\n    print(\n        \"\\n\"\n        + \"=\" * 80\n    )\n\n    print(\n        f\"FOLD {fold}\"\n    )\n\n    print(\n        \"=\" * 80\n    )\n\n    train_dataset = (\n        KneeSwinDataset(\n            train_df,\n            train_series_groups,\n            TRAIN_SERIES_DIR,\n            training=True\n        )\n    )\n\n    valid_dataset = (\n        KneeSwinDataset(\n            valid_df,\n            train_series_groups,\n            TRAIN_SERIES_DIR,\n            training=False\n        )\n    )\n\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=BATCH_SIZE,\n        shuffle=True,\n        num_workers=NUM_WORKERS,\n        pin_memory=True,\n        persistent_workers=(\n            NUM_WORKERS > 0\n        )\n    )\n\n    valid_loader = DataLoader(\n        valid_dataset,\n        batch_size=BATCH_SIZE,\n        shuffle=False,\n        num_workers=NUM_WORKERS,\n        pin_memory=True,\n        persistent_workers=(\n            NUM_WORKERS > 0\n        )\n    )\n\n    model = SwinKneeModel(\n        num_classes=NUM_CLASSES,\n        pretrained=True\n    )\n\n    model = model.to(\n        DEVICE\n    )\n\n    # --------------------------------------------------------\n    # BCE loss.\n    #\n    # Because the evaluation metric is ROC-AUC, probabilities\n    # rather than hard predictions are used.\n    # --------------------------------------------------------\n\n    criterion = (\n        nn.BCEWithLogitsLoss()\n    )\n\n    # --------------------------------------------------------\n    # AdamW\n    # --------------------------------------------------------\n\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=LR,\n        weight_decay=WEIGHT_DECAY\n    )\n\n    # --------------------------------------------------------\n    # Warmup + cosine.\n    # --------------------------------------------------------\n\n    total_steps = (\n        len(train_loader)\n        * EPOCHS\n        // ACCUM_STEPS\n    )\n\n    warmup_steps = max(\n        10,\n        int(\n            total_steps * 0.10\n        )\n    )\n\n    def lr_lambda(step):\n\n        if step < warmup_steps:\n\n            return (\n                float(step + 1)\n                / float(warmup_steps)\n            )\n\n        progress = (\n            step - warmup_steps\n        ) / max(\n            1,\n            total_steps\n            - warmup_steps\n        )\n\n        return 0.5 * (\n            1.0\n            + math.cos(\n                math.pi * progress\n            )\n        )\n\n    scheduler = (\n        torch.optim.lr_scheduler.LambdaLR(\n            optimizer,\n            lr_lambda\n        )\n    )\n\n    scaler = torch.cuda.amp.GradScaler(\n        enabled=(\n            USE_AMP\n            and DEVICE.type == \"cuda\"\n        )\n    )\n\n    best_auc = -np.inf\n\n    best_state = None\n\n    global_step = 0\n\n    # --------------------------------------------------------\n    # Epoch loop\n    # --------------------------------------------------------\n\n    for epoch in range(\n        EPOCHS\n    ):\n\n        model.train()\n\n        optimizer.zero_grad(\n            set_to_none=True\n        )\n\n        running_loss = 0.0\n\n        for step, (\n            images,\n            targets\n        ) in enumerate(\n            train_loader\n        ):\n\n            images = images.to(\n                DEVICE,\n                non_blocking=True\n            )\n\n            targets = targets.to(\n                DEVICE,\n                non_blocking=True\n            )\n\n            with torch.cuda.amp.autocast(\n                enabled=(\n                    USE_AMP\n                    and DEVICE.type == \"cuda\"\n                )\n            ):\n\n                logits = model(\n                    images\n                )\n\n                loss = criterion(\n                    logits,\n                    targets\n                )\n\n                loss = (\n                    loss\n                    / ACCUM_STEPS\n                )\n\n            scaler.scale(\n                loss\n            ).backward()\n\n            if (\n                (step + 1)\n                % ACCUM_STEPS\n                == 0\n            ):\n\n                scaler.unscale_(\n                    optimizer\n                )\n\n                torch.nn.utils.clip_grad_norm_(\n                    model.parameters(),\n                    1.0\n                )\n\n                scaler.step(\n                    optimizer\n                )\n\n                scaler.update()\n\n                optimizer.zero_grad(\n                    set_to_none=True\n                )\n\n                scheduler.step()\n\n                global_step += 1\n\n            running_loss += (\n                loss.item()\n                * ACCUM_STEPS\n            )\n\n        # ----------------------------------------------------\n        # Validation\n        # ----------------------------------------------------\n\n        model.eval()\n\n        valid_true = []\n        valid_pred = []\n\n        with torch.no_grad():\n\n            for images, targets in (\n                valid_loader\n            ):\n\n                images = images.to(\n                    DEVICE,\n                    non_blocking=True\n                )\n\n                with torch.cuda.amp.autocast(\n                    enabled=(\n                        USE_AMP\n                        and DEVICE.type == \"cuda\"\n                    )\n                ):\n\n                    logits = model(\n                        images\n                    )\n\n                probabilities = (\n                    torch.sigmoid(\n                        logits\n                    )\n                    .float()\n                    .cpu()\n                    .numpy()\n                )\n\n                valid_pred.append(\n                    probabilities\n                )\n\n                valid_true.append(\n                    targets.numpy()\n                )\n\n        y_true = np.concatenate(\n            valid_true,\n            axis=0\n        )\n\n        y_pred = np.concatenate(\n            valid_pred,\n            axis=0\n        )\n\n        macro, individual = (\n            calculate_auc(\n                y_true,\n                y_pred\n            )\n        )\n\n        train_loss = (\n            running_loss\n            / len(train_loader)\n        )\n\n        print(\n            f\"Epoch {epoch+1:02d}/{EPOCHS} | \"\n            f\"loss {train_loss:.5f} | \"\n            f\"AUC {macro:.6f}\"\n        )\n\n        if macro > best_auc:\n\n            best_auc = macro\n\n            best_state = {\n                k: v.detach()\n                .cpu()\n                .clone()\n                for k, v\n                in model.state_dict().items()\n            }\n\n            torch.save(\n                best_state,\n                f\"/kaggle/working/\"\n                f\"swin_fold_{fold}.pth\"\n            )\n\n    # --------------------------------------------------------\n    # Restore best checkpoint.\n    # --------------------------------------------------------\n\n    model.load_state_dict(\n        best_state\n    )\n\n    return (\n        model,\n        best_auc\n    )\n\n\n# ============================================================\n# 23. 5-FOLD CROSS VALIDATION\n# ============================================================\n\nkf = KFold(\n    n_splits=N_FOLDS,\n    shuffle=True,\n    random_state=SEED\n)\n\noof_predictions = np.zeros(\n    (\n        len(gold),\n        NUM_CLASSES\n    ),\n    dtype=np.float32\n)\n\noof_targets = gold[\n    TARGETS\n].values.astype(\n    np.float32\n)\n\nfold_scores = []\n\n\nfor fold, (\n    train_idx,\n    valid_idx\n) in enumerate(\n    kf.split(gold)\n):\n\n    fold_train = gold.iloc[\n        train_idx\n    ].reset_index(\n        drop=True\n    )\n\n    fold_valid = gold.iloc[\n        valid_idx\n    ].reset_index(\n        drop=True\n    )\n\n    model, fold_score = (\n        train_fold(\n            fold_train,\n            fold_valid,\n            fold\n        )\n    )\n\n    fold_scores.append(\n        fold_score\n    )\n\n    # --------------------------------------------------------\n    # Generate OOF predictions.\n    # --------------------------------------------------------\n\n    valid_dataset = (\n        KneeSwinDataset(\n            fold_valid,\n            train_series_groups,\n            TRAIN_SERIES_DIR,\n            training=False\n        )\n    )\n\n    valid_loader = DataLoader(\n        valid_dataset,\n        batch_size=BATCH_SIZE,\n        shuffle=False,\n        num_workers=NUM_WORKERS,\n        pin_memory=True\n    )\n\n    model.eval()\n\n    fold_predictions = []\n\n    with torch.no_grad():\n\n        for images, _ in (\n            valid_loader\n        ):\n\n            images = images.to(\n                DEVICE,\n                non_blocking=True\n            )\n\n            with torch.cuda.amp.autocast(\n                enabled=(\n                    USE_AMP\n                    and DEVICE.type == \"cuda\"\n                )\n            ):\n\n                logits = model(\n                    images\n                )\n\n            probabilities = (\n                torch.sigmoid(\n                    logits\n                )\n                .float()\n                .cpu()\n                .numpy()\n            )\n\n            fold_predictions.append(\n                probabilities\n            )\n\n    fold_predictions = np.concatenate(\n        fold_predictions,\n        axis=0\n    )\n\n    oof_predictions[\n        valid_idx\n    ] = fold_predictions\n\n    del model\n    gc.collect()\n\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\n# ============================================================\n# 24. OOF SCORE\n# ============================================================\n\noverall_auc, per_target_auc = (\n    calculate_auc(\n        oof_targets,\n        oof_predictions\n    )\n)\n\nprint(\n    \"\\n\"\n    + \"=\" * 80\n)\n\nprint(\n    \"FINAL OOF RESULTS\"\n)\n\nprint(\n    \"=\" * 80\n)\n\nprint(\n    f\"\\nMacro ROC-AUC: \"\n    f\"{overall_auc:.6f}\"\n)\n\nprint()\n\nfor target, score in zip(\n    TARGETS,\n    per_target_auc\n):\n\n    print(\n        f\"{target:20s} \"\n        f\"{score:.6f}\"\n    )\n\nprint(\n    \"\\nFold scores:\"\n)\n\nfor i, score in enumerate(\n    fold_scores\n):\n\n    print(\n        f\"Fold {i}: \"\n        f\"{score:.6f}\"\n    )\n\n\n# ============================================================\n# 25. SAVE OOF\n# ============================================================\n\noof_df = gold[\n    [\n        \"StudyInstanceUID\"\n    ]\n].copy()\n\nfor i, target in enumerate(\n    TARGETS\n):\n\n    oof_df[target] = (\n        oof_predictions[\n            :, i\n        ]\n    )\n\noof_df.to_csv(\n    \"/kaggle/working/\"\n    \"swin_oof.csv\",\n    index=False\n)\n\n\n# ============================================================\n# 26. FINAL MODEL\n#\n# Train on all 58 gold-labelled studies.\n# ============================================================\n\nprint(\n    \"\\n\"\n    + \"=\" * 80\n)\n\nprint(\n    \"FINAL SWIN TRAINING\"\n)\n\nprint(\n    \"=\" * 80\n)\n\n\nfull_dataset = (\n    KneeSwinDataset(\n        gold,\n        train_series_groups,\n        TRAIN_SERIES_DIR,\n        training=True\n    )\n)\n\nfull_loader = DataLoader(\n    full_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=NUM_WORKERS,\n    pin_memory=True,\n    persistent_workers=(\n        NUM_WORKERS > 0\n    )\n)\n\n\nfinal_model = SwinKneeModel(\n    num_classes=NUM_CLASSES,\n    pretrained=True\n)\n\nfinal_model = final_model.to(\n    DEVICE\n)\n\n\ncriterion = (\n    nn.BCEWithLogitsLoss()\n)\n\noptimizer = torch.optim.AdamW(\n    final_model.parameters(),\n    lr=LR,\n    weight_decay=WEIGHT_DECAY\n)\n\ntotal_steps = (\n    len(full_loader)\n    * EPOCHS\n    // ACCUM_STEPS\n)\n\nscheduler = (\n    torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer,\n        T_max=max(\n            1,\n            total_steps\n        )\n    )\n)\n\nscaler = torch.cuda.amp.GradScaler(\n    enabled=(\n        USE_AMP\n        and DEVICE.type == \"cuda\"\n    )\n)\n\n\nfor epoch in range(\n    EPOCHS\n):\n\n    final_model.train()\n\n    optimizer.zero_grad(\n        set_to_none=True\n    )\n\n    running_loss = 0.0\n\n    for step, (\n        images,\n        targets\n    ) in enumerate(\n        full_loader\n    ):\n\n        images = images.to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        targets = targets.to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        with torch.cuda.amp.autocast(\n            enabled=(\n                USE_AMP\n                and DEVICE.type == \"cuda\"\n            )\n        ):\n\n            logits = final_model(\n                images\n            )\n\n            loss = criterion(\n                logits,\n                targets\n            )\n\n            loss = (\n                loss\n                / ACCUM_STEPS\n            )\n\n        scaler.scale(\n            loss\n        ).backward()\n\n        if (\n            (step + 1)\n            % ACCUM_STEPS\n            == 0\n        ):\n\n            scaler.unscale_(\n                optimizer\n            )\n\n            torch.nn.utils.clip_grad_norm_(\n                final_model.parameters(),\n                1.0\n            )\n\n            scaler.step(\n                optimizer\n            )\n\n            scaler.update()\n\n            optimizer.zero_grad(\n                set_to_none=True\n            )\n\n            scheduler.step()\n\n        running_loss += (\n            loss.item()\n            * ACCUM_STEPS\n        )\n\n    print(\n        f\"Final epoch \"\n        f\"{epoch+1:02d}/{EPOCHS} | \"\n        f\"loss=\"\n        f\"{running_loss / len(full_loader):.5f}\"\n    )\n\n\ntorch.save(\n    final_model.state_dict(),\n    \"/kaggle/working/\"\n    \"swin_final.pth\"\n)\n\n\n# ============================================================\n# 27. TEST DATASET\n# ============================================================\n\ntest_dataset = (\n    KneeSwinDataset(\n        test,\n        test_series_groups,\n        TEST_SERIES_DIR,\n        training=False\n    )\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=True,\n    persistent_workers=(\n        NUM_WORKERS > 0\n    )\n)\n\n\n# ============================================================\n# 28. TEST INFERENCE\n# ============================================================\n\nprint(\n    \"\\n\"\n    + \"=\" * 80\n)\n\nprint(\n    \"TEST INFERENCE\"\n)\n\nprint(\n    \"=\" * 80\n)\n\n\nfinal_model.eval()\n\ntest_predictions = []\n\ntest_ids = []\n\n\nwith torch.no_grad():\n\n    for images, ids in (\n        test_loader\n    ):\n\n        images = images.to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        with torch.cuda.amp.autocast(\n            enabled=(\n                USE_AMP\n                and DEVICE.type == \"cuda\"\n            )\n        ):\n\n            logits = final_model(\n                images\n            )\n\n        probabilities = (\n            torch.sigmoid(\n                logits\n            )\n            .float()\n            .cpu()\n            .numpy()\n        )\n\n        test_predictions.append(\n            probabilities\n        )\n\n        test_ids.extend(\n            list(ids)\n        )\n\n\ntest_predictions = np.concatenate(\n    test_predictions,\n    axis=0\n)\n\n\n# ============================================================\n# 29. SUBMISSION\n# ============================================================\n\nsubmission = pd.DataFrame(\n    {\n        \"StudyInstanceUID\":\n            test_ids\n    }\n)\n\n\nfor i, target in enumerate(\n    TARGETS\n):\n\n    submission[target] = (\n        test_predictions[:, i]\n    )\n\n\nsubmission = submission[\n    [\n        \"StudyInstanceUID\"\n    ] + TARGETS\n]\n\n\n# ============================================================\n# 30. SAFETY CHECKS\n# ============================================================\n\nassert (\n    len(submission)\n    == len(test)\n)\n\nassert (\n    list(submission.columns)\n    ==\n    [\n        \"StudyInstanceUID\"\n    ] + TARGETS\n)\n\nassert (\n    submission[TARGETS]\n    .notna()\n    .all()\n    .all()\n)\n\nassert (\n    submission[TARGETS]\n    .values >= 0\n).all()\n\nassert (\n    submission[TARGETS]\n    .values <= 1\n).all()\n\n\n# ============================================================\n# 31. SAVE\n# ============================================================\n\nSUBMISSION_PATH = (\n    \"/kaggle/working/\"\n    \"submission.csv\"\n)\n\nsubmission.to_csv(\n    SUBMISSION_PATH,\n    index=False\n)\n\n\n# ============================================================\n# 32. FINAL REPORT\n# ============================================================\n\nprint(\n    \"\\n\"\n    + \"=\" * 80\n)\n\nprint(\n    \"SUBMISSION READY\"\n)\n\nprint(\n    \"=\" * 80\n)\n\nprint(\n    \"\\nFile:\"\n)\n\nprint(\n    SUBMISSION_PATH\n)\n\nprint(\n    \"\\nShape:\",\n    submission.shape\n)\n\nprint(\n    \"\\nColumns:\"\n)\n\nprint(\n    submission.columns.tolist()\n)\n\nprint(\n    \"\\nFirst 5 rows:\"\n)\n\ndisplay(\n    submission.head()\n)\n\nprint(\n    \"\\nPrediction statistics:\"\n)\n\ndisplay(\n    submission[\n        TARGETS\n    ].describe().T\n)\n\nprint(\n    \"\\nDONE.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T09:02:23.709216Z","iopub.execute_input":"2026-09-14T09:02:23.709747Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}