{"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":"import os\nimport re\nimport gc\nimport math\nimport random\nimport warnings\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\n\nimport pydicom\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nimport torchvision\nfrom torchvision import transforms\n\nwarnings.filterwarnings(\"ignore\")\n\nprint(\"PyTorch:\", torch.__version__)\nprint(\"Torchvision:\", torchvision.__version__)\nprint(\"CUDA:\", torch.cuda.is_available())\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nif torch.cuda.is_available():\n    print(\"GPU:\", torch.cuda.get_device_name(0))\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:01:30.273066Z","iopub.execute_input":"2026-08-13T04:01:30.273731Z","iopub.status.idle":"2026-08-13T04:01:40.804887Z","shell.execute_reply.started":"2026-08-13T04:01:30.273704Z","shell.execute_reply":"2026-08-13T04:01:40.804191Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 42\n\ndef seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n\nseed_everything(SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:03:10.844468Z","iopub.execute_input":"2026-08-13T04:03:10.845111Z","iopub.status.idle":"2026-08-13T04:03:10.854664Z","shell.execute_reply.started":"2026-08-13T04:03:10.845073Z","shell.execute_reply":"2026-08-13T04:03:10.853869Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\n\nINPUT_ROOT = Path(\"/kaggle/input\")\n\ntrain_csv_candidates = list(INPUT_ROOT.rglob(\"train.csv\"))\n\nif not train_csv_candidates:\n    raise FileNotFoundError(\"train.csv not found\")\n\nDATA_DIR = train_csv_candidates[0].parent\n\nTRAIN_CSV = DATA_DIR / \"train.csv\"\nTRAIN_SERIES_CSV = DATA_DIR / \"train_series.csv\"\n\nTEST_CSV = DATA_DIR / \"test.csv\"\nTEST_SERIES_CSV = DATA_DIR / \"test_series.csv\"\n\nSAMPLE_SUBMISSION = DATA_DIR / \"sample_submission.csv\"\n\nTRAIN_SERIES_DIR = DATA_DIR / \"train_series\"\nTEST_SERIES_DIR = DATA_DIR / \"test_series\"\n\nprint(DATA_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:03:21.104233Z","iopub.execute_input":"2026-08-13T04:03:21.104512Z","iopub.status.idle":"2026-08-13T04:09:31.020609Z","shell.execute_reply.started":"2026-08-13T04:03:21.10449Z","shell.execute_reply":"2026-08-13T04:09:31.019689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = pd.read_csv(TRAIN_CSV)\ntrain_series = pd.read_csv(TRAIN_SERIES_CSV)\n\ntest = pd.read_csv(TEST_CSV)\ntest_series = pd.read_csv(TEST_SERIES_CSV)\n\nsample_submission = pd.read_csv(SAMPLE_SUBMISSION)\n\nLABELS = [\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\nprint(\"Train:\", train.shape)\nprint(\"Train series:\", train_series.shape)\nprint(\"Test:\", test.shape)\nprint(\"Test series:\", test_series.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:10:24.596496Z","iopub.execute_input":"2026-08-13T04:10:24.597034Z","iopub.status.idle":"2026-08-13T04:10:24.856803Z","shell.execute_reply.started":"2026-08-13T04:10:24.597003Z","shell.execute_reply":"2026-08-13T04:10:24.855836Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"series_meta = train_series.copy()\n\nseries_meta[\"study_path\"] = (\n    series_meta[\"StudyInstanceUID\"]\n    .astype(str)\n)\n\nseries_meta[\"series_path\"] = (\n    series_meta[\"StudyInstanceUID\"].astype(str)\n    + \"/\"\n    + series_meta[\"SeriesInstanceUID\"].astype(str)\n)\n\nseries_meta.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:10:49.262004Z","iopub.execute_input":"2026-08-13T04:10:49.262291Z","iopub.status.idle":"2026-08-13T04:10:49.311736Z","shell.execute_reply.started":"2026-08-13T04:10:49.262267Z","shell.execute_reply":"2026-08-13T04:10:49.311026Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NEGATIVE_PATTERNS = [\n    r\"\\bno\\b\",\n    r\"\\bnot\\b\",\n    r\"\\bwithout\\b\",\n    r\"\\bnegative for\\b\",\n    r\"\\bno evidence of\\b\",\n    r\"\\babsence of\\b\",\n    r\"\\bintact\\b\",\n    r\"\\bnormal\\b\",\n    r\"\\bunremarkable\\b\",\n    r\"\\bpreserved\\b\",\n    r\"\\bno definite\\b\",\n    r\"\\bno significant\\b\",\n]\n\nTARGET_PATTERNS = {\n\n    \"ACL\": [\n        r\"\\bacl\\b\",\n        r\"anterior cruciate ligament\"\n    ],\n\n    \"MCL\": [\n        r\"\\bmcl\\b\",\n        r\"medial collateral ligament\"\n    ],\n\n    \"Medial Meniscus\": [\n        r\"medial meniscus\"\n    ],\n\n    \"Lateral Meniscus\": [\n        r\"lateral meniscus\"\n    ],\n\n    \"Medial OA\": [\n        r\"medial compartment.*(?:osteoarthritis|arthrosis|degenerative)\",\n        r\"medial.*(?:osteoarthritis|arthrosis)\"\n    ],\n\n    \"Lateral OA\": [\n        r\"lateral compartment.*(?:osteoarthritis|arthrosis|degenerative)\",\n        r\"lateral.*(?:osteoarthritis|arthrosis)\"\n    ],\n\n    \"PF OA\": [\n        r\"patellofemoral.*(?:osteoarthritis|arthrosis|degenerative)\",\n        r\"patellofemoral.*(?:cartilage loss|chondral)\"\n    ],\n\n    \"Effusion\": [\n        r\"\\beffusion\\b\",\n        r\"joint fluid\"\n    ],\n\n    \"Synovitis\": [\n        r\"\\bsynovitis\\b\",\n        r\"synovial thickening\"\n    ],\n\n    \"Baker's\": [\n        r\"baker.?s cyst\",\n        r\"popliteal cyst\"\n    ],\n\n    \"Contusion\": [\n        r\"bone contusion\",\n        r\"bone bruise\",\n        r\"marrow edema\",\n        r\"bone marrow edema\"\n    ],\n\n    \"Fracture\": [\n        r\"\\bfracture\\b\",\n        r\"\\bfractured\\b\"\n    ]\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:11:01.381613Z","iopub.execute_input":"2026-08-13T04:11:01.382417Z","iopub.status.idle":"2026-08-13T04:11:01.388791Z","shell.execute_reply.started":"2026-08-13T04:11:01.382388Z","shell.execute_reply":"2026-08-13T04:11:01.387907Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def normalize_report(text):\n    text = str(text).lower()\n    text = text.replace(\"\\n\", \" \")\n    text = re.sub(r\"\\s+\", \" \", text)\n    return text.strip()\n\n\ndef split_sentences(text):\n    return re.split(r\"(?<=[.!?;])\\s+\", text)\n\n\ndef sentence_has_negation(sentence, match_start):\n    prefix = sentence[:match_start]\n\n    # Restrict negation search to nearby context.\n    prefix = prefix[-100:]\n\n    for pattern in NEGATIVE_PATTERNS:\n        if re.search(pattern, prefix):\n            return True\n\n    return False\n\n\ndef weak_label_report(report):\n\n    report = normalize_report(report)\n\n    result = {}\n\n    sentences = split_sentences(report)\n\n    for label, patterns in TARGET_PATTERNS.items():\n\n        found_positive = False\n        found_negative = False\n\n        for sentence in sentences:\n\n            for pattern in patterns:\n\n                match = re.search(pattern, sentence)\n\n                if match:\n\n                    if sentence_has_negation(\n                        sentence,\n                        match.start()\n                    ):\n                        found_negative = True\n                    else:\n                        found_positive = True\n\n        if found_positive and not found_negative:\n            result[label] = 1.0\n\n        elif found_negative and not found_positive:\n            result[label] = 0.0\n\n        else:\n            result[label] = np.nan\n\n    return result","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:11:31.207007Z","iopub.execute_input":"2026-08-13T04:11:31.207554Z","iopub.status.idle":"2026-08-13T04:11:31.215002Z","shell.execute_reply.started":"2026-08-13T04:11:31.207525Z","shell.execute_reply":"2026-08-13T04:11:31.214038Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"weak_labels = []\n\nfor _, row in train.iterrows():\n\n    result = weak_label_report(row[\"Report\"])\n\n    result[\"StudyInstanceUID\"] = row[\"StudyInstanceUID\"]\n\n    weak_labels.append(result)\n\nweak_labels = pd.DataFrame(weak_labels)\n\nweak_labels = weak_labels[\n    [\"StudyInstanceUID\"] + LABELS\n]\n\nprint(weak_labels.shape)\n\ndisplay(weak_labels.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:11:51.116201Z","iopub.execute_input":"2026-08-13T04:11:51.116554Z","iopub.status.idle":"2026-08-13T04:11:53.459124Z","shell.execute_reply.started":"2026-08-13T04:11:51.116527Z","shell.execute_reply":"2026-08-13T04:11:53.458434Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_data = train[\n    [\"StudyInstanceUID\"] + LABELS\n].copy()\n\nlabel_data = label_data.merge(\n    weak_labels,\n    on=\"StudyInstanceUID\",\n    how=\"left\",\n    suffixes=(\"_official\", \"_weak\")\n)\n\nfor label in LABELS:\n\n    official = label_data[f\"{label}_official\"]\n    weak = label_data[f\"{label}_weak\"]\n\n    # Official labels have priority.\n    label_data[label] = official.combine_first(weak)\n\nlabel_data = label_data[\n    [\"StudyInstanceUID\"] + LABELS\n]\n\ndisplay(label_data.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:12:04.561343Z","iopub.execute_input":"2026-08-13T04:12:04.561967Z","iopub.status.idle":"2026-08-13T04:12:04.597313Z","shell.execute_reply.started":"2026-08-13T04:12:04.561938Z","shell.execute_reply":"2026-08-13T04:12:04.596703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"coverage = pd.DataFrame({\n    \"Known labels\":\n        label_data[LABELS].notna().sum(),\n\n    \"Unknown labels\":\n        label_data[LABELS].isna().sum()\n})\n\ncoverage[\"Coverage %\"] = (\n    coverage[\"Known labels\"] /\n    len(label_data) * 100\n).round(2)\n\ndisplay(coverage)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:12:15.636521Z","iopub.execute_input":"2026-08-13T04:12:15.637374Z","iopub.status.idle":"2026-08-13T04:12:15.65598Z","shell.execute_reply.started":"2026-08-13T04:12:15.637333Z","shell.execute_reply":"2026-08-13T04:12:15.655092Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_array = label_data[LABELS].values.astype(np.float32)\n\nlabel_mask = ~np.isnan(label_array)\n\nlabel_array = np.nan_to_num(\n    label_array,\n    nan=0.0\n).astype(np.float32)\n\nlabel_mask = label_mask.astype(np.float32)\n\nprint(\"Known target entries:\",\n      label_mask.sum())\n\nprint(\"Total target entries:\",\n      label_mask.size)\n\nprint(\n    \"Coverage:\",\n    label_mask.mean() * 100,\n    \"%\"\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:12:31.207742Z","iopub.execute_input":"2026-08-13T04:12:31.208024Z","iopub.status.idle":"2026-08-13T04:12:31.218418Z","shell.execute_reply.started":"2026-08-13T04:12:31.208001Z","shell.execute_reply":"2026-08-13T04:12:31.217419Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def series_priority(row):\n\n    score = 0\n\n    if row[\"Fluid_Sensitive\"] == 1:\n        score += 5\n\n    if row[\"Fat_Suppression\"] == 1:\n        score += 3\n\n    if row[\"Anatomical_Plane\"] == \"Sagittal\":\n        score += 3\n\n    elif row[\"Anatomical_Plane\"] == \"Coronal\":\n        score += 2\n\n    elif row[\"Anatomical_Plane\"] == \"Axial\":\n        score += 1\n\n    return score\n\n\ntrain_series[\"priority\"] = train_series.apply(\n    series_priority,\n    axis=1\n)\n\ntest_series[\"priority\"] = test_series.apply(\n    series_priority,\n    axis=1\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:12:43.356999Z","iopub.execute_input":"2026-08-13T04:12:43.357426Z","iopub.status.idle":"2026-08-13T04:12:43.582599Z","shell.execute_reply.started":"2026-08-13T04:12:43.357397Z","shell.execute_reply":"2026-08-13T04:12:43.581783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MAX_SERIES_PER_STUDY = 4\n\n\ndef select_series(series_df):\n\n    selected = {}\n\n    for study_id, group in series_df.groupby(\n        \"StudyInstanceUID\"\n    ):\n\n        group = group.sort_values(\n            \"priority\",\n            ascending=False\n        )\n\n        selected[study_id] = group.head(\n            MAX_SERIES_PER_STUDY\n        ).to_dict(\"records\")\n\n    return selected\n\n\ntrain_selected = select_series(train_series)\ntest_selected = select_series(test_series)\n\nprint(\n    \"Training studies with series:\",\n    len(train_selected)\n)\n\nprint(\n    \"Test studies with series:\",\n    len(test_selected)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:13:20.102039Z","iopub.execute_input":"2026-08-13T04:13:20.102833Z","iopub.status.idle":"2026-08-13T04:13:23.74397Z","shell.execute_reply.started":"2026-08-13T04:13:20.102804Z","shell.execute_reply":"2026-08-13T04:13:23.743098Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def read_dicom_image(path):\n\n    ds = pydicom.dcmread(path)\n\n    image = ds.pixel_array.astype(np.float32)\n\n    # Apply rescale slope/intercept when present.\n    slope = float(\n        getattr(ds, \"RescaleSlope\", 1.0)\n    )\n\n    intercept = float(\n        getattr(ds, \"RescaleIntercept\", 0.0)\n    )\n\n    image = image * slope + intercept\n\n    # MONOCHROME1 needs inversion.\n    if getattr(\n        ds,\n        \"PhotometricInterpretation\",\n        \"\"\n    ) == \"MONOCHROME1\":\n\n        image = image.max() - image\n\n    # Robust intensity normalization.\n    low, high = np.percentile(\n        image,\n        [1, 99]\n    )\n\n    if high > low:\n\n        image = np.clip(\n            image,\n            low,\n            high\n        )\n\n        image = (\n            image - low\n        ) / (high - low)\n\n    else:\n\n        image = np.zeros_like(image)\n\n    return image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:13:39.988402Z","iopub.execute_input":"2026-08-13T04:13:39.989096Z","iopub.status.idle":"2026-08-13T04:13:39.995274Z","shell.execute_reply.started":"2026-08-13T04:13:39.989066Z","shell.execute_reply":"2026-08-13T04:13:39.994328Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def dicom_sort_key(path):\n\n    try:\n\n        ds = pydicom.dcmread(\n            path,\n            stop_before_pixels=True\n        )\n\n        if hasattr(\n            ds,\n            \"ImagePositionPatient\"\n        ):\n\n            pos = ds.ImagePositionPatient\n\n            if hasattr(\n                ds,\n                \"ImageOrientationPatient\"\n            ):\n\n                orientation = np.array(\n                    ds.ImageOrientationPatient,\n                    dtype=float\n                )\n\n                row = orientation[:3]\n                col = orientation[3:]\n\n                normal = np.cross(\n                    row,\n                    col\n                )\n\n                location = np.dot(\n                    np.array(pos, dtype=float),\n                    normal\n                )\n\n                return (\n                    0,\n                    float(location)\n                )\n\n        return (\n            1,\n            float(\n                getattr(\n                    ds,\n                    \"InstanceNumber\",\n                    0\n                )\n            )\n        )\n\n    except Exception:\n\n        return (\n            2,\n            str(path)\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:13:52.842311Z","iopub.execute_input":"2026-08-13T04:13:52.842788Z","iopub.status.idle":"2026-08-13T04:13:52.849342Z","shell.execute_reply.started":"2026-08-13T04:13:52.84276Z","shell.execute_reply":"2026-08-13T04:13:52.848495Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SLICES_PER_SERIES = 8\n\n\ndef get_series_files(series_path):\n\n    files = list(\n        Path(series_path).glob(\"*.dcm\")\n    )\n\n    if not files:\n        return []\n\n    files = sorted(\n        files,\n        key=dicom_sort_key\n    )\n\n    return files\n\n\ndef sample_series_files(\n    series_path,\n    n=SLICES_PER_SERIES\n):\n\n    files = get_series_files(series_path)\n\n    if not files:\n        return []\n\n    if len(files) <= n:\n        return files\n\n    indices = np.linspace(\n        0,\n        len(files) - 1,\n        n\n    ).astype(int)\n\n    return [\n        files[i]\n        for i in indices\n    ]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:14:08.271716Z","iopub.execute_input":"2026-08-13T04:14:08.272579Z","iopub.status.idle":"2026-08-13T04:14:08.278791Z","shell.execute_reply.started":"2026-08-13T04:14:08.272552Z","shell.execute_reply":"2026-08-13T04:14:08.277981Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMG_SIZE = 224\n\nimage_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize(\n        (IMG_SIZE, IMG_SIZE)\n    ),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:14:17.557179Z","iopub.execute_input":"2026-08-13T04:14:17.557844Z","iopub.status.idle":"2026-08-13T04:14:17.563223Z","shell.execute_reply.started":"2026-08-13T04:14:17.55781Z","shell.execute_reply":"2026-08-13T04:14:17.562301Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class KneeStudyDataset(Dataset):\n\n    def __init__(\n        self,\n        studies,\n        selected_series,\n        series_root,\n        labels=None,\n        masks=None,\n        training=True\n    ):\n\n        self.studies = list(studies)\n\n        self.selected_series = selected_series\n\n        self.series_root = Path(\n            series_root\n        )\n\n        self.labels = labels\n        self.masks = masks\n\n        self.training = training\n\n\n    def __len__(self):\n        return len(self.studies)\n\n\n    def _load_study_images(\n        self,\n        study_id\n    ):\n\n        all_images = []\n\n        series_list = self.selected_series.get(\n            study_id,\n            []\n        )\n\n        for series_info in series_list:\n\n            series_id = (\n                series_info[\"SeriesInstanceUID\"]\n            )\n\n            series_path = (\n                self.series_root\n                / str(study_id)\n                / str(series_id)\n            )\n\n            files = sample_series_files(\n                series_path\n            )\n\n            for file in files:\n\n                try:\n\n                    image = read_dicom_image(\n                        file\n                    )\n\n                    tensor = image_transform(\n                        image\n                    )\n\n                    # grayscale -> 3 channels\n                    tensor = tensor.repeat(\n                        3,\n                        1,\n                        1\n                    )\n\n                    all_images.append(\n                        tensor\n                    )\n\n                except Exception:\n                    continue\n\n        # Safety fallback\n        if not all_images:\n\n            return torch.zeros(\n                1,\n                3,\n                IMG_SIZE,\n                IMG_SIZE\n            )\n\n        return torch.stack(\n            all_images\n        )\n\n\n    def __getitem__(self, idx):\n\n        study_id = self.studies[idx]\n\n        images = self._load_study_images(\n            study_id\n        )\n\n        if self.labels is not None:\n\n            label_idx = self.studies.index(\n                study_id\n            )\n\n            target = torch.tensor(\n                self.labels[label_idx],\n                dtype=torch.float32\n            )\n\n            mask = torch.tensor(\n                self.masks[label_idx],\n                dtype=torch.float32\n            )\n\n            return (\n                images,\n                target,\n                mask,\n                study_id\n            )\n\n        return images, study_id","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:14:29.003765Z","iopub.execute_input":"2026-08-13T04:14:29.00421Z","iopub.status.idle":"2026-08-13T04:14:29.014124Z","shell.execute_reply.started":"2026-08-13T04:14:29.00418Z","shell.execute_reply":"2026-08-13T04:14:29.013336Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\nofficial_mask = train[LABELS].notna().all(axis=1)\n\nofficial_ids = train.loc[\n    official_mask,\n    \"StudyInstanceUID\"\n].tolist()\n\ntrain_ids_all = train[\n    \"StudyInstanceUID\"\n].tolist()\n\ntrain_ids, valid_ids = train_test_split(\n    train_ids_all,\n    test_size=0.15,\n    random_state=SEED\n)\n\nprint(\"Training studies:\", len(train_ids))\nprint(\"Validation studies:\", len(valid_ids))\nprint(\"Official fully-labelled:\", len(official_ids))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:14:46.842069Z","iopub.execute_input":"2026-08-13T04:14:46.843049Z","iopub.status.idle":"2026-08-13T04:14:47.870196Z","shell.execute_reply.started":"2026-08-13T04:14:46.843008Z","shell.execute_reply":"2026-08-13T04:14:47.869414Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_lookup = label_data.set_index(\n    \"StudyInstanceUID\"\n)\n\ndef get_targets(study_ids):\n\n    targets = []\n    masks = []\n\n    for study_id in study_ids:\n\n        row = label_lookup.loc[\n            study_id,\n            LABELS\n        ].values.astype(np.float32)\n\n        mask = ~np.isnan(row)\n\n        row = np.nan_to_num(\n            row,\n            nan=0.0\n        )\n\n        targets.append(row)\n        masks.append(mask.astype(np.float32))\n\n    return (\n        np.array(targets),\n        np.array(masks)\n    )\n\n\ntrain_targets, train_masks = get_targets(\n    train_ids\n)\n\nvalid_targets, valid_masks = get_targets(\n    valid_ids\n)\n\nprint(train_targets.shape)\nprint(train_masks.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:14:57.235781Z","iopub.execute_input":"2026-08-13T04:14:57.236221Z","iopub.status.idle":"2026-08-13T04:14:58.895397Z","shell.execute_reply.started":"2026-08-13T04:14:57.236196Z","shell.execute_reply":"2026-08-13T04:14:58.89471Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = KneeStudyDataset(\n    studies=train_ids,\n    selected_series=train_selected,\n    series_root=TRAIN_SERIES_DIR,\n    labels=train_targets,\n    masks=train_masks,\n    training=True\n)\n\nvalid_dataset = KneeStudyDataset(\n    studies=valid_ids,\n    selected_series=train_selected,\n    series_root=TRAIN_SERIES_DIR,\n    labels=valid_targets,\n    masks=valid_masks,\n    training=False\n)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=1,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\nvalid_loader = DataLoader(\n    valid_dataset,\n    batch_size=1,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:15:10.801567Z","iopub.execute_input":"2026-08-13T04:15:10.802358Z","iopub.status.idle":"2026-08-13T04:15:10.807621Z","shell.execute_reply.started":"2026-08-13T04:15:10.802328Z","shell.execute_reply":"2026-08-13T04:15:10.80672Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision.models import (\n    efficientnet_b0,\n    EfficientNet_B0_Weights\n)\n\nweights = EfficientNet_B0_Weights.DEFAULT\n\nmodel = efficientnet_b0(\n    weights=weights\n)\n\nin_features = model.classifier[\n    1\n].in_features\n\nmodel.classifier = nn.Sequential(\n    nn.Dropout(0.30),\n    nn.Linear(\n        in_features,\n        len(LABELS)\n    )\n)\n\nmodel = model.to(DEVICE)\n\nprint(\"Model ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:15:26.426476Z","iopub.execute_input":"2026-08-13T04:15:26.42702Z","iopub.status.idle":"2026-08-13T04:15:27.409741Z","shell.execute_reply.started":"2026-08-13T04:15:26.426991Z","shell.execute_reply":"2026-08-13T04:15:27.408932Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class StudyModel(nn.Module):\n\n    def __init__(\n        self,\n        backbone\n    ):\n\n        super().__init__()\n\n        self.backbone = backbone\n\n        self.feature_extractor = nn.Sequential(\n            *list(\n                self.backbone.children()\n            )[:-1]\n        )\n\n        self.pool = nn.AdaptiveAvgPool2d(\n            1\n        )\n\n        feature_dim = 1280\n\n        self.head = nn.Sequential(\n            nn.LayerNorm(feature_dim),\n            nn.Dropout(0.30),\n            nn.Linear(\n                feature_dim,\n                len(LABELS)\n            )\n        )\n\n\n    def forward(self, x):\n\n        # x:\n        # [number_of_slices, 3, H, W]\n\n        features = self.feature_extractor(x)\n\n        features = self.pool(\n            features\n        )\n\n        features = features.flatten(1)\n\n        # Study-level mean pooling\n        study_feature = features.mean(\n            dim=0,\n            keepdim=True\n        )\n\n        output = self.head(\n            study_feature\n        )\n\n        return output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:15:48.037104Z","iopub.execute_input":"2026-08-13T04:15:48.037873Z","iopub.status.idle":"2026-08-13T04:15:48.044315Z","shell.execute_reply.started":"2026-08-13T04:15:48.037843Z","shell.execute_reply":"2026-08-13T04:15:48.043272Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"study_model = StudyModel(\n    model\n).to(DEVICE)\n\nprint(study_model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:16:48.317132Z","iopub.execute_input":"2026-08-13T04:16:48.317771Z","iopub.status.idle":"2026-08-13T04:16:48.336375Z","shell.execute_reply.started":"2026-08-13T04:16:48.317739Z","shell.execute_reply":"2026-08-13T04:16:48.335713Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MaskedBCELoss(nn.Module):\n\n    def __init__(self):\n        super().__init__()\n\n\n    def forward(\n        self,\n        logits,\n        targets,\n        mask\n    ):\n\n        loss = F.binary_cross_entropy_with_logits(\n            logits,\n            targets,\n            reduction=\"none\"\n        )\n\n        loss = loss * mask\n\n        denominator = mask.sum().clamp_min(\n            1.0\n        )\n\n        return loss.sum() / denominator","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:16:58.4083Z","iopub.execute_input":"2026-08-13T04:16:58.408785Z","iopub.status.idle":"2026-08-13T04:16:58.414274Z","shell.execute_reply.started":"2026-08-13T04:16:58.408756Z","shell.execute_reply":"2026-08-13T04:16:58.413374Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = MaskedBCELoss()\n\noptimizer = torch.optim.AdamW(\n    study_model.parameters(),\n    lr=2e-4,\n    weight_decay=1e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=3\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:17:00.98231Z","iopub.execute_input":"2026-08-13T04:17:00.983149Z","iopub.status.idle":"2026-08-13T04:17:00.988533Z","shell.execute_reply.started":"2026-08-13T04:17:00.983109Z","shell.execute_reply":"2026-08-13T04:17:00.987924Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 3\n\n\ndef train_one_epoch(\n    model,\n    loader\n):\n\n    model.train()\n\n    running_loss = 0.0\n    count = 0\n\n    for batch in loader:\n\n        images, targets, masks, study_ids = batch\n\n        images = images.squeeze(0).to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        targets = targets.to(\n            DEVICE\n        )\n\n        masks = masks.to(\n            DEVICE\n        )\n\n        optimizer.zero_grad(\n            set_to_none=True\n        )\n\n        logits = model(\n            images\n        )\n\n        loss = criterion(\n            logits,\n            targets,\n            masks\n        )\n\n        loss.backward()\n\n        torch.nn.utils.clip_grad_norm_(\n            model.parameters(),\n            1.0\n        )\n\n        optimizer.step()\n\n        running_loss += (\n            loss.item()\n        )\n\n        count += 1\n\n    return running_loss / max(\n        count,\n        1\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:17:18.282838Z","iopub.execute_input":"2026-08-13T04:17:18.283687Z","iopub.status.idle":"2026-08-13T04:17:18.290158Z","shell.execute_reply.started":"2026-08-13T04:17:18.283615Z","shell.execute_reply":"2026-08-13T04:17:18.289136Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score\n\n\n@torch.no_grad()\ndef validate(\n    model,\n    loader\n):\n\n    model.eval()\n\n    predictions = []\n    targets_all = []\n    masks_all = []\n\n    for batch in loader:\n\n        images, targets, masks, study_ids = batch\n\n        images = images.squeeze(0).to(\n            DEVICE\n        )\n\n        logits = model(\n            images\n        )\n\n        probs = torch.sigmoid(\n            logits\n        )\n\n        predictions.append(\n            probs.cpu().numpy()[0]\n        )\n\n        targets_all.append(\n            targets.numpy()[0]\n        )\n\n        masks_all.append(\n            masks.numpy()[0]\n        )\n\n    predictions = np.array(\n        predictions\n    )\n\n    targets_all = np.array(\n        targets_all\n    )\n\n    masks_all = np.array(\n        masks_all\n    )\n\n    aucs = []\n\n    for i, label in enumerate(\n        LABELS\n    ):\n\n        valid = masks_all[:, i] > 0\n\n        if valid.sum() < 2:\n            continue\n\n        y_true = targets_all[\n            valid,\n            i\n        ]\n\n        y_pred = predictions[\n            valid,\n            i\n        ]\n\n        if len(\n            np.unique(y_true)\n        ) < 2:\n            continue\n\n        auc = roc_auc_score(\n            y_true,\n            y_pred\n        )\n\n        aucs.append(auc)\n\n        print(\n            f\"{label:20s}: {auc:.4f}\"\n        )\n\n    mean_auc = (\n        np.mean(aucs)\n        if aucs\n        else 0\n    )\n\n    return mean_auc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T04:17:31.317782Z","iopub.execute_input":"2026-08-13T04:17:31.318583Z","iopub.status.idle":"2026-08-13T04:17:31.326295Z","shell.execute_reply.started":"2026-08-13T04:17:31.318554Z","shell.execute_reply":"2026-08-13T04:17:31.325686Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_auc = -1\n\nfor epoch in range(EPOCHS):\n\n    print(\n        f\"\\n========== EPOCH {epoch + 1}/{EPOCHS} ==========\"\n    )\n\n    train_loss = train_one_epoch(\n        study_model,\n        train_loader\n    )\n\n    print(\n        f\"Train loss: {train_loss:.5f}\"\n    )\n\n    val_auc = validate(\n        study_model,\n        valid_loader\n    )\n\n    print(\n        f\"\\nMean validation ROC-AUC: \"\n        f\"{val_auc:.5f}\"\n    )\n\n    scheduler.step()\n\n    if val_auc > best_auc:\n\n        best_auc = val_auc\n\n        torch.save(\n            study_model.state_dict(),\n            \"/kaggle/working/best_model.pth\"\n        )\n\n        print(\n            \"Saved best model.\"\n        )\n\n    gc.collect()\n\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_model_path = (\n    \"/kaggle/working/best_model.pth\"\n)\n\nif os.path.exists(\n    best_model_path\n):\n\n    study_model.load_state_dict(\n        torch.load(\n            best_model_path,\n            map_location=DEVICE\n        )\n    )\n\nprint(\"Best model loaded.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_ids = test[\n    \"StudyInstanceUID\"\n].tolist()\n\ntest_dataset = KneeStudyDataset(\n    studies=test_ids,\n    selected_series=test_selected,\n    series_root=TEST_SERIES_DIR,\n    labels=None,\n    masks=None,\n    training=False\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=1,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef predict_test(\n    model,\n    loader\n):\n\n    model.eval()\n\n    predictions = []\n    ids = []\n\n    for batch in loader:\n\n        images, study_id = batch\n\n        images = images.squeeze(0).to(\n            DEVICE\n        )\n\n        logits = model(\n            images\n        )\n\n        probs = torch.sigmoid(\n            logits\n        )\n\n        predictions.append(\n            probs.cpu().numpy()[0]\n        )\n\n        ids.append(\n            study_id[0]\n        )\n\n    return (\n        np.array(predictions),\n        ids\n    )\n\n\ntest_predictions, prediction_ids = predict_test(\n    study_model,\n    test_loader\n)\n\nprint(\n    \"Prediction shape:\",\n    test_predictions.shape\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = pd.DataFrame(\n    test_predictions,\n    columns=LABELS\n)\n\nsubmission.insert(\n    0,\n    \"StudyInstanceUID\",\n    prediction_ids\n)\n\n# Make absolutely sure ordering follows test.csv.\nsubmission = (\n    test[\n        [\"StudyInstanceUID\"]\n    ]\n    .merge(\n        submission,\n        on=\"StudyInstanceUID\",\n        how=\"left\"\n    )\n)\n\nprint(submission.shape)\n\ndisplay(submission.head())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"assert len(submission) == len(test)\n\nassert (\n    submission[\"StudyInstanceUID\"]\n    .equals(\n        test[\"StudyInstanceUID\"]\n    )\n)\n\nfor label in LABELS:\n\n    assert submission[label].between(\n        0,\n        1\n    ).all()\n\nprint(\"All submission checks passed.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SUBMISSION_PATH = (\n    \"/kaggle/working/submission.csv\"\n)\n\nsubmission.to_csv(\n    SUBMISSION_PATH,\n    index=False\n)\n\nprint(\n    \"Saved:\",\n    SUBMISSION_PATH\n)\n\ndisplay(\n    pd.read_csv(\n        SUBMISSION_PATH\n    ).head()\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}