{"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\n\nINPUT_DIR = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n\nprint(\"DATASET PATH:\")\nprint(INPUT_DIR)\nprint(\"\\nFILES / FOLDERS:\")\nprint(\"=\" * 70)\n\nfor item in os.listdir(INPUT_DIR):\n    print(item)\n\nimport pandas as pd\nimport os\n\nINPUT_DIR = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n\n# Load metadata only\ntrain = pd.read_csv(f\"{INPUT_DIR}/train.csv\")\ntrain_series = pd.read_csv(f\"{INPUT_DIR}/train_series.csv\")\ntest = pd.read_csv(f\"{INPUT_DIR}/test.csv\")\ntest_series = pd.read_csv(f\"{INPUT_DIR}/test_series.csv\")\nsample_submission = pd.read_csv(f\"{INPUT_DIR}/sample_submission.csv\")\n\nprint(\"=\" * 80)\nprint(\"RSNA KNEE ABNORMALITY DETECTION - DATASET OVERVIEW\")\nprint(\"=\" * 80)\n\nprint(f\"\\nTrain studies       : {len(train):,}\")\nprint(f\"Train series        : {len(train_series):,}\")\nprint(f\"Test studies        : {len(test):,}\")\nprint(f\"Test series         : {len(test_series):,}\")\nprint(f\"Submission rows     : {len(sample_submission):,}\")\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAIN SHAPE\")\nprint(\"=\" * 80)\nprint(train.shape)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAIN SERIES SHAPE\")\nprint(\"=\" * 80)\nprint(train_series.shape)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TEST SHAPE\")\nprint(\"=\" * 80)\nprint(test.shape)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TEST SERIES SHAPE\")\nprint(\"=\" * 80)\nprint(test_series.shape)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAIN COLUMNS\")\nprint(\"=\" * 80)\n\nfor i, col in enumerate(train.columns, 1):\n    print(f\"{i:2}. {col}\")\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAIN SERIES COLUMNS\")\nprint(\"=\" * 80)\n\nfor i, col in enumerate(train_series.columns, 1):\n    print(f\"{i:2}. {col}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:49.80865Z","iopub.execute_input":"2026-08-19T04:11:49.808911Z","iopub.status.idle":"2026-08-19T04:11:49.828917Z","shell.execute_reply.started":"2026-08-19T04:11:49.808878Z","shell.execute_reply":"2026-08-19T04:11:49.828306Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_columns = [\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(\"=\" * 80)\nprint(\"LABEL DISTRIBUTION\")\nprint(\"=\" * 80)\n\nfor label in label_columns:\n    if label in train.columns:\n        counts = train[label].value_counts(dropna=False)\n        \n        print(f\"\\n{label}\")\n        print(\"-\" * 40)\n        print(counts)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:51.262668Z","iopub.execute_input":"2026-08-19T04:11:51.262977Z","iopub.status.idle":"2026-08-19T04:11:51.29171Z","shell.execute_reply.started":"2026-08-19T04:11:51.262943Z","shell.execute_reply":"2026-08-19T04:11:51.290993Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# Show 10 random reports\nprint(\"=\" * 80)\nprint(\"SAMPLE RADIOLOGY REPORTS\")\nprint(\"=\" * 80)\n\nsamples = train.sample(10, random_state=42)\n\nfor i, (_, row) in enumerate(samples.iterrows(), 1):\n    print(f\"\\n{'=' * 80}\")\n    print(f\"REPORT {i}\")\n    print(f\"StudyInstanceUID: {row['StudyInstanceUID']}\")\n    print(f\"{'=' * 80}\")\n    print(row[\"Report\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:51.292723Z","iopub.execute_input":"2026-08-19T04:11:51.293207Z","iopub.status.idle":"2026-08-19T04:11:51.31806Z","shell.execute_reply.started":"2026-08-19T04:11:51.293173Z","shell.execute_reply":"2026-08-19T04:11:51.31722Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_columns = [\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\nlabeled = train[train[label_columns].notna().all(axis=1)].copy()\n\nprint(\"Number of labeled studies:\", len(labeled))\n\nfor i, (_, row) in enumerate(labeled.iterrows(), 1):\n\n    print(\"\\n\" + \"=\" * 100)\n    print(f\"LABELED STUDY {i}\")\n    print(\"=\" * 100)\n\n    print(\"StudyInstanceUID:\", row[\"StudyInstanceUID\"])\n\n    print(\"\\nLABELS:\")\n    for label in label_columns:\n        print(f\"{label:20s}: {int(row[label])}\")\n\n    print(\"\\nREPORT:\")\n    print(row[\"Report\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:51.318977Z","iopub.execute_input":"2026-08-19T04:11:51.320367Z","iopub.status.idle":"2026-08-19T04:11:51.354229Z","shell.execute_reply.started":"2026-08-19T04:11:51.320341Z","shell.execute_reply":"2026-08-19T04:11:51.353001Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_columns = [\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\nlabeled = train[train[label_columns].notna().all(axis=1)].copy()\n\nprint(\"Number of labeled studies:\", len(labeled))\n\nfor i, (_, row) in enumerate(labeled.iterrows(), 1):\n\n    print(\"\\n\" + \"=\" * 100)\n    print(f\"LABELED STUDY {i}\")\n    print(\"=\" * 100)\n\n    print(\"StudyInstanceUID:\", row[\"StudyInstanceUID\"])\n\n    print(\"\\nLABELS:\")\n    for label in label_columns:\n        print(f\"{label:20s}: {int(row[label])}\")\n\n    print(\"\\nREPORT:\")\n    print(row[\"Report\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:51.355522Z","iopub.execute_input":"2026-08-19T04:11:51.355941Z","iopub.status.idle":"2026-08-19T04:11:51.398405Z","shell.execute_reply.started":"2026-08-19T04:11:51.355909Z","shell.execute_reply":"2026-08-19T04:11:51.396786Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport re\n\nINPUT_DIR = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n\ntrain = pd.read_csv(f\"{INPUT_DIR}/train.csv\")\n\nlabel_columns = [\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\nlabeled = train[train[label_columns].notna().all(axis=1)].copy()\n\nprint(\"Total studies:\", len(train))\nprint(\"Labeled studies:\", len(labeled))\nprint(\"Unlabeled studies:\", len(train) - len(labeled))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:51.404185Z","iopub.execute_input":"2026-08-19T04:11:51.405293Z","iopub.status.idle":"2026-08-19T04:11:51.525047Z","shell.execute_reply.started":"2026-08-19T04:11:51.405213Z","shell.execute_reply":"2026-08-19T04:11:51.523783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Basic report statistics\n\ntrain[\"Report\"] = train[\"Report\"].fillna(\"\").astype(str)\n\ntrain[\"report_chars\"] = train[\"Report\"].str.len()\ntrain[\"report_words\"] = train[\"Report\"].str.split().str.len()\n\nprint(\"REPORT STATISTICS\")\nprint(\"=\" * 60)\n\nprint(train[[\"report_chars\", \"report_words\"]].describe())\n\nprint(\"\\nEmpty reports:\", (train[\"Report\"].str.strip() == \"\").sum())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:51.526202Z","iopub.execute_input":"2026-08-19T04:11:51.526578Z","iopub.status.idle":"2026-08-19T04:11:51.647657Z","shell.execute_reply.started":"2026-08-19T04:11:51.526538Z","shell.execute_reply":"2026-08-19T04:11:51.646805Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport re\n\nINPUT_DIR = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n\ntrain = pd.read_csv(f\"{INPUT_DIR}/train.csv\")\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\nlabeled = train[train[LABELS].notna().all(axis=1)].copy()\n\nprint(\"Labeled studies:\", len(labeled))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:51.648608Z","iopub.execute_input":"2026-08-19T04:11:51.648873Z","iopub.status.idle":"2026-08-19T04:11:51.765857Z","shell.execute_reply.started":"2026-08-19T04:11:51.648844Z","shell.execute_reply":"2026-08-19T04:11:51.764957Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initial English terminology only.\n# This is a BASELINE, not our final extractor.\n\nTERMS = {\n    \"ACL\": [\n        \"acl\",\n        \"anterior cruciate ligament\"\n    ],\n\n    \"MCL\": [\n        \"mcl\",\n        \"medial collateral ligament\"\n    ],\n\n    \"Medial Meniscus\": [\n        \"medial meniscus\"\n    ],\n\n    \"Lateral Meniscus\": [\n        \"lateral meniscus\"\n    ],\n\n    \"Medial OA\": [\n        \"medial compartment osteoarthritis\",\n        \"medial compartment oa\"\n    ],\n\n    \"Lateral OA\": [\n        \"lateral compartment osteoarthritis\",\n        \"lateral compartment oa\"\n    ],\n\n    \"PF OA\": [\n        \"patellofemoral osteoarthritis\",\n        \"patellofemoral compartment osteoarthritis\"\n    ],\n\n    \"Effusion\": [\n        \"joint effusion\",\n        \"knee effusion\"\n    ],\n\n    \"Synovitis\": [\n        \"synovitis\"\n    ],\n\n    \"Baker's\": [\n        \"baker's cyst\",\n        \"baker cyst\",\n        \"popliteal cyst\"\n    ],\n\n    \"Contusion\": [\n        \"bone contusion\",\n        \"bone bruise\"\n    ],\n\n    \"Fracture\": [\n        \"fracture\"\n    ]\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:51.767027Z","iopub.execute_input":"2026-08-19T04:11:51.767381Z","iopub.status.idle":"2026-08-19T04:11:51.772631Z","shell.execute_reply.started":"2026-08-19T04:11:51.767358Z","shell.execute_reply":"2026-08-19T04:11:51.771723Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def contains_term(text, terms):\n    text = str(text).lower()\n\n    for term in terms:\n        if term.lower() in text:\n            return 1\n\n    return 0\n\n\nfor label, terms in TERMS.items():\n    labeled[f\"weak_{label}\"] = labeled[\"Report\"].apply(\n        lambda x: contains_term(x, terms)\n    )\n\nprint(\"Weak-label term detection completed.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:51.773602Z","iopub.execute_input":"2026-08-19T04:11:51.773909Z","iopub.status.idle":"2026-08-19T04:11:51.796619Z","shell.execute_reply.started":"2026-08-19T04:11:51.773875Z","shell.execute_reply":"2026-08-19T04:11:51.795959Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score\n\nresults = []\n\nfor label in LABELS:\n\n    y_true = labeled[label].astype(int)\n\n    y_pred = labeled[f\"weak_{label}\"].astype(int)\n\n    results.append({\n        \"Label\": label,\n        \"Accuracy\": accuracy_score(y_true, y_pred),\n        \"Precision\": precision_score(y_true, y_pred, zero_division=0),\n        \"Recall\": recall_score(y_true, y_pred, zero_division=0),\n        \"F1\": f1_score(y_true, y_pred, zero_division=0)\n    })\n\nresults_df = pd.DataFrame(results)\n\ndisplay(results_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:51.797789Z","iopub.execute_input":"2026-08-19T04:11:51.798131Z","iopub.status.idle":"2026-08-19T04:11:53.394481Z","shell.execute_reply.started":"2026-08-19T04:11:51.798102Z","shell.execute_reply":"2026-08-19T04:11:53.393492Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\n\n# ============================================================\n# RSNA KNEE ABNORMALITY DETECTION\n# PHASE 1 — DATASET INTELLIGENCE\n# ============================================================\n\nINPUT_DIR = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n\nTRAIN_CSV = f\"{INPUT_DIR}/train.csv\"\nTRAIN_SERIES_CSV = f\"{INPUT_DIR}/train_series.csv\"\n\ntrain = pd.read_csv(TRAIN_CSV)\ntrain_series = pd.read_csv(TRAIN_SERIES_CSV)\n\nprint(\"=\" * 80)\nprint(\"DATASET LOADED\")\nprint(\"=\" * 80)\n\nprint(f\"Training studies : {len(train):,}\")\nprint(f\"Training series  : {len(train_series):,}\")\n\nprint(\"\\nColumns in train_series:\")\nprint(train_series.columns.tolist())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:53.395659Z","iopub.execute_input":"2026-08-19T04:11:53.396133Z","iopub.status.idle":"2026-08-19T04:11:53.54452Z","shell.execute_reply.started":"2026-08-19T04:11:53.396108Z","shell.execute_reply":"2026-08-19T04:11:53.543831Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# SERIES PER STUDY\n# ============================================================\n\nseries_per_study = (\n    train_series\n    .groupby(\"StudyInstanceUID\")\n    .size()\n    .rename(\"Number_of_Series\")\n)\n\nprint(\"=\" * 80)\nprint(\"SERIES PER STUDY\")\nprint(\"=\" * 80)\n\nprint(series_per_study.describe())\n\nprint(\"\\nMinimum series :\", series_per_study.min())\nprint(\"Maximum series :\", series_per_study.max())\nprint(\"Mean series    :\", round(series_per_study.mean(), 2))\nprint(\"Median series  :\", series_per_study.median())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:53.545545Z","iopub.execute_input":"2026-08-19T04:11:53.545935Z","iopub.status.idle":"2026-08-19T04:11:53.564002Z","shell.execute_reply.started":"2026-08-19T04:11:53.545903Z","shell.execute_reply":"2026-08-19T04:11:53.563394Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# ANATOMICAL PLANE DISTRIBUTION\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"ANATOMICAL PLANE DISTRIBUTION\")\nprint(\"=\" * 80)\n\nplane_counts = train_series[\"Anatomical_Plane\"].value_counts(dropna=False)\n\nprint(plane_counts)\n\nprint(\"\\nPercentage:\")\nprint(\n    (train_series[\"Anatomical_Plane\"]\n     .value_counts(normalize=True, dropna=False) * 100)\n     .round(2)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:53.565611Z","iopub.execute_input":"2026-08-19T04:11:53.566448Z","iopub.status.idle":"2026-08-19T04:11:53.576052Z","shell.execute_reply.started":"2026-08-19T04:11:53.566413Z","shell.execute_reply":"2026-08-19T04:11:53.575224Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# FLUID-SENSITIVE DISTRIBUTION\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"FLUID-SENSITIVE DISTRIBUTION\")\nprint(\"=\" * 80)\n\nfluid_counts = train_series[\"Fluid_Sensitive\"].value_counts(dropna=False)\n\nprint(fluid_counts)\n\nprint(\"\\nPercentage:\")\nprint(\n    (train_series[\"Fluid_Sensitive\"]\n     .value_counts(normalize=True, dropna=False) * 100)\n     .round(2)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:53.577025Z","iopub.execute_input":"2026-08-19T04:11:53.57738Z","iopub.status.idle":"2026-08-19T04:11:53.599017Z","shell.execute_reply.started":"2026-08-19T04:11:53.577357Z","shell.execute_reply":"2026-08-19T04:11:53.598276Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# FAT SUPPRESSION DISTRIBUTION\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"FAT-SUPPRESSION DISTRIBUTION\")\nprint(\"=\" * 80)\n\nfat_counts = train_series[\"Fat_Suppression\"].value_counts(dropna=False)\n\nprint(fat_counts)\n\nprint(\"\\nPercentage:\")\nprint(\n    (train_series[\"Fat_Suppression\"]\n     .value_counts(normalize=True, dropna=False) * 100)\n     .round(2)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:53.600114Z","iopub.execute_input":"2026-08-19T04:11:53.601154Z","iopub.status.idle":"2026-08-19T04:11:53.615514Z","shell.execute_reply.started":"2026-08-19T04:11:53.601122Z","shell.execute_reply":"2026-08-19T04:11:53.614187Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# COMBINED SERIES METADATA\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"SERIES METADATA COMBINATIONS\")\nprint(\"=\" * 80)\n\ncombination_counts = (\n    train_series\n    .groupby(\n        [\"Anatomical_Plane\", \"Fluid_Sensitive\", \"Fat_Suppression\"],\n        dropna=False\n    )\n    .size()\n    .reset_index(name=\"Series_Count\")\n    .sort_values(\"Series_Count\", ascending=False)\n)\n\ndisplay(combination_counts)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:53.616641Z","iopub.execute_input":"2026-08-19T04:11:53.617128Z","iopub.status.idle":"2026-08-19T04:11:53.645859Z","shell.execute_reply.started":"2026-08-19T04:11:53.617093Z","shell.execute_reply":"2026-08-19T04:11:53.64532Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# SERIES COUNT DISTRIBUTION\n# ============================================================\n\ndistribution = (\n    series_per_study\n    .value_counts()\n    .sort_index()\n    .rename_axis(\"Number_of_Series\")\n    .reset_index(name=\"Number_of_Studies\")\n)\n\ndisplay(distribution)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:53.646744Z","iopub.execute_input":"2026-08-19T04:11:53.646951Z","iopub.status.idle":"2026-08-19T04:11:53.65865Z","shell.execute_reply.started":"2026-08-19T04:11:53.646931Z","shell.execute_reply":"2026-08-19T04:11:53.657883Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# SELECT ONE STUDY\n# ============================================================\n\nstudy_id = train_series[\"StudyInstanceUID\"].iloc[0]\n\nstudy_series = train_series[\n    train_series[\"StudyInstanceUID\"] == study_id\n].copy()\n\nprint(\"=\" * 80)\nprint(\"SELECTED STUDY\")\nprint(\"=\" * 80)\n\nprint(\"StudyInstanceUID:\")\nprint(study_id)\n\nprint(\"\\nNumber of series:\", len(study_series))\n\ndisplay(study_series)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:53.659657Z","iopub.execute_input":"2026-08-19T04:11:53.659946Z","iopub.status.idle":"2026-08-19T04:11:53.685912Z","shell.execute_reply.started":"2026-08-19T04:11:53.659915Z","shell.execute_reply":"2026-08-19T04:11:53.685122Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# INSPECT ONE STUDY'S DICOM FOLDERS\n# ============================================================\n\nTRAIN_DICOM_DIR = f\"{INPUT_DIR}/train_series\"\n\nstudy_path = os.path.join(TRAIN_DICOM_DIR, study_id)\n\nprint(\"Study path:\")\nprint(study_path)\n\nprint(\"\\nExists:\", os.path.exists(study_path))\n\nif os.path.exists(study_path):\n    print(\"\\nSeries folders:\")\n    for item in os.listdir(study_path):\n        print(item)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:53.68699Z","iopub.execute_input":"2026-08-19T04:11:53.687313Z","iopub.status.idle":"2026-08-19T04:11:53.704675Z","shell.execute_reply.started":"2026-08-19T04:11:53.68728Z","shell.execute_reply":"2026-08-19T04:11:53.703868Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# COUNT DICOM SLICES FOR ONE STUDY\n# ============================================================\n\nseries_slice_counts = []\n\nfor series_id in os.listdir(study_path):\n\n    series_path = os.path.join(study_path, series_id)\n\n    if not os.path.isdir(series_path):\n        continue\n\n    dicom_files = [\n        f for f in os.listdir(series_path)\n        if f.lower().endswith(\".dcm\")\n    ]\n\n    series_slice_counts.append({\n        \"StudyInstanceUID\": study_id,\n        \"SeriesInstanceUID\": series_id,\n        \"Number_of_Slices\": len(dicom_files)\n    })\n\nslice_df = pd.DataFrame(series_slice_counts)\n\ndisplay(slice_df)\n\nprint(\"\\nTotal DICOM slices in this study:\",\n      slice_df[\"Number_of_Slices\"].sum())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:53.709494Z","iopub.execute_input":"2026-08-19T04:11:53.709896Z","iopub.status.idle":"2026-08-19T04:11:53.741418Z","shell.execute_reply.started":"2026-08-19T04:11:53.709864Z","shell.execute_reply":"2026-08-19T04:11:53.74065Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 11 — DICOM SLICE COUNT DISTRIBUTION\n# ============================================================\n\nimport os\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nTRAIN_DICOM_DIR = f\"{INPUT_DIR}/train_series\"\n\nslice_records = []\n\nfor study_id, series_id in tqdm(\n    train_series[[\"StudyInstanceUID\", \"SeriesInstanceUID\"]].itertuples(index=False),\n    total=len(train_series),\n    desc=\"Counting DICOM slices\"\n):\n    \n    series_path = os.path.join(\n        TRAIN_DICOM_DIR,\n        study_id,\n        series_id\n    )\n\n    if not os.path.exists(series_path):\n        continue\n\n    count = 0\n\n    for filename in os.listdir(series_path):\n        if filename.lower().endswith(\".dcm\"):\n            count += 1\n\n    slice_records.append({\n        \"StudyInstanceUID\": study_id,\n        \"SeriesInstanceUID\": series_id,\n        \"Number_of_Slices\": count\n    })\n\nslice_counts = pd.DataFrame(slice_records)\n\nprint(\"=\" * 80)\nprint(\"DICOM SLICE COUNT ANALYSIS\")\nprint(\"=\" * 80)\n\nprint(\"Series analyzed:\", len(slice_counts))\n\nprint(\"\\nStatistics:\")\nprint(slice_counts[\"Number_of_Slices\"].describe())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:11:53.742141Z","iopub.execute_input":"2026-08-19T04:11:53.742503Z","iopub.status.idle":"2026-08-19T04:14:09.84749Z","shell.execute_reply.started":"2026-08-19T04:11:53.74248Z","shell.execute_reply":"2026-08-19T04:14:09.846614Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport os","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:14:09.848608Z","iopub.execute_input":"2026-08-19T04:14:09.849203Z","iopub.status.idle":"2026-08-19T04:14:10.40386Z","shell.execute_reply.started":"2026-08-19T04:14:09.849167Z","shell.execute_reply":"2026-08-19T04:14:10.403312Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 12 — FIND DICOM FILES IN THE SELECTED STUDY\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"DICOM IMAGE INSPECTION\")\nprint(\"=\" * 80)\n\nall_dicom_files = []\n\nfor series_id in os.listdir(study_path):\n\n    series_path = os.path.join(study_path, series_id)\n\n    if not os.path.isdir(series_path):\n        continue\n\n    files = [\n        os.path.join(series_path, f)\n        for f in os.listdir(series_path)\n        if f.lower().endswith(\".dcm\")\n    ]\n\n    for f in files:\n        all_dicom_files.append({\n            \"series_id\": series_id,\n            \"file\": f\n        })\n\nprint(\"Total DICOM files in selected study:\", len(all_dicom_files))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:14:10.404746Z","iopub.execute_input":"2026-08-19T04:14:10.405073Z","iopub.status.idle":"2026-08-19T04:14:10.417337Z","shell.execute_reply.started":"2026-08-19T04:14:10.405049Z","shell.execute_reply":"2026-08-19T04:14:10.416677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# READ ONE DICOM\n# ============================================================\n\nsample_file = all_dicom_files[len(all_dicom_files)//2][\"file\"]\n\nds = pydicom.dcmread(sample_file)\n\nprint(\"=\" * 80)\nprint(\"DICOM METADATA\")\nprint(\"=\" * 80)\n\nprint(\"File:\")\nprint(sample_file)\n\nprint(\"\\nRows:\", getattr(ds, \"Rows\", \"N/A\"))\nprint(\"Columns:\", getattr(ds, \"Columns\", \"N/A\"))\n\nprint(\"Pixel Spacing:\", getattr(ds, \"PixelSpacing\", \"N/A\"))\nprint(\"Slice Thickness:\", getattr(ds, \"SliceThickness\", \"N/A\"))\nprint(\"Instance Number:\", getattr(ds, \"InstanceNumber\", \"N/A\"))\nprint(\"Image Position:\", getattr(ds, \"ImagePositionPatient\", \"N/A\"))\nprint(\"Image Orientation:\", getattr(ds, \"ImageOrientationPatient\", \"N/A\"))\n\nprint(\"\\nHas pixel data:\", hasattr(ds, \"PixelData\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:14:10.418199Z","iopub.execute_input":"2026-08-19T04:14:10.418906Z","iopub.status.idle":"2026-08-19T04:14:10.447982Z","shell.execute_reply.started":"2026-08-19T04:14:10.418873Z","shell.execute_reply":"2026-08-19T04:14:10.447176Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# DISPLAY ONE REAL MRI SLICE\n# ============================================================\n\nimage = ds.pixel_array\n\nprint(\"Image shape:\", image.shape)\nprint(\"Data type:\", image.dtype)\nprint(\"Minimum intensity:\", image.min())\nprint(\"Maximum intensity:\", image.max())\n\nplt.figure(figsize=(7, 7))\nplt.imshow(image, cmap=\"gray\")\nplt.axis(\"off\")\nplt.title(\"RSNA Knee MRI — DICOM Slice\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:14:10.448827Z","iopub.execute_input":"2026-08-19T04:14:10.449106Z","iopub.status.idle":"2026-08-19T04:14:10.71506Z","shell.execute_reply.started":"2026-08-19T04:14:10.449084Z","shell.execute_reply":"2026-08-19T04:14:10.714201Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 15 — VISUALIZE ALL SERIES FROM ONE STUDY\n# ============================================================\n\nimport pydicom\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport os\n\ndef get_sorted_dicom_files(series_path):\n    \"\"\"\n    Read DICOM files and sort them by InstanceNumber.\n    \"\"\"\n    files = []\n\n    for filename in os.listdir(series_path):\n        if filename.lower().endswith(\".dcm\"):\n            filepath = os.path.join(series_path, filename)\n\n            try:\n                ds_tmp = pydicom.dcmread(\n                    filepath,\n                    stop_before_pixels=True\n                )\n\n                instance_number = getattr(\n                    ds_tmp,\n                    \"InstanceNumber\",\n                    0\n                )\n\n                files.append(\n                    (instance_number, filepath)\n                )\n\n            except Exception:\n                pass\n\n    files.sort(key=lambda x: x[0])\n\n    return [x[1] for x in files]\n\n\n# ------------------------------------------------------------\n# Get series information\n# ------------------------------------------------------------\n\nseries_info = train_series[\n    train_series[\"StudyInstanceUID\"] == study_id\n].copy()\n\nprint(\"Study:\", study_id)\nprint(\"Number of series:\", len(series_info))\n\ndisplay(series_info)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:14:10.716109Z","iopub.execute_input":"2026-08-19T04:14:10.716659Z","iopub.status.idle":"2026-08-19T04:14:10.730419Z","shell.execute_reply.started":"2026-08-19T04:14:10.716632Z","shell.execute_reply":"2026-08-19T04:14:10.729733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# DISPLAY REPRESENTATIVE SLICES FROM EVERY SERIES\n# ============================================================\n\nfor _, row in series_info.iterrows():\n\n    series_id = row[\"SeriesInstanceUID\"]\n\n    series_path = os.path.join(\n        TRAIN_DICOM_DIR,\n        study_id,\n        series_id\n    )\n\n    files = get_sorted_dicom_files(series_path)\n\n    if len(files) == 0:\n        continue\n\n    # Select 5 representative positions\n    indices = np.linspace(\n        0,\n        len(files) - 1,\n        5\n    ).astype(int)\n\n    fig, axes = plt.subplots(\n        1, 5,\n        figsize=(20, 4)\n    )\n\n    for ax, idx in zip(axes, indices):\n\n        ds_img = pydicom.dcmread(files[idx])\n\n        image = ds_img.pixel_array\n\n        ax.imshow(image, cmap=\"gray\")\n        ax.set_title(\n            f\"Slice {idx + 1}/{len(files)}\"\n        )\n        ax.axis(\"off\")\n\n    fig.suptitle(\n        f\"{row['Anatomical_Plane']} | \"\n        f\"Fluid={row['Fluid_Sensitive']} | \"\n        f\"FatSupp={row['Fat_Suppression']} | \"\n        f\"{len(files)} slices\",\n        fontsize=14\n    )\n\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:14:10.731389Z","iopub.execute_input":"2026-08-19T04:14:10.73173Z","iopub.status.idle":"2026-08-19T04:14:15.26668Z","shell.execute_reply.started":"2026-08-19T04:14:10.731707Z","shell.execute_reply":"2026-08-19T04:14:15.265967Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 16 — DICOM GEOMETRY ANALYSIS\n# ONE REPRESENTATIVE DICOM PER SERIES\n# ============================================================\n\nimport os\nimport pydicom\nimport pandas as pd\nimport numpy as np\nfrom tqdm.auto import tqdm\n\nrecords = []\n\nfor row in tqdm(\n    train_series.itertuples(index=False),\n    total=len(train_series),\n    desc=\"Inspecting DICOM metadata\"\n):\n\n    study_id = row.StudyInstanceUID\n    series_id = row.SeriesInstanceUID\n\n    series_path = os.path.join(\n        TRAIN_DICOM_DIR,\n        study_id,\n        series_id\n    )\n\n    if not os.path.exists(series_path):\n        continue\n\n    dicom_files = [\n        f for f in os.listdir(series_path)\n        if f.lower().endswith(\".dcm\")\n    ]\n\n    if len(dicom_files) == 0:\n        continue\n\n    # Take one representative DICOM\n    filepath = os.path.join(\n        series_path,\n        dicom_files[len(dicom_files)//2]\n    )\n\n    try:\n\n        ds = pydicom.dcmread(\n            filepath,\n            stop_before_pixels=True\n        )\n\n        pixel_spacing = getattr(\n            ds,\n            \"PixelSpacing\",\n            [np.nan, np.nan]\n        )\n\n        records.append({\n            \"StudyInstanceUID\": study_id,\n            \"SeriesInstanceUID\": series_id,\n\n            \"Plane\": row.Anatomical_Plane,\n            \"Fluid\": row.Fluid_Sensitive,\n            \"FatSupp\": row.Fat_Suppression,\n\n            \"Rows\": getattr(ds, \"Rows\", np.nan),\n            \"Columns\": getattr(ds, \"Columns\", np.nan),\n\n            \"PixelSpacing_Y\": float(pixel_spacing[0]),\n            \"PixelSpacing_X\": float(pixel_spacing[1]),\n\n            \"SliceThickness\": float(\n                getattr(ds, \"SliceThickness\", np.nan)\n            ),\n\n            \"InstanceNumber\": getattr(\n                ds,\n                \"InstanceNumber\",\n                np.nan\n            )\n        })\n\n    except Exception as e:\n        pass\n\n\ndicom_meta = pd.DataFrame(records)\n\nprint(\"=\" * 80)\nprint(\"DICOM GEOMETRY ANALYSIS\")\nprint(\"=\" * 80)\n\nprint(\"Series successfully inspected:\",\n      len(dicom_meta))\n\ndisplay(dicom_meta.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:14:15.267727Z","iopub.execute_input":"2026-08-19T04:14:15.268039Z","iopub.status.idle":"2026-08-19T04:17:42.222937Z","shell.execute_reply.started":"2026-08-19T04:14:15.268002Z","shell.execute_reply":"2026-08-19T04:17:42.222071Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# IMAGE DIMENSIONS\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"IMAGE DIMENSIONS\")\nprint(\"=\" * 80)\n\ndimension_counts = (\n    dicom_meta\n    .groupby([\"Rows\", \"Columns\"])\n    .size()\n    .reset_index(name=\"Series_Count\")\n    .sort_values(\"Series_Count\", ascending=False)\n)\n\ndisplay(dimension_counts)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:19.175473Z","iopub.execute_input":"2026-08-19T04:44:19.175771Z","iopub.status.idle":"2026-08-19T04:44:19.191182Z","shell.execute_reply.started":"2026-08-19T04:44:19.175748Z","shell.execute_reply":"2026-08-19T04:44:19.190362Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# PIXEL SPACING\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"PIXEL SPACING\")\nprint(\"=\" * 80)\n\nprint(\n    dicom_meta[\n        [\"PixelSpacing_Y\", \"PixelSpacing_X\"]\n    ].describe()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:19.192481Z","iopub.execute_input":"2026-08-19T04:44:19.192821Z","iopub.status.idle":"2026-08-19T04:44:19.215807Z","shell.execute_reply.started":"2026-08-19T04:44:19.192796Z","shell.execute_reply":"2026-08-19T04:44:19.214891Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# SLICE THICKNESS\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"SLICE THICKNESS\")\nprint(\"=\" * 80)\n\nprint(\n    dicom_meta[\"SliceThickness\"].describe()\n)\n\nprint(\"\\nMost common slice thickness values:\")\n\ndisplay(\n    dicom_meta[\"SliceThickness\"]\n    .value_counts()\n    .head(20)\n    .rename_axis(\"SliceThickness\")\n    .reset_index(name=\"Series_Count\")\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:19.216646Z","iopub.execute_input":"2026-08-19T04:44:19.216932Z","iopub.status.idle":"2026-08-19T04:44:19.231945Z","shell.execute_reply.started":"2026-08-19T04:44:19.216909Z","shell.execute_reply":"2026-08-19T04:44:19.231142Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# MISSING SERIES METADATA\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"MISSING METADATA\")\nprint(\"=\" * 80)\n\nprint(\n    train_series[\n        [\"Fluid_Sensitive\",\n         \"Fat_Suppression\",\n         \"Anatomical_Plane\"]\n    ].isna().sum()\n)\n\nprint(\"\\nRows with missing metadata:\")\n\ndisplay(\n    train_series[\n        train_series[\n            [\"Fluid_Sensitive\",\n             \"Fat_Suppression\",\n             \"Anatomical_Plane\"]\n        ].isna().any(axis=1)\n    ].head(20)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:19.233496Z","iopub.execute_input":"2026-08-19T04:44:19.233781Z","iopub.status.idle":"2026-08-19T04:44:19.247337Z","shell.execute_reply.started":"2026-08-19T04:44:19.23376Z","shell.execute_reply":"2026-08-19T04:44:19.246432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 21 — DICOM NORMALIZATION TEST\n# ============================================================\n\nimport cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport pydicom\n\n\ndef normalize_mri(image):\n    \"\"\"\n    Robust MRI intensity normalization.\n\n    Uses percentile clipping to reduce the influence\n    of extreme background/outlier intensities.\n    \"\"\"\n\n    image = image.astype(np.float32)\n\n    # Remove extreme low/high intensity outliers\n    low = np.percentile(image, 1)\n    high = np.percentile(image, 99)\n\n    image = np.clip(image, low, high)\n\n    # Scale to 0-1\n    if high > low:\n        image = (image - low) / (high - low)\n    else:\n        image = np.zeros_like(image)\n\n    return image\n\n\ndef resize_mri(image, size=256):\n    \"\"\"\n    Resize MRI slice to square CNN input.\n    \"\"\"\n\n    return cv2.resize(\n        image,\n        (size, size),\n        interpolation=cv2.INTER_AREA\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:19.248237Z","iopub.execute_input":"2026-08-19T04:44:19.248596Z","iopub.status.idle":"2026-08-19T04:44:19.257181Z","shell.execute_reply.started":"2026-08-19T04:44:19.248556Z","shell.execute_reply":"2026-08-19T04:44:19.256635Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 22 — LOAD ACTUAL PIXEL DATA AND TEST PREPROCESSING\n# ============================================================\n\nimport pydicom\nimport cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\n\n# ------------------------------------------------------------\n# Reload the DICOM WITH pixel data\n# ------------------------------------------------------------\n\nds_pixel = pydicom.dcmread(sample_file)\n\nprint(\"=\" * 80)\nprint(\"DICOM PIXEL DATA\")\nprint(\"=\" * 80)\n\nprint(\"File:\")\nprint(sample_file)\n\nprint(\"Rows:\", ds_pixel.Rows)\nprint(\"Columns:\", ds_pixel.Columns)\nprint(\"Has PixelData:\", \"PixelData\" in ds_pixel)\n\n# ------------------------------------------------------------\n# Get actual image\n# ------------------------------------------------------------\n\noriginal = ds_pixel.pixel_array\n\nprint(\"\\nOriginal image:\")\nprint(\"Shape :\", original.shape)\nprint(\"dtype :\", original.dtype)\nprint(\"Min   :\", original.min())\nprint(\"Max   :\", original.max())\n\n\n# ============================================================\n# NORMALIZATION\n# ============================================================\n\ndef normalize_mri(image):\n\n    image = image.astype(np.float32)\n\n    # Robust percentile clipping\n    low = np.percentile(image, 1)\n    high = np.percentile(image, 99)\n\n    image = np.clip(image, low, high)\n\n    # Normalize to 0–1\n    if high > low:\n        image = (image - low) / (high - low)\n    else:\n        image = np.zeros_like(image)\n\n    return image\n\n\n# ============================================================\n# RESIZE\n# ============================================================\n\ndef resize_mri(image, size=256):\n\n    return cv2.resize(\n        image,\n        (size, size),\n        interpolation=cv2.INTER_AREA\n    )\n\n\nnormalized = normalize_mri(original)\n\nresized = resize_mri(\n    normalized,\n    size=256\n)\n\n\n# ============================================================\n# RESULTS\n# ============================================================\n\nprint(\"\\nNormalized image:\")\nprint(\"Shape :\", normalized.shape)\nprint(\"dtype :\", normalized.dtype)\nprint(\"Min   :\", normalized.min())\nprint(\"Max   :\", normalized.max())\n\nprint(\"\\nResized image:\")\nprint(\"Shape :\", resized.shape)\nprint(\"dtype :\", resized.dtype)\nprint(\"Min   :\", resized.min())\nprint(\"Max   :\", resized.max())\n\n\n# ============================================================\n# VISUALIZATION\n# ============================================================\n\nplt.figure(figsize=(15, 5))\n\nplt.subplot(1, 3, 1)\n\nplt.imshow(\n    original,\n    cmap=\"gray\"\n)\n\nplt.title(\n    f\"Original\\n{original.shape[0]} × {original.shape[1]}\"\n)\n\nplt.axis(\"off\")\n\n\nplt.subplot(1, 3, 2)\n\nplt.imshow(\n    normalized,\n    cmap=\"gray\",\n    vmin=0,\n    vmax=1\n)\n\nplt.title(\"Normalized\\n0–1\")\n\nplt.axis(\"off\")\n\n\nplt.subplot(1, 3, 3)\n\nplt.imshow(\n    resized,\n    cmap=\"gray\",\n    vmin=0,\n    vmax=1\n)\n\nplt.title(\"Resized\\n256 × 256\")\n\nplt.axis(\"off\")\n\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:19.258114Z","iopub.execute_input":"2026-08-19T04:44:19.258539Z","iopub.status.idle":"2026-08-19T04:44:19.695966Z","shell.execute_reply.started":"2026-08-19T04:44:19.258494Z","shell.execute_reply":"2026-08-19T04:44:19.69532Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 23 — VERIFY DICOM SLICE ORDER\n# ============================================================\n\ndef inspect_slice_order(study_id, series_id):\n\n    series_path = os.path.join(\n        TRAIN_DICOM_DIR,\n        study_id,\n        series_id\n    )\n\n    files = [\n        os.path.join(series_path, f)\n        for f in os.listdir(series_path)\n        if f.lower().endswith(\".dcm\")\n    ]\n\n    records = []\n\n    for filepath in files:\n\n        try:\n            dcm = pydicom.dcmread(\n                filepath,\n                stop_before_pixels=True\n            )\n\n            instance = getattr(\n                dcm,\n                \"InstanceNumber\",\n                np.nan\n            )\n\n            position = getattr(\n                dcm,\n                \"ImagePositionPatient\",\n                [np.nan, np.nan, np.nan]\n            )\n\n            records.append({\n                \"file\": filepath,\n                \"InstanceNumber\": instance,\n                \"X\": float(position[0]),\n                \"Y\": float(position[1]),\n                \"Z\": float(position[2])\n            })\n\n        except Exception:\n            pass\n\n    df = pd.DataFrame(records)\n\n    return df\n\n\n# Use the first series of our selected study\ntest_series_id = series_info.iloc[0][\"SeriesInstanceUID\"]\n\norder_df = inspect_slice_order(\n    study_id,\n    test_series_id\n)\n\nprint(\"=\" * 80)\nprint(\"SLICE ORDER INSPECTION\")\nprint(\"=\" * 80)\n\nprint(\"Total slices:\", len(order_df))\n\ndisplay(order_df.head(10))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:19.697493Z","iopub.execute_input":"2026-08-19T04:44:19.697698Z","iopub.status.idle":"2026-08-19T04:44:19.767022Z","shell.execute_reply.started":"2026-08-19T04:44:19.697678Z","shell.execute_reply":"2026-08-19T04:44:19.766271Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Compare InstanceNumber ordering with physical Z position\n\nprint(\"=\" * 80)\nprint(\"INSTANCE NUMBER RANGE\")\nprint(\"=\" * 80)\n\nprint(\n    order_df[\"InstanceNumber\"].min(),\n    \"→\",\n    order_df[\"InstanceNumber\"].max()\n)\n\nprint(\"\\nPhysical position ranges:\")\n\nprint(\"X:\", order_df[\"X\"].min(), \"→\", order_df[\"X\"].max())\nprint(\"Y:\", order_df[\"Y\"].min(), \"→\", order_df[\"Y\"].max())\nprint(\"Z:\", order_df[\"Z\"].min(), \"→\", order_df[\"Z\"].max())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:19.768021Z","iopub.execute_input":"2026-08-19T04:44:19.768669Z","iopub.status.idle":"2026-08-19T04:44:19.775616Z","shell.execute_reply.started":"2026-08-19T04:44:19.768646Z","shell.execute_reply":"2026-08-19T04:44:19.775021Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 24 — ROBUST PHYSICAL SLICE SORTING\n# ============================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport pydicom\n\n\ndef get_sorted_dicom_files_physical(series_path):\n    \"\"\"\n    Sort DICOM slices using physical position.\n\n    The slice coordinate is calculated by projecting\n    ImagePositionPatient onto the slice normal obtained\n    from ImageOrientationPatient.\n    \"\"\"\n\n    records = []\n\n    for filename in os.listdir(series_path):\n\n        if not filename.lower().endswith(\".dcm\"):\n            continue\n\n        filepath = os.path.join(series_path, filename)\n\n        try:\n\n            ds = pydicom.dcmread(\n                filepath,\n                stop_before_pixels=True\n            )\n\n            position = np.array(\n                ds.ImagePositionPatient,\n                dtype=np.float64\n            )\n\n            orientation = np.array(\n                ds.ImageOrientationPatient,\n                dtype=np.float64\n            )\n\n            # First 3 values = row direction\n            row_direction = orientation[:3]\n\n            # Last 3 values = column direction\n            column_direction = orientation[3:]\n\n            # Slice normal\n            normal = np.cross(\n                row_direction,\n                column_direction\n            )\n\n            # Normalize\n            normal = normal / np.linalg.norm(normal)\n\n            # Physical position along slice direction\n            slice_coordinate = np.dot(\n                position,\n                normal\n            )\n\n            instance_number = getattr(\n                ds,\n                \"InstanceNumber\",\n                -1\n            )\n\n            records.append({\n                \"file\": filepath,\n                \"InstanceNumber\": instance_number,\n                \"X\": position[0],\n                \"Y\": position[1],\n                \"Z\": position[2],\n                \"SliceCoordinate\": slice_coordinate\n            })\n\n        except Exception as e:\n            print(\n                \"Error:\",\n                filepath,\n                e\n            )\n\n    df = pd.DataFrame(records)\n\n    # Sort using physical slice coordinate\n    df = df.sort_values(\n        \"SliceCoordinate\"\n    ).reset_index(drop=True)\n\n    return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:19.776802Z","iopub.execute_input":"2026-08-19T04:44:19.777078Z","iopub.status.idle":"2026-08-19T04:44:19.790058Z","shell.execute_reply.started":"2026-08-19T04:44:19.777058Z","shell.execute_reply":"2026-08-19T04:44:19.789501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# TEST PHYSICAL SORTING\n# ============================================================\n\nseries_path = os.path.join(\n    TRAIN_DICOM_DIR,\n    study_id,\n    test_series_id\n)\n\nsorted_df = get_sorted_dicom_files_physical(\n    series_path\n)\n\nprint(\"=\" * 80)\nprint(\"PHYSICALLY SORTED SLICES\")\nprint(\"=\" * 80)\n\nprint(\"Number of slices:\", len(sorted_df))\n\ndisplay(\n    sorted_df[\n        [\n            \"InstanceNumber\",\n            \"X\",\n            \"Y\",\n            \"Z\",\n            \"SliceCoordinate\"\n        ]\n    ].head(15)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:19.791155Z","iopub.execute_input":"2026-08-19T04:44:19.791767Z","iopub.status.idle":"2026-08-19T04:44:19.855319Z","shell.execute_reply.started":"2026-08-19T04:44:19.791744Z","shell.execute_reply":"2026-08-19T04:44:19.854686Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\nInstance numbers after physical sorting:\")\n\nprint(\n    sorted_df[\"InstanceNumber\"].tolist()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:19.856014Z","iopub.execute_input":"2026-08-19T04:44:19.856193Z","iopub.status.idle":"2026-08-19T04:44:19.860768Z","shell.execute_reply.started":"2026-08-19T04:44:19.856176Z","shell.execute_reply":"2026-08-19T04:44:19.859949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 27 — VERIFY SLICE SORTING FOR ALL 3 PLANES\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"VERIFYING PHYSICAL SORTING — ALL PLANES\")\nprint(\"=\" * 80)\n\n# ------------------------------------------------------------\n# Find one series from each plane\n# ------------------------------------------------------------\n\nplane_examples = {}\n\nfor plane in [\"Sagittal\", \"Coronal\", \"Axial\"]:\n\n    rows = series_info[\n        series_info[\"Anatomical_Plane\"] == plane\n    ]\n\n    if len(rows) > 0:\n        plane_examples[plane] = rows.iloc[0][\"SeriesInstanceUID\"]\n\n\n# ------------------------------------------------------------\n# Test each plane\n# ------------------------------------------------------------\n\nfor plane, series_id in plane_examples.items():\n\n    series_path = os.path.join(\n        TRAIN_DICOM_DIR,\n        study_id,\n        series_id\n    )\n\n    sorted_plane_df = get_sorted_dicom_files_physical(\n        series_path\n    )\n\n    coordinates = sorted_plane_df[\n        \"SliceCoordinate\"\n    ].values\n\n    differences = np.diff(coordinates)\n\n    print(\"\\n\" + \"-\" * 80)\n    print(f\"PLANE: {plane}\")\n    print(\"-\" * 80)\n\n    print(\"Series:\", series_id)\n    print(\"Number of slices:\", len(sorted_plane_df))\n\n    print(\n        \"Coordinate start:\",\n        coordinates[0]\n    )\n\n    print(\n        \"Coordinate end:\",\n        coordinates[-1]\n    )\n\n    print(\n        \"Minimum coordinate difference:\",\n        differences.min()\n    )\n\n    print(\n        \"Maximum coordinate difference:\",\n        differences.max()\n    )\n\n    print(\n        \"Monotonically increasing:\",\n        np.all(differences >= 0)\n    )\n\n    print(\n        \"Instance numbers:\",\n        sorted_plane_df[\"InstanceNumber\"].tolist()\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:19.861838Z","iopub.execute_input":"2026-08-19T04:44:19.862043Z","iopub.status.idle":"2026-08-19T04:44:20.023825Z","shell.execute_reply.started":"2026-08-19T04:44:19.862021Z","shell.execute_reply":"2026-08-19T04:44:20.02322Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 28 — ADAPTIVE SLICE SAMPLER\n# ============================================================\n\ndef sample_slice_indices(num_slices, target_slices=12):\n    \"\"\"\n    Select approximately evenly distributed slices\n    across the complete MRI series.\n\n    The first and last slices are included.\n    \"\"\"\n\n    if num_slices <= 0:\n        return []\n\n    # If the series already has <= target slices,\n    # use every slice.\n    if num_slices <= target_slices:\n        return np.arange(num_slices)\n\n    # Evenly distributed sampling\n    indices = np.linspace(\n        0,\n        num_slices - 1,\n        target_slices\n    )\n\n    indices = np.round(indices).astype(int)\n\n    # Remove any accidental duplicates\n    indices = np.unique(indices)\n\n    return indices","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:20.025987Z","iopub.execute_input":"2026-08-19T04:44:20.026589Z","iopub.status.idle":"2026-08-19T04:44:20.031093Z","shell.execute_reply.started":"2026-08-19T04:44:20.026566Z","shell.execute_reply":"2026-08-19T04:44:20.030461Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# TEST ADAPTIVE SAMPLING\n# ============================================================\n\ntest_sizes = [\n    11,\n    20,\n    24,\n    30,\n    50,\n    100,\n    200,\n    320\n]\n\nprint(\"=\" * 80)\nprint(\"ADAPTIVE SLICE SAMPLING\")\nprint(\"=\" * 80)\n\nfor n in test_sizes:\n\n    indices = sample_slice_indices(\n        n,\n        target_slices=12\n    )\n\n    print(\n        f\"\\nOriginal slices: {n:3d}\"\n    )\n\n    print(\n        f\"Selected slices: {len(indices):2d}\"\n    )\n\n    print(\n        \"Indices:\",\n        indices.tolist()\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:20.032776Z","iopub.execute_input":"2026-08-19T04:44:20.033146Z","iopub.status.idle":"2026-08-19T04:44:20.048411Z","shell.execute_reply.started":"2026-08-19T04:44:20.033124Z","shell.execute_reply":"2026-08-19T04:44:20.04766Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 30 — VISUALIZE SELECTED SLICES\n# ============================================================\n\n# Use the sagittal example\nplane = \"Sagittal\"\n\nseries_id = plane_examples[plane]\n\nseries_path = os.path.join(\n    TRAIN_DICOM_DIR,\n    study_id,\n    series_id\n)\n\nsorted_df = get_sorted_dicom_files_physical(\n    series_path\n)\n\nindices = sample_slice_indices(\n    len(sorted_df),\n    target_slices=12\n)\n\nfig, axes = plt.subplots(\n    3,\n    4,\n    figsize=(12, 10)\n)\n\nfor ax, idx in zip(\n    axes.ravel(),\n    indices\n):\n\n    filepath = sorted_df.iloc[idx][\"file\"]\n\n    dcm = pydicom.dcmread(filepath)\n\n    image = dcm.pixel_array\n\n    image = normalize_mri(image)\n\n    ax.imshow(\n        image,\n        cmap=\"gray\",\n        vmin=0,\n        vmax=1\n    )\n\n    ax.set_title(\n        f\"Slice {idx + 1}/{len(sorted_df)}\"\n    )\n\n    ax.axis(\"off\")\n\n\nplt.suptitle(\n    f\"{plane} — 12 Representative Slices\",\n    fontsize=16\n)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:20.049553Z","iopub.execute_input":"2026-08-19T04:44:20.050136Z","iopub.status.idle":"2026-08-19T04:44:21.567814Z","shell.execute_reply.started":"2026-08-19T04:44:20.050112Z","shell.execute_reply":"2026-08-19T04:44:21.566479Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 31 — COMPARE SLICE SAMPLING STRATEGIES\n# ============================================================\n\ndef sample_uniform(num_slices, target_slices=12):\n\n    if num_slices <= target_slices:\n        return np.arange(num_slices)\n\n    return np.round(\n        np.linspace(\n            0,\n            num_slices - 1,\n            target_slices\n        )\n    ).astype(int)\n\n\ndef sample_central(\n    num_slices,\n    target_slices=12,\n    margin=0.10\n):\n    \"\"\"\n    Sample uniformly from the central portion\n    of the series.\n    \"\"\"\n\n    if num_slices <= target_slices:\n        return np.arange(num_slices)\n\n    start = int(\n        round(num_slices * margin)\n    )\n\n    end = int(\n        round(num_slices * (1 - margin))\n    ) - 1\n\n    indices = np.round(\n        np.linspace(\n            start,\n            end,\n            target_slices\n        )\n    ).astype(int)\n\n    return np.unique(indices)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:21.56909Z","iopub.execute_input":"2026-08-19T04:44:21.569702Z","iopub.status.idle":"2026-08-19T04:44:21.575992Z","shell.execute_reply.started":"2026-08-19T04:44:21.569665Z","shell.execute_reply":"2026-08-19T04:44:21.575167Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# PRINT COMPARISON\n# ============================================================\n\nn = 24\n\nuniform = sample_uniform(n, 12)\n\ncentral_80 = sample_central(\n    n,\n    12,\n    margin=0.10\n)\n\ncentral_70 = sample_central(\n    n,\n    12,\n    margin=0.15\n)\n\nprint(\"=\" * 80)\nprint(\"SAMPLING COMPARISON — 24 SLICES\")\nprint(\"=\" * 80)\n\nprint(\"\\nUniform:\")\nprint(uniform.tolist())\n\nprint(\"\\nCentral 80%:\")\nprint(central_80.tolist())\n\nprint(\"\\nCentral 70%:\")\nprint(central_70.tolist())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:21.577725Z","iopub.execute_input":"2026-08-19T04:44:21.578313Z","iopub.status.idle":"2026-08-19T04:44:21.591293Z","shell.execute_reply.started":"2026-08-19T04:44:21.578281Z","shell.execute_reply":"2026-08-19T04:44:21.590429Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# VISUAL COMPARISON\n# ============================================================\n\nstrategies = {\n    \"Uniform\": uniform,\n    \"Central 80%\": central_80,\n    \"Central 70%\": central_70\n}\n\nfig, axes = plt.subplots(\n    3,\n    12,\n    figsize=(24, 7)\n)\n\nfor row_idx, (name, indices) in enumerate(\n    strategies.items()\n):\n\n    for col_idx in range(12):\n\n        ax = axes[row_idx, col_idx]\n\n        idx = indices[col_idx]\n\n        filepath = sorted_df.iloc[idx][\"file\"]\n\n        dcm = pydicom.dcmread(filepath)\n\n        image = normalize_mri(\n            dcm.pixel_array\n        )\n\n        ax.imshow(\n            image,\n            cmap=\"gray\",\n            vmin=0,\n            vmax=1\n        )\n\n        ax.set_title(\n            f\"{idx + 1}/{len(sorted_df)}\",\n            fontsize=9\n        )\n\n        ax.axis(\"off\")\n\n    axes[row_idx, 0].set_ylabel(\n        name,\n        fontsize=12\n    )\n\nplt.suptitle(\n    \"MRI Slice Sampling Comparison\",\n    fontsize=16\n)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:21.592337Z","iopub.execute_input":"2026-08-19T04:44:21.592642Z","iopub.status.idle":"2026-08-19T04:44:25.179759Z","shell.execute_reply.started":"2026-08-19T04:44:21.592609Z","shell.execute_reply":"2026-08-19T04:44:25.177153Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 33 — FINAL FIRST-VERSION SLICE SAMPLER\n# ============================================================\n\ndef sample_mri_slices(\n    num_slices,\n    target_slices=12,\n    central_fraction=0.80\n):\n    \"\"\"\n    Select representative MRI slices.\n\n    First version:\n    - Uses the central 80% of the series\n    - Samples approximately uniformly\n    - Produces up to 12 slices\n    \"\"\"\n\n    if num_slices <= 0:\n        return np.array([], dtype=int)\n\n    # If short series, keep all slices\n    if num_slices <= target_slices:\n        return np.arange(num_slices)\n\n    # Calculate central region\n    margin = (1.0 - central_fraction) / 2.0\n\n    start = int(\n        round(num_slices * margin)\n    )\n\n    end = int(\n        round(num_slices * (1.0 - margin))\n    ) - 1\n\n    # Safety\n    start = max(0, start)\n    end = min(num_slices - 1, end)\n\n    # Uniformly sample\n    indices = np.linspace(\n        start,\n        end,\n        target_slices\n    )\n\n    indices = np.round(indices).astype(int)\n\n    return np.unique(indices)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:25.180702Z","iopub.execute_input":"2026-08-19T04:44:25.180909Z","iopub.status.idle":"2026-08-19T04:44:25.187324Z","shell.execute_reply.started":"2026-08-19T04:44:25.180889Z","shell.execute_reply":"2026-08-19T04:44:25.186504Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# TEST FINAL SAMPLER\n# ============================================================\n\nfor n in [11, 20, 24, 30, 50, 100, 200, 320]:\n\n    indices = sample_mri_slices(\n        n,\n        target_slices=12,\n        central_fraction=0.80\n    )\n\n    print(\n        f\"{n:3d} slices → \"\n        f\"{len(indices):2d} selected → \"\n        f\"{indices.tolist()}\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:25.18837Z","iopub.execute_input":"2026-08-19T04:44:25.18856Z","iopub.status.idle":"2026-08-19T04:44:25.206309Z","shell.execute_reply.started":"2026-08-19T04:44:25.188541Z","shell.execute_reply":"2026-08-19T04:44:25.205711Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 34 — CENTRAL 80% SAMPLING FOR ALL 3 PLANES\n# ============================================================\n\nfig, axes = plt.subplots(\n    3,\n    12,\n    figsize=(24, 7)\n)\n\nfor row, (plane, series_id) in enumerate(\n    plane_examples.items()\n):\n\n    series_path = os.path.join(\n        TRAIN_DICOM_DIR,\n        study_id,\n        series_id\n    )\n\n    sorted_df = get_sorted_dicom_files_physical(\n        series_path\n    )\n\n    indices = sample_mri_slices(\n        len(sorted_df),\n        target_slices=12,\n        central_fraction=0.80\n    )\n\n    for col in range(12):\n\n        ax = axes[row, col]\n\n        # Some short series may have fewer than 12 slices\n        if col >= len(indices):\n            ax.axis(\"off\")\n            continue\n\n        idx = indices[col]\n\n        filepath = sorted_df.iloc[idx][\"file\"]\n\n        dcm = pydicom.dcmread(filepath)\n\n        image = normalize_mri(\n            dcm.pixel_array\n        )\n\n        ax.imshow(\n            image,\n            cmap=\"gray\",\n            vmin=0,\n            vmax=1\n        )\n\n        ax.set_title(\n            f\"{idx + 1}/{len(sorted_df)}\",\n            fontsize=8\n        )\n\n        ax.axis(\"off\")\n\n    axes[row, 0].set_ylabel(\n        plane,\n        fontsize=13\n    )\n\nplt.suptitle(\n    \"Central 80% — Representative MRI Slices\",\n    fontsize=16\n)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:25.207264Z","iopub.execute_input":"2026-08-19T04:44:25.207541Z","iopub.status.idle":"2026-08-19T04:44:28.749441Z","shell.execute_reply.started":"2026-08-19T04:44:25.207511Z","shell.execute_reply":"2026-08-19T04:44:28.748478Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 35 — BUILD STUDY / SERIES MANIFEST\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"BUILDING STUDY-SERIES MANIFEST\")\nprint(\"=\" * 80)\n\nseries_manifest = train_series.copy()\n\n# Count DICOM slices per series\nslice_counts = (\n    dicom_meta\n    .groupby(\n        [\"StudyInstanceUID\", \"SeriesInstanceUID\"]\n    )\n    .size()\n    .rename(\"Number_of_Slices\")\n    .reset_index()\n)\n\n# Merge slice count\nseries_manifest = series_manifest.merge(\n    slice_counts,\n    on=[\n        \"StudyInstanceUID\",\n        \"SeriesInstanceUID\"\n    ],\n    how=\"left\"\n)\n\nprint(\"Studies:\",\n      series_manifest[\"StudyInstanceUID\"].nunique())\n\nprint(\"Series:\",\n      len(series_manifest))\n\nprint(\"Missing slice counts:\",\n      series_manifest[\"Number_of_Slices\"].isna().sum())\n\ndisplay(\n    series_manifest.head(10)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:28.750486Z","iopub.execute_input":"2026-08-19T04:44:28.750802Z","iopub.status.idle":"2026-08-19T04:44:28.814671Z","shell.execute_reply.started":"2026-08-19T04:44:28.750778Z","shell.execute_reply":"2026-08-19T04:44:28.813888Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 36 — ADD DICOM PATHS\n# ============================================================\n\nseries_manifest[\"SeriesPath\"] = (\n    TRAIN_DICOM_DIR\n    + \"/\"\n    + series_manifest[\"StudyInstanceUID\"].astype(str)\n    + \"/\"\n    + series_manifest[\"SeriesInstanceUID\"].astype(str)\n)\n\nprint(\"=\" * 80)\nprint(\"SERIES PATH CHECK\")\nprint(\"=\" * 80)\n\ndisplay(\n    series_manifest[\n        [\n            \"StudyInstanceUID\",\n            \"SeriesInstanceUID\",\n            \"Anatomical_Plane\",\n            \"Fluid_Sensitive\",\n            \"Fat_Suppression\",\n            \"Number_of_Slices\",\n            \"SeriesPath\"\n        ]\n    ].head()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:28.81574Z","iopub.execute_input":"2026-08-19T04:44:28.816008Z","iopub.status.idle":"2026-08-19T04:44:28.8402Z","shell.execute_reply.started":"2026-08-19T04:44:28.815986Z","shell.execute_reply":"2026-08-19T04:44:28.839326Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 37 — COMPLETE STUDY STRUCTURE\n# ============================================================\n\nexample_study = (\n    series_manifest[\n        \"StudyInstanceUID\"\n    ].iloc[0]\n)\n\nstudy_manifest = series_manifest[\n    series_manifest[\"StudyInstanceUID\"] == example_study\n].copy()\n\nprint(\"=\" * 80)\nprint(\"COMPLETE STUDY\")\nprint(\"=\" * 80)\n\nprint(\"Study ID:\")\nprint(example_study)\n\nprint(\"\\nNumber of series:\",\n      len(study_manifest))\n\ndisplay(\n    study_manifest[\n        [\n            \"SeriesInstanceUID\",\n            \"Anatomical_Plane\",\n            \"Fluid_Sensitive\",\n            \"Fat_Suppression\",\n            \"Number_of_Slices\"\n        ]\n    ]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:28.841216Z","iopub.execute_input":"2026-08-19T04:44:28.841592Z","iopub.status.idle":"2026-08-19T04:44:28.855648Z","shell.execute_reply.started":"2026-08-19T04:44:28.841569Z","shell.execute_reply":"2026-08-19T04:44:28.85482Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 38 — ESTIMATE SELECTED IMAGE COUNT\n# ============================================================\n\ndef number_of_selected_slices(\n    num_slices,\n    target_slices=12\n):\n\n    return min(\n        num_slices,\n        target_slices\n    )\n\n\nstudy_manifest[\"Selected_Slices\"] = (\n    study_manifest[\"Number_of_Slices\"]\n    .apply(number_of_selected_slices)\n)\n\nprint(\"=\" * 80)\nprint(\"STUDY IMAGE PROCESSING ESTIMATE\")\nprint(\"=\" * 80)\n\ndisplay(\n    study_manifest[\n        [\n            \"Anatomical_Plane\",\n            \"Fluid_Sensitive\",\n            \"Fat_Suppression\",\n            \"Number_of_Slices\",\n            \"Selected_Slices\"\n        ]\n    ]\n)\n\nprint(\n    \"\\nTotal selected images for study:\",\n    study_manifest[\"Selected_Slices\"].sum()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:28.856637Z","iopub.execute_input":"2026-08-19T04:44:28.85701Z","iopub.status.idle":"2026-08-19T04:44:28.879152Z","shell.execute_reply.started":"2026-08-19T04:44:28.856963Z","shell.execute_reply":"2026-08-19T04:44:28.878515Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 35 — CORRECT ACTUAL DICOM SLICE COUNTS\n# ============================================================\n\nimport os\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nprint(\"=\" * 80)\nprint(\"COUNTING ACTUAL DICOM SLICES\")\nprint(\"=\" * 80)\n\nactual_slice_counts = []\n\nfor row in tqdm(\n    train_series.itertuples(index=False),\n    total=len(train_series),\n    desc=\"Counting DICOM slices\"\n):\n\n    study_id_current = row.StudyInstanceUID\n    series_id_current = row.SeriesInstanceUID\n\n    series_path = os.path.join(\n        TRAIN_DICOM_DIR,\n        study_id_current,\n        series_id_current\n    )\n\n    try:\n\n        count = sum(\n            1\n            for filename in os.listdir(series_path)\n            if filename.lower().endswith(\".dcm\")\n        )\n\n    except Exception:\n\n        count = 0\n\n    actual_slice_counts.append({\n        \"StudyInstanceUID\": study_id_current,\n        \"SeriesInstanceUID\": series_id_current,\n        \"Number_of_Slices\": count\n    })\n\n\nactual_slice_counts = pd.DataFrame(\n    actual_slice_counts\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"SLICE COUNT RESULT\")\nprint(\"=\" * 80)\n\nprint(\n    \"Series analyzed:\",\n    len(actual_slice_counts)\n)\n\nprint(\n    \"Missing/zero series:\",\n    (actual_slice_counts[\"Number_of_Slices\"] == 0).sum()\n)\n\nprint(\n    \"\\nSlice statistics:\"\n)\n\nprint(\n    actual_slice_counts[\n        \"Number_of_Slices\"\n    ].describe()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:28.880076Z","iopub.execute_input":"2026-08-19T04:44:28.88047Z","iopub.status.idle":"2026-08-19T04:44:49.232409Z","shell.execute_reply.started":"2026-08-19T04:44:28.880439Z","shell.execute_reply":"2026-08-19T04:44:49.231784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 36 — BUILD CORRECT SERIES MANIFEST\n# ============================================================\n\nseries_manifest = train_series.copy()\n\nseries_manifest = series_manifest.merge(\n    actual_slice_counts,\n    on=[\n        \"StudyInstanceUID\",\n        \"SeriesInstanceUID\"\n    ],\n    how=\"left\"\n)\n\nseries_manifest[\"SeriesPath\"] = (\n    TRAIN_DICOM_DIR\n    + \"/\"\n    + series_manifest[\"StudyInstanceUID\"].astype(str)\n    + \"/\"\n    + series_manifest[\"SeriesInstanceUID\"].astype(str)\n)\n\nprint(\"=\" * 80)\nprint(\"CORRECTED STUDY-SERIES MANIFEST\")\nprint(\"=\" * 80)\n\nprint(\n    \"Studies:\",\n    series_manifest[\"StudyInstanceUID\"].nunique()\n)\n\nprint(\n    \"Series:\",\n    len(series_manifest)\n)\n\nprint(\n    \"Missing slice counts:\",\n    series_manifest[\"Number_of_Slices\"].isna().sum()\n)\n\ndisplay(\n    series_manifest[\n        [\n            \"StudyInstanceUID\",\n            \"SeriesInstanceUID\",\n            \"Fluid_Sensitive\",\n            \"Fat_Suppression\",\n            \"Anatomical_Plane\",\n            \"Number_of_Slices\"\n        ]\n    ].head(10)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:49.23347Z","iopub.execute_input":"2026-08-19T04:44:49.233863Z","iopub.status.idle":"2026-08-19T04:44:49.273385Z","shell.execute_reply.started":"2026-08-19T04:44:49.233828Z","shell.execute_reply":"2026-08-19T04:44:49.272762Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 40 — STUDY-LEVEL MRI LOADER\n# ============================================================\n\ndef load_series_images(\n    series_path,\n    target_slices=12,\n    image_size=256,\n    central_fraction=0.80\n):\n    \"\"\"\n    Load one MRI series.\n\n    Pipeline:\n        DICOM files\n        -> physical sorting\n        -> Central 80% sampling\n        -> pixel loading\n        -> normalization\n        -> resize\n\n    Returns:\n        images: [N, 256, 256]\n        metadata: information about selected slices\n    \"\"\"\n\n    # --------------------------------------------------------\n    # 1. Get physically sorted DICOM files\n    # --------------------------------------------------------\n\n    sorted_df = get_sorted_dicom_files_physical(\n        series_path\n    )\n\n    num_slices = len(sorted_df)\n\n    if num_slices == 0:\n        return np.empty(\n            (0, image_size, image_size),\n            dtype=np.float32\n        ), []\n\n    # --------------------------------------------------------\n    # 2. Select representative slices\n    # --------------------------------------------------------\n\n    selected_indices = sample_mri_slices(\n        num_slices,\n        target_slices=target_slices,\n        central_fraction=central_fraction\n    )\n\n    images = []\n    metadata = []\n\n    # --------------------------------------------------------\n    # 3. Load selected DICOM images\n    # --------------------------------------------------------\n\n    for idx in selected_indices:\n\n        filepath = sorted_df.iloc[idx][\"file\"]\n\n        try:\n\n            ds = pydicom.dcmread(\n                filepath\n            )\n\n            image = ds.pixel_array\n\n            # ----------------------------------------------\n            # Normalize\n            # ----------------------------------------------\n\n            image = normalize_mri(\n                image\n            )\n\n            # ----------------------------------------------\n            # Resize\n            # ----------------------------------------------\n\n            image = cv2.resize(\n                image,\n                (\n                    image_size,\n                    image_size\n                ),\n                interpolation=cv2.INTER_AREA\n            )\n\n            images.append(\n                image.astype(np.float32)\n            )\n\n            metadata.append({\n                \"index\": int(idx),\n                \"InstanceNumber\": int(\n                    sorted_df.iloc[idx][\"InstanceNumber\"]\n                ),\n                \"SliceCoordinate\": float(\n                    sorted_df.iloc[idx][\"SliceCoordinate\"]\n                )\n            })\n\n        except Exception as e:\n\n            print(\n                f\"Failed to read: {filepath}\"\n            )\n\n            print(\n                \"Error:\",\n                e\n            )\n\n    # --------------------------------------------------------\n    # 4. Convert to NumPy\n    # --------------------------------------------------------\n\n    if len(images) == 0:\n\n        images_array = np.empty(\n            (\n                0,\n                image_size,\n                image_size\n            ),\n            dtype=np.float32\n        )\n\n    else:\n\n        images_array = np.stack(\n            images\n        ).astype(np.float32)\n\n    return (\n        images_array,\n        metadata\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:49.274371Z","iopub.execute_input":"2026-08-19T04:44:49.274616Z","iopub.status.idle":"2026-08-19T04:44:49.283087Z","shell.execute_reply.started":"2026-08-19T04:44:49.274594Z","shell.execute_reply":"2026-08-19T04:44:49.282389Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 41 — STUDY LOADER\n# ============================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport pydicom\nfrom PIL import Image\n\n\ndef normalize_mri(image):\n    \"\"\"\n    Normalize MRI pixel values to [0, 1].\n    \"\"\"\n    image = image.astype(np.float32)\n\n    min_val = image.min()\n    max_val = image.max()\n\n    if max_val > min_val:\n        image = (image - min_val) / (max_val - min_val)\n    else:\n        image = np.zeros_like(image, dtype=np.float32)\n\n    return image\n\n\ndef sample_slice_indices(num_slices, target_slices=12):\n    \"\"\"\n    Select evenly distributed slices from a series.\n\n    Uses the complete series while avoiding only-edge\n    sampling when enough slices are available.\n    \"\"\"\n\n    if num_slices <= 0:\n        return np.array([], dtype=int)\n\n    if num_slices <= target_slices:\n        return np.arange(num_slices)\n\n    # Central 80% sampling\n    start = int(round(0.10 * (num_slices - 1)))\n    end = int(round(0.90 * (num_slices - 1)))\n\n    indices = np.linspace(\n        start,\n        end,\n        target_slices\n    )\n\n    indices = np.round(indices).astype(int)\n\n    return np.unique(indices)\n\n\ndef load_dicom_series(\n    series_path,\n    target_slices=12,\n    image_size=(256, 256)\n):\n    \"\"\"\n    Load one DICOM series.\n\n    Steps:\n    1. Read DICOM files\n    2. Sort physically using ImagePositionPatient\n    3. Select representative slices\n    4. Normalize\n    5. Resize to 256 x 256\n    \"\"\"\n\n    dicom_files = []\n\n    for filename in os.listdir(series_path):\n\n        if filename.lower().endswith(\".dcm\"):\n\n            filepath = os.path.join(\n                series_path,\n                filename\n            )\n\n            try:\n                ds = pydicom.dcmread(\n                    filepath,\n                    stop_before_pixels=False\n                )\n\n                if not hasattr(ds, \"PixelData\"):\n                    continue\n\n                # Physical coordinate\n                position = getattr(\n                    ds,\n                    \"ImagePositionPatient\",\n                    None\n                )\n\n                if position is not None:\n                    coordinate = float(position[2])\n                else:\n                    coordinate = float(\n                        getattr(\n                            ds,\n                            \"InstanceNumber\",\n                            0\n                        )\n                    )\n\n                dicom_files.append(\n                    {\n                        \"file\": filepath,\n                        \"coordinate\": coordinate,\n                        \"instance\": int(\n                            getattr(\n                                ds,\n                                \"InstanceNumber\",\n                                0\n                            )\n                        )\n                    }\n                )\n\n            except Exception:\n                continue\n\n    if len(dicom_files) == 0:\n        raise RuntimeError(\n            f\"No valid DICOM files found:\\n{series_path}\"\n        )\n\n    # --------------------------------------------------------\n    # PHYSICAL SORTING\n    # --------------------------------------------------------\n\n    dicom_files = sorted(\n        dicom_files,\n        key=lambda x: x[\"coordinate\"]\n    )\n\n    # --------------------------------------------------------\n    # SLICE SELECTION\n    # --------------------------------------------------------\n\n    selected_indices = sample_slice_indices(\n        len(dicom_files),\n        target_slices\n    )\n\n    selected_files = [\n        dicom_files[i]\n        for i in selected_indices\n    ]\n\n    images = []\n\n    # --------------------------------------------------------\n    # READ + NORMALIZE + RESIZE\n    # --------------------------------------------------------\n\n    for item in selected_files:\n\n        ds = pydicom.dcmread(\n            item[\"file\"],\n            stop_before_pixels=False\n        )\n\n        image = ds.pixel_array\n\n        image = normalize_mri(image)\n\n        image = Image.fromarray(image)\n\n        image = image.resize(\n            image_size,\n            Image.Resampling.BILINEAR\n        )\n\n        image = np.asarray(\n            image,\n            dtype=np.float32\n        )\n\n        images.append(image)\n\n    return np.stack(images, axis=0)\n\n\ndef load_study(\n    study_id,\n    series_manifest,\n    target_slices=12,\n    image_size=(256, 256)\n):\n    \"\"\"\n    Load every MRI series belonging to one study.\n    \"\"\"\n\n    study_rows = series_manifest[\n        series_manifest[\"StudyInstanceUID\"].astype(str)\n        == str(study_id)\n    ].copy()\n\n    if len(study_rows) == 0:\n        raise ValueError(\n            f\"Study not found: {study_id}\"\n        )\n\n    loaded_series = []\n\n    for _, row in study_rows.iterrows():\n\n        series_path = row[\"SeriesPath\"]\n\n        images = load_dicom_series(\n            series_path=series_path,\n            target_slices=target_slices,\n            image_size=image_size\n        )\n\n        loaded_series.append(\n            {\n                \"StudyInstanceUID\":\n                    row[\"StudyInstanceUID\"],\n\n                \"SeriesInstanceUID\":\n                    row[\"SeriesInstanceUID\"],\n\n                \"Fluid_Sensitive\":\n                    row[\"Fluid_Sensitive\"],\n\n                \"Fat_Suppression\":\n                    row[\"Fat_Suppression\"],\n\n                \"Anatomical_Plane\":\n                    row[\"Anatomical_Plane\"],\n\n                \"Number_of_Slices\":\n                    int(row[\"Number_of_Slices\"]),\n\n                \"Images\":\n                    images\n            }\n        )\n\n    return loaded_series\n\n\nprint(\"=\" * 80)\nprint(\"STUDY LOADER DEFINED\")\nprint(\"=\" * 80)\n\nprint(\"load_dicom_series() : OK\")\nprint(\"load_study()        : OK\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:49.284065Z","iopub.execute_input":"2026-08-19T04:44:49.284513Z","iopub.status.idle":"2026-08-19T04:44:49.305446Z","shell.execute_reply.started":"2026-08-19T04:44:49.284482Z","shell.execute_reply":"2026-08-19T04:44:49.304816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 40 — COMPLETE SELF-CONTAINED STUDY LOADER\n# ============================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport cv2\n\n\n# ------------------------------------------------------------\n# 1. NORMALIZE MRI\n# ------------------------------------------------------------\n\ndef normalize_mri(image):\n\n    image = image.astype(np.float32)\n\n    min_val = image.min()\n    max_val = image.max()\n\n    if max_val > min_val:\n        image = (\n            image - min_val\n        ) / (\n            max_val - min_val\n        )\n    else:\n        image = np.zeros_like(image)\n\n    return image\n\n\n# ------------------------------------------------------------\n# 2. GET DICOM FILES\n# ------------------------------------------------------------\n\ndef get_dicom_files(series_path):\n\n    files = []\n\n    for filename in os.listdir(series_path):\n\n        if filename.lower().endswith(\".dcm\"):\n\n            files.append(\n                os.path.join(\n                    series_path,\n                    filename\n                )\n            )\n\n    return files\n\n\n# ------------------------------------------------------------\n# 3. PHYSICAL SLICE SORTING\n# ------------------------------------------------------------\n\ndef get_sorted_dicom_files_physical(series_path):\n\n    dicom_records = []\n\n    files = get_dicom_files(series_path)\n\n    for filepath in files:\n\n        try:\n\n            ds = pydicom.dcmread(\n                filepath,\n                stop_before_pixels=True\n            )\n\n            position = getattr(\n                ds,\n                \"ImagePositionPatient\",\n                None\n            )\n\n            orientation = getattr(\n                ds,\n                \"ImageOrientationPatient\",\n                None\n            )\n\n            instance_number = getattr(\n                ds,\n                \"InstanceNumber\",\n                0\n            )\n\n            # ----------------------------------------------\n            # Calculate physical slice coordinate\n            # ----------------------------------------------\n\n            if (\n                position is not None\n                and orientation is not None\n                and len(position) >= 3\n                and len(orientation) >= 6\n            ):\n\n                row_cosine = np.array(\n                    orientation[:3],\n                    dtype=np.float64\n                )\n\n                col_cosine = np.array(\n                    orientation[3:6],\n                    dtype=np.float64\n                )\n\n                normal = np.cross(\n                    row_cosine,\n                    col_cosine\n                )\n\n                position_array = np.array(\n                    position[:3],\n                    dtype=np.float64\n                )\n\n                slice_coordinate = np.dot(\n                    position_array,\n                    normal\n                )\n\n            else:\n\n                slice_coordinate = float(\n                    instance_number\n                )\n\n            dicom_records.append({\n\n                \"file\": filepath,\n\n                \"InstanceNumber\":\n                    int(instance_number),\n\n                \"SliceCoordinate\":\n                    float(slice_coordinate)\n            })\n\n        except Exception:\n\n            continue\n\n    if len(dicom_records) == 0:\n\n        return pd.DataFrame(\n            columns=[\n                \"file\",\n                \"InstanceNumber\",\n                \"SliceCoordinate\"\n            ]\n        )\n\n    sorted_df = pd.DataFrame(\n        dicom_records\n    ).sort_values(\n        \"SliceCoordinate\"\n    ).reset_index(\n        drop=True\n    )\n\n    return sorted_df\n\n\n# ------------------------------------------------------------\n# 4. CENTRAL 80% SLICE SAMPLING\n# ------------------------------------------------------------\n\ndef sample_mri_slices(\n    num_slices,\n    target_slices=12,\n    central_fraction=0.80\n):\n\n    if num_slices <= 0:\n\n        return np.array(\n            [],\n            dtype=int\n        )\n\n    # ----------------------------------------------\n    # If series is small, use all slices\n    # ----------------------------------------------\n\n    if num_slices <= target_slices:\n\n        return np.arange(\n            num_slices\n        )\n\n    # ----------------------------------------------\n    # Central region\n    # ----------------------------------------------\n\n    margin = (\n        1.0 - central_fraction\n    ) / 2.0\n\n    start = int(\n        np.floor(\n            num_slices * margin\n        )\n    )\n\n    end = int(\n        np.ceil(\n            num_slices * (1.0 - margin)\n        )\n    ) - 1\n\n    # ----------------------------------------------\n    # Uniformly sample central region\n    # ----------------------------------------------\n\n    indices = np.linspace(\n        start,\n        end,\n        target_slices\n    )\n\n    indices = np.round(\n        indices\n    ).astype(int)\n\n    indices = np.unique(\n        indices\n    )\n\n    return indices\n\n\n# ------------------------------------------------------------\n# 5. LOAD ONE SERIES\n# ------------------------------------------------------------\n\ndef load_series_images(\n    series_path,\n    target_slices=12,\n    image_size=256,\n    central_fraction=0.80\n):\n\n    sorted_df = (\n        get_sorted_dicom_files_physical(\n            series_path\n        )\n    )\n\n    num_slices = len(\n        sorted_df\n    )\n\n    if num_slices == 0:\n\n        return (\n            np.empty(\n                (\n                    0,\n                    image_size,\n                    image_size\n                ),\n                dtype=np.float32\n            ),\n            []\n        )\n\n    selected_indices = sample_mri_slices(\n        num_slices=num_slices,\n        target_slices=target_slices,\n        central_fraction=central_fraction\n    )\n\n    images = []\n\n    metadata = []\n\n    for idx in selected_indices:\n\n        filepath = (\n            sorted_df.iloc[idx][\"file\"]\n        )\n\n        try:\n\n            # ------------------------------------------\n            # Read complete DICOM\n            # ------------------------------------------\n\n            ds = pydicom.dcmread(\n                filepath\n            )\n\n            # ------------------------------------------\n            # Pixel data\n            # ------------------------------------------\n\n            image = ds.pixel_array\n\n            # ------------------------------------------\n            # Normalize\n            # ------------------------------------------\n\n            image = normalize_mri(\n                image\n            )\n\n            # ------------------------------------------\n            # Resize\n            # ------------------------------------------\n\n            image = cv2.resize(\n                image,\n                (\n                    image_size,\n                    image_size\n                ),\n                interpolation=cv2.INTER_AREA\n            )\n\n            images.append(\n                image.astype(\n                    np.float32\n                )\n            )\n\n            metadata.append({\n\n                \"index\":\n                    int(idx),\n\n                \"InstanceNumber\":\n                    int(\n                        sorted_df.iloc[idx]\n                        [\"InstanceNumber\"]\n                    ),\n\n                \"SliceCoordinate\":\n                    float(\n                        sorted_df.iloc[idx]\n                        [\"SliceCoordinate\"]\n                    ),\n\n                \"file\":\n                    filepath\n            })\n\n        except Exception as e:\n\n            print(\n                \"Failed:\",\n                filepath\n            )\n\n            print(\n                \"Error:\",\n                e\n            )\n\n    if len(images) > 0:\n\n        images_array = np.stack(\n            images\n        ).astype(\n            np.float32\n        )\n\n    else:\n\n        images_array = np.empty(\n            (\n                0,\n                image_size,\n                image_size\n            ),\n            dtype=np.float32\n        )\n\n    return (\n        images_array,\n        metadata\n    )\n\n\n# ------------------------------------------------------------\n# 6. LOAD COMPLETE STUDY\n# ------------------------------------------------------------\n\ndef load_study(\n    study_id,\n    series_manifest,\n    target_slices=12,\n    image_size=256,\n    central_fraction=0.80\n):\n\n    study_rows = (\n        series_manifest[\n            series_manifest[\n                \"StudyInstanceUID\"\n            ] == study_id\n        ]\n        .copy()\n    )\n\n    study_data = []\n\n    for _, row in study_rows.iterrows():\n\n        images, slice_metadata = (\n            load_series_images(\n                series_path=row[\n                    \"SeriesPath\"\n                ],\n\n                target_slices=\n                    target_slices,\n\n                image_size=\n                    image_size,\n\n                central_fraction=\n                    central_fraction\n            )\n        )\n\n        study_data.append({\n\n            \"StudyInstanceUID\":\n                row[\n                    \"StudyInstanceUID\"\n                ],\n\n            \"SeriesInstanceUID\":\n                row[\n                    \"SeriesInstanceUID\"\n                ],\n\n            \"Anatomical_Plane\":\n                row[\n                    \"Anatomical_Plane\"\n                ],\n\n            \"Fluid_Sensitive\":\n                int(\n                    row[\n                        \"Fluid_Sensitive\"\n                    ]\n                ),\n\n            \"Fat_Suppression\":\n                int(\n                    row[\n                        \"Fat_Suppression\"\n                    ]\n                ),\n\n            \"Number_of_Slices\":\n                int(\n                    row[\n                        \"Number_of_Slices\"\n                    ]\n                ),\n\n            \"Images\":\n                images,\n\n            \"Slice_Metadata\":\n                slice_metadata\n        })\n\n    return study_data\n\n\nprint(\"=\" * 80)\nprint(\"STUDY LOADER READY\")\nprint(\"=\" * 80)\n\nprint(\"normalize_mri()                ✓\")\nprint(\"get_sorted_dicom_files_physical() ✓\")\nprint(\"sample_mri_slices()            ✓\")\nprint(\"load_series_images()           ✓\")\nprint(\"load_study()                   ✓\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:49.30634Z","iopub.execute_input":"2026-08-19T04:44:49.306605Z","iopub.status.idle":"2026-08-19T04:44:49.327886Z","shell.execute_reply.started":"2026-08-19T04:44:49.306574Z","shell.execute_reply":"2026-08-19T04:44:49.327377Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 41 — TEST COMPLETE STUDY LOADER\n# ============================================================\n\nexample_study = (\n    series_manifest[\n        \"StudyInstanceUID\"\n    ].iloc[0]\n)\n\nprint(\"=\" * 80)\nprint(\"LOADING STUDY\")\nprint(\"=\" * 80)\n\nprint(\n    \"Study ID:\",\n    example_study\n)\n\nloaded_study = load_study(\n    study_id=example_study,\n    series_manifest=series_manifest,\n    target_slices=12,\n    image_size=256,\n    central_fraction=0.80\n)\n\nprint(\n    \"\\nSeries loaded:\",\n    len(loaded_study)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:49.332303Z","iopub.execute_input":"2026-08-19T04:44:49.332494Z","iopub.status.idle":"2026-08-19T04:44:49.775067Z","shell.execute_reply.started":"2026-08-19T04:44:49.332476Z","shell.execute_reply":"2026-08-19T04:44:49.774385Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 42 — VERIFY LOADED STUDY\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"LOADED STUDY SUMMARY\")\nprint(\"=\" * 80)\n\ntotal_selected = 0\n\nfor i, series in enumerate(\n    loaded_study,\n    start=1\n):\n\n    images = series[\"Images\"]\n\n    n_selected = len(images)\n\n    total_selected += n_selected\n\n    print(f\"\\nSeries {i}\")\n    print(\"-\" * 50)\n\n    print(\n        \"Plane           :\",\n        series[\"Anatomical_Plane\"]\n    )\n\n    print(\n        \"Fluid           :\",\n        series[\"Fluid_Sensitive\"]\n    )\n\n    print(\n        \"Fat Suppression :\",\n        series[\"Fat_Suppression\"]\n    )\n\n    print(\n        \"Original slices :\",\n        series[\"Number_of_Slices\"]\n    )\n\n    print(\n        \"Selected slices :\",\n        n_selected\n    )\n\n    print(\n        \"Image shape     :\",\n        images.shape\n    )\n\nprint(\"\\n\" + \"=\" * 80)\n\nprint(\n    \"TOTAL SELECTED IMAGES:\",\n    total_selected\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:49.775873Z","iopub.execute_input":"2026-08-19T04:44:49.776157Z","iopub.status.idle":"2026-08-19T04:44:49.783053Z","shell.execute_reply.started":"2026-08-19T04:44:49.776135Z","shell.execute_reply":"2026-08-19T04:44:49.782279Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 43 — VISUALIZE COMPLETE LOADED STUDY\n# ============================================================\n\nimport matplotlib.pyplot as plt\nimport numpy as np\n\nprint(\"=\" * 80)\nprint(\"VISUALIZING COMPLETE LOADED STUDY\")\nprint(\"=\" * 80)\n\nfig, axes = plt.subplots(\n    nrows=len(loaded_study),\n    ncols=6,\n    figsize=(18, 3.5 * len(loaded_study))\n)\n\nfor row_idx, series in enumerate(loaded_study):\n\n    images = series[\"Images\"]\n\n    # Select 6 evenly distributed slices\n    display_indices = np.linspace(\n        0,\n        len(images) - 1,\n        6\n    ).round().astype(int)\n\n    for col_idx, image_idx in enumerate(display_indices):\n\n        ax = axes[row_idx, col_idx]\n\n        ax.imshow(\n            images[image_idx],\n            cmap=\"gray\",\n            vmin=0,\n            vmax=1\n        )\n\n        ax.set_title(\n            f\"{series['Anatomical_Plane']} | \"\n            f\"Slice {image_idx + 1}/{len(images)}\",\n            fontsize=10\n        )\n\n        ax.axis(\"off\")\n\n    # Add series information on left\n    axes[row_idx, 0].set_ylabel(\n        f\"Series {row_idx + 1}\\n\"\n        f\"Fluid={series['Fluid_Sensitive']}\\n\"\n        f\"FatSup={series['Fat_Suppression']}\",\n        fontsize=11\n    )\n\nplt.suptitle(\n    \"RSNA Knee MRI — Complete Loaded Study\",\n    fontsize=18\n)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:49.784022Z","iopub.execute_input":"2026-08-19T04:44:49.784364Z","iopub.status.idle":"2026-08-19T04:44:52.054925Z","shell.execute_reply.started":"2026-08-19T04:44:49.784341Z","shell.execute_reply":"2026-08-19T04:44:52.0517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 44 — BUILD STUDY-LEVEL DATASET SUMMARY\n# ============================================================\n\nLABEL_COLUMNS = [\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\n# ------------------------------------------------------------\n# Series statistics per study\n# ------------------------------------------------------------\n\nstudy_summary = (\n    series_manifest\n    .groupby(\"StudyInstanceUID\")\n    .agg(\n        Number_of_Series=(\n            \"SeriesInstanceUID\",\n            \"nunique\"\n        ),\n\n        Total_Slices=(\n            \"Number_of_Slices\",\n            \"sum\"\n        ),\n\n        Mean_Slices_Per_Series=(\n            \"Number_of_Slices\",\n            \"mean\"\n        ),\n\n        Min_Slices_Per_Series=(\n            \"Number_of_Slices\",\n            \"min\"\n        ),\n\n        Max_Slices_Per_Series=(\n            \"Number_of_Slices\",\n            \"max\"\n        )\n    )\n    .reset_index()\n)\n\n# ------------------------------------------------------------\n# Count planes\n# ------------------------------------------------------------\n\nplane_table = pd.crosstab(\n    series_manifest[\"StudyInstanceUID\"],\n    series_manifest[\"Anatomical_Plane\"]\n).reset_index()\n\nplane_table = plane_table.rename(\n    columns={\n        \"Sagittal\": \"Sagittal_Series\",\n        \"Coronal\": \"Coronal_Series\",\n        \"Axial\": \"Axial_Series\"\n    }\n)\n\n# ------------------------------------------------------------\n# Merge\n# ------------------------------------------------------------\n\nstudy_summary = study_summary.merge(\n    plane_table,\n    on=\"StudyInstanceUID\",\n    how=\"left\"\n)\n\n# ------------------------------------------------------------\n# Add labels and reports\n# ------------------------------------------------------------\n\nstudy_summary = study_summary.merge(\n    train[\n        [\"StudyInstanceUID\", \"Report\"] + LABEL_COLUMNS\n    ],\n    on=\"StudyInstanceUID\",\n    how=\"left\"\n)\n\n# ------------------------------------------------------------\n# Number of labeled abnormalities\n# ------------------------------------------------------------\n\nstudy_summary[\"Number_of_Labeled_Abnormalities\"] = (\n    study_summary[LABEL_COLUMNS]\n    .notna()\n    .sum(axis=1)\n)\n\n# ------------------------------------------------------------\n# Display\n# ------------------------------------------------------------\n\nprint(\"=\" * 80)\nprint(\"STUDY-LEVEL DATASET SUMMARY\")\nprint(\"=\" * 80)\n\nprint(\n    \"Studies:\",\n    len(study_summary)\n)\n\nprint(\n    \"Columns:\",\n    len(study_summary.columns)\n)\n\ndisplay(\n    study_summary.head()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.056007Z","iopub.execute_input":"2026-08-19T04:44:52.056654Z","iopub.status.idle":"2026-08-19T04:44:52.237579Z","shell.execute_reply.started":"2026-08-19T04:44:52.056628Z","shell.execute_reply":"2026-08-19T04:44:52.236728Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 45 — LABELED STUDY ANALYSIS\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"LABELED vs UNLABELED STUDIES\")\nprint(\"=\" * 80)\n\n# A study is labeled if at least one of the 12 labels is present\nlabeled_mask = (\n    study_summary[LABEL_COLUMNS]\n    .notna()\n    .any(axis=1)\n)\n\nlabeled_studies = study_summary[\n    labeled_mask\n].copy()\n\nunlabeled_studies = study_summary[\n    ~labeled_mask\n].copy()\n\nprint(\n    \"Total studies     :\",\n    len(study_summary)\n)\n\nprint(\n    \"Labeled studies   :\",\n    len(labeled_studies)\n)\n\nprint(\n    \"Unlabeled studies :\",\n    len(unlabeled_studies)\n)\n\nprint(\"\\nExpected:\")\nprint(\"Labeled   = 58\")\nprint(\"Unlabeled = 4349\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.238593Z","iopub.execute_input":"2026-08-19T04:44:52.239278Z","iopub.status.idle":"2026-08-19T04:44:52.249444Z","shell.execute_reply.started":"2026-08-19T04:44:52.239211Z","shell.execute_reply":"2026-08-19T04:44:52.248769Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 46 — LABEL COUNTS\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"EXPLICIT LABEL DISTRIBUTION\")\nprint(\"=\" * 80)\n\nlabel_results = []\n\nfor label in LABEL_COLUMNS:\n\n    positive = (\n        labeled_studies[label] == 1\n    ).sum()\n\n    negative = (\n        labeled_studies[label] == 0\n    ).sum()\n\n    unknown = (\n        labeled_studies[label].isna()\n    ).sum()\n\n    label_results.append({\n\n        \"Abnormality\": label,\n\n        \"Positive\": int(\n            positive\n        ),\n\n        \"Negative\": int(\n            negative\n        ),\n\n        \"Unknown\": int(\n            unknown\n        ),\n\n        \"Positive_%\": round(\n            positive /\n            max(positive + negative, 1)\n            * 100,\n            2\n        )\n    })\n\nlabel_distribution = pd.DataFrame(\n    label_results\n)\n\ndisplay(\n    label_distribution\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.25042Z","iopub.execute_input":"2026-08-19T04:44:52.251374Z","iopub.status.idle":"2026-08-19T04:44:52.268923Z","shell.execute_reply.started":"2026-08-19T04:44:52.251339Z","shell.execute_reply":"2026-08-19T04:44:52.268318Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 47 — INSPECT LABELED REPORTS\n# ============================================================\n\ndisplay_columns = [\n    \"StudyInstanceUID\",\n    \"Report\"\n] + LABEL_COLUMNS\n\nlabeled_report_table = (\n    labeled_studies[\n        display_columns\n    ]\n    .reset_index(drop=True)\n)\n\nprint(\"=\" * 80)\nprint(\"LABELED STUDIES + RADIOLOGY REPORTS\")\nprint(\"=\" * 80)\n\nprint(\n    \"Number of labeled studies:\",\n    len(labeled_report_table)\n)\n\ndisplay(\n    labeled_report_table\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.269636Z","iopub.execute_input":"2026-08-19T04:44:52.269818Z","iopub.status.idle":"2026-08-19T04:44:52.326793Z","shell.execute_reply.started":"2026-08-19T04:44:52.269801Z","shell.execute_reply":"2026-08-19T04:44:52.32594Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 48 — EXACT 12-ABNORMALITY LABEL STATISTICS\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"12 ABNORMALITY LABEL STATISTICS — 58 LABELED STUDIES\")\nprint(\"=\" * 80)\n\nlabel_stats = []\n\nfor label in LABEL_COLUMNS:\n\n    positive = int((labeled_studies[label] == 1).sum())\n    negative = int((labeled_studies[label] == 0).sum())\n    unknown = int(labeled_studies[label].isna().sum())\n\n    label_stats.append({\n        \"Abnormality\": label,\n        \"Positive\": positive,\n        \"Negative\": negative,\n        \"Unknown\": unknown,\n        \"Labeled_Total\": positive + negative\n    })\n\nlabel_stats_df = pd.DataFrame(label_stats)\n\ndisplay(label_stats_df)\n\nprint(\"\\nTotal positive labels:\")\nprint(\n    label_stats_df[\"Positive\"].sum()\n)\n\nprint(\"\\nTotal negative labels:\")\nprint(\n    label_stats_df[\"Negative\"].sum()\n)\n\nprint(\"\\nTotal label entries:\")\nprint(\n    label_stats_df[\"Labeled_Total\"].sum()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.327824Z","iopub.execute_input":"2026-08-19T04:44:52.328102Z","iopub.status.idle":"2026-08-19T04:44:52.347968Z","shell.execute_reply.started":"2026-08-19T04:44:52.328073Z","shell.execute_reply":"2026-08-19T04:44:52.347413Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 49 — ABNORMALITIES PER STUDY\n# ============================================================\n\nlabeled_studies[\"Total_Abnormalities\"] = (\n    labeled_studies[LABEL_COLUMNS]\n    .sum(axis=1)\n)\n\nprint(\"=\" * 80)\nprint(\"NUMBER OF ABNORMALITIES PER LABELED STUDY\")\nprint(\"=\" * 80)\n\nprint(\n    labeled_studies[\"Total_Abnormalities\"]\n    .describe()\n)\n\nprint(\"\\nDistribution:\")\n\nabnormality_distribution = (\n    labeled_studies[\"Total_Abnormalities\"]\n    .value_counts()\n    .sort_index()\n    .rename_axis(\"Number_of_Abnormalities\")\n    .reset_index(name=\"Number_of_Studies\")\n)\n\ndisplay(abnormality_distribution)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.348982Z","iopub.execute_input":"2026-08-19T04:44:52.349587Z","iopub.status.idle":"2026-08-19T04:44:52.370528Z","shell.execute_reply.started":"2026-08-19T04:44:52.349565Z","shell.execute_reply":"2026-08-19T04:44:52.369847Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 50 — ABNORMALITY CORRELATION\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"ABNORMALITY LABEL CORRELATION\")\nprint(\"=\" * 80)\n\ncorrelation_matrix = (\n    labeled_studies[LABEL_COLUMNS]\n    .corr()\n)\n\ndisplay(\n    correlation_matrix.round(2)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.371472Z","iopub.execute_input":"2026-08-19T04:44:52.371685Z","iopub.status.idle":"2026-08-19T04:44:52.395967Z","shell.execute_reply.started":"2026-08-19T04:44:52.371665Z","shell.execute_reply":"2026-08-19T04:44:52.395428Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 51 — FINAL LABEL BALANCE\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"FINAL LABEL BALANCE — 58 LABELED STUDIES\")\nprint(\"=\" * 80)\n\nlabel_balance = []\n\nfor label in LABEL_COLUMNS:\n\n    positive = int(\n        (labeled_studies[label] == 1).sum()\n    )\n\n    negative = int(\n        (labeled_studies[label] == 0).sum()\n    )\n\n    total = positive + negative\n\n    positive_percent = (\n        positive / total * 100\n        if total > 0 else 0\n    )\n\n    negative_percent = (\n        negative / total * 100\n        if total > 0 else 0\n    )\n\n    label_balance.append({\n        \"Abnormality\": label,\n        \"Positive\": positive,\n        \"Negative\": negative,\n        \"Positive_%\": round(positive_percent, 2),\n        \"Negative_%\": round(negative_percent, 2)\n    })\n\nlabel_balance_df = pd.DataFrame(label_balance)\n\ndisplay(label_balance_df)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TOTALS\")\nprint(\"=\" * 80)\n\nprint(\n    \"Total positive labels :\",\n    label_balance_df[\"Positive\"].sum()\n)\n\nprint(\n    \"Total negative labels :\",\n    label_balance_df[\"Negative\"].sum()\n)\n\nprint(\n    \"Total label entries   :\",\n    label_balance_df[\"Positive\"].sum()\n    + label_balance_df[\"Negative\"].sum()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.396874Z","iopub.execute_input":"2026-08-19T04:44:52.397374Z","iopub.status.idle":"2026-08-19T04:44:52.421123Z","shell.execute_reply.started":"2026-08-19T04:44:52.397352Z","shell.execute_reply":"2026-08-19T04:44:52.42056Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 52 — STUDY-LEVEL TRAIN / VALIDATION SPLIT\n# ============================================================\n\nfrom sklearn.model_selection import train_test_split\n\nprint(\"=\" * 80)\nprint(\"STUDY-LEVEL TRAIN / VALIDATION SPLIT\")\nprint(\"=\" * 80)\n\n# ------------------------------------------------------------\n# 1. Confirm labeled studies\n# ------------------------------------------------------------\n\nprint(\"Total labeled studies:\", len(labeled_studies))\n\n# ------------------------------------------------------------\n# 2. Remove duplicate studies if any\n# ------------------------------------------------------------\n\nlabeled_studies = (\n    labeled_studies\n    .drop_duplicates(subset=[\"StudyInstanceUID\"])\n    .reset_index(drop=True)\n)\n\nprint(\n    \"Unique labeled studies:\",\n    labeled_studies[\"StudyInstanceUID\"].nunique()\n)\n\n# ------------------------------------------------------------\n# 3. Train / validation split\n# ------------------------------------------------------------\n\ntrain_studies, val_studies = train_test_split(\n    labeled_studies,\n    test_size=0.20,\n    random_state=42,\n    shuffle=True\n)\n\ntrain_studies = train_studies.reset_index(drop=True)\nval_studies = val_studies.reset_index(drop=True)\n\n# ------------------------------------------------------------\n# 4. Display split sizes\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"SPLIT SUMMARY\")\nprint(\"=\" * 80)\n\nprint(\"Total studies     :\", len(labeled_studies))\nprint(\"Training studies  :\", len(train_studies))\nprint(\"Validation studies:\", len(val_studies))\n\nprint(\"\\nTraining percentage  :\",\n      round(len(train_studies) / len(labeled_studies) * 100, 2),\n      \"%\")\n\nprint(\"Validation percentage:\",\n      round(len(val_studies) / len(labeled_studies) * 100, 2),\n      \"%\")\n\n# ------------------------------------------------------------\n# 5. Check study leakage\n# ------------------------------------------------------------\n\ntrain_ids = set(\n    train_studies[\"StudyInstanceUID\"]\n)\n\nval_ids = set(\n    val_studies[\"StudyInstanceUID\"]\n)\n\noverlap = train_ids.intersection(val_ids)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"LEAKAGE CHECK\")\nprint(\"=\" * 80)\n\nprint(\"Training studies :\", len(train_ids))\nprint(\"Validation studies:\", len(val_ids))\nprint(\"Overlapping studies:\", len(overlap))\n\nif len(overlap) == 0:\n    print(\"STATUS: PASS — No study leakage detected.\")\nelse:\n    print(\"STATUS: FAIL — Study leakage detected!\")\n\n# ------------------------------------------------------------\n# 6. Label distribution in each split\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAINING LABEL DISTRIBUTION\")\nprint(\"=\" * 80)\n\ntrain_distribution = train_studies[LABEL_COLUMNS].sum().sort_values(\n    ascending=False\n)\n\ndisplay(\n    train_distribution.to_frame(\"Positive_Count\")\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"VALIDATION LABEL DISTRIBUTION\")\nprint(\"=\" * 80)\n\nval_distribution = val_studies[LABEL_COLUMNS].sum().sort_values(\n    ascending=False\n)\n\ndisplay(\n    val_distribution.to_frame(\"Positive_Count\")\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.42202Z","iopub.execute_input":"2026-08-19T04:44:52.422319Z","iopub.status.idle":"2026-08-19T04:44:52.44736Z","shell.execute_reply.started":"2026-08-19T04:44:52.422289Z","shell.execute_reply":"2026-08-19T04:44:52.446627Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 53 — MULTI-LABEL STRATIFIED TRAIN / VALIDATION SPLIT\n# ============================================================\n\nimport numpy as np\nimport pandas as pd\n\nprint(\"=\" * 80)\nprint(\"MULTI-LABEL STRATIFIED SPLIT\")\nprint(\"=\" * 80)\n\n# ------------------------------------------------------------\n# LABEL COLUMNS\n# ------------------------------------------------------------\n\nprint(\"Number of labels:\", len(LABEL_COLUMNS))\nprint(\"Labels:\")\n\nfor i, label in enumerate(LABEL_COLUMNS, 1):\n    print(f\"{i:2d}. {label}\")\n\n# ------------------------------------------------------------\n# PREPARE X AND Y\n# ------------------------------------------------------------\n\nX = labeled_studies[\n    [\"StudyInstanceUID\"]\n].copy()\n\nY = (\n    labeled_studies[LABEL_COLUMNS]\n    .fillna(0)\n    .astype(int)\n)\n\nprint(\"\\nDataset shape:\")\nprint(\"X:\", X.shape)\nprint(\"Y:\", Y.shape)\n\n# ------------------------------------------------------------\n# CHECK TOTAL LABEL COUNTS\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TOTAL LABEL DISTRIBUTION\")\nprint(\"=\" * 80)\n\ndisplay(\n    Y.sum()\n    .sort_values(ascending=False)\n    .to_frame(\"Total_Positive\")\n)\n\n# ------------------------------------------------------------\n# LOAD MULTI-LABEL STRATIFIER\n# ------------------------------------------------------------\n\ntry:\n\n    from iterstrat.ml_stratifiers import (\n        MultilabelStratifiedShuffleSplit\n    )\n\n    print(\"\\nMulti-label stratifier available.\")\n\nexcept ImportError:\n\n    print(\"\\nInstalling iterative-stratification...\")\n\n    !pip install -q iterative-stratification\n\n    from iterstrat.ml_stratifiers import (\n        MultilabelStratifiedShuffleSplit\n    )\n\n# ------------------------------------------------------------\n# CREATE STRATIFIED SPLIT\n# ------------------------------------------------------------\n\nmsss = MultilabelStratifiedShuffleSplit(\n    n_splits=1,\n    test_size=0.20,\n    random_state=42\n)\n\ntrain_idx, val_idx = next(\n    msss.split(X, Y)\n)\n\n# ------------------------------------------------------------\n# CREATE DATASETS\n# ------------------------------------------------------------\n\ntrain_studies = (\n    labeled_studies\n    .iloc[train_idx]\n    .reset_index(drop=True)\n)\n\nval_studies = (\n    labeled_studies\n    .iloc[val_idx]\n    .reset_index(drop=True)\n)\n\n# ------------------------------------------------------------\n# SPLIT SUMMARY\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STRATIFIED SPLIT SUMMARY\")\nprint(\"=\" * 80)\n\nprint(\"Total studies     :\", len(labeled_studies))\nprint(\"Training studies  :\", len(train_studies))\nprint(\"Validation studies:\", len(val_studies))\n\n# ------------------------------------------------------------\n# LEAKAGE CHECK\n# ------------------------------------------------------------\n\ntrain_ids = set(\n    train_studies[\"StudyInstanceUID\"]\n)\n\nval_ids = set(\n    val_studies[\"StudyInstanceUID\"]\n)\n\noverlap = train_ids.intersection(val_ids)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STUDY LEAKAGE CHECK\")\nprint(\"=\" * 80)\n\nprint(\"Training unique studies :\", len(train_ids))\nprint(\"Validation unique studies:\", len(val_ids))\nprint(\"Overlapping studies     :\", len(overlap))\n\nif len(overlap) == 0:\n    print(\"STATUS: PASS\")\nelse:\n    print(\"STATUS: FAIL\")\n\n# ------------------------------------------------------------\n# LABEL DISTRIBUTION — TRAIN\n# ------------------------------------------------------------\n\ntrain_Y = (\n    train_studies[LABEL_COLUMNS]\n    .fillna(0)\n    .astype(int)\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAINING LABEL DISTRIBUTION\")\nprint(\"=\" * 80)\n\ntrain_counts = (\n    train_Y.sum()\n    .sort_values(ascending=False)\n)\n\ndisplay(\n    train_counts.to_frame(\n        \"Training_Positive\"\n    )\n)\n\n# ------------------------------------------------------------\n# LABEL DISTRIBUTION — VALIDATION\n# ------------------------------------------------------------\n\nval_Y = (\n    val_studies[LABEL_COLUMNS]\n    .fillna(0)\n    .astype(int)\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"VALIDATION LABEL DISTRIBUTION\")\nprint(\"=\" * 80)\n\nval_counts = (\n    val_Y.sum()\n    .sort_values(ascending=False)\n)\n\ndisplay(\n    val_counts.to_frame(\n        \"Validation_Positive\"\n    )\n)\n\n# ------------------------------------------------------------\n# COMPARISON TABLE\n# ------------------------------------------------------------\n\ncomparison = pd.DataFrame({\n    \"Total\": Y.sum(),\n    \"Train\": train_Y.sum(),\n    \"Validation\": val_Y.sum()\n})\n\ncomparison[\"Train_%\"] = (\n    comparison[\"Train\"]\n    / comparison[\"Total\"]\n    * 100\n).round(1)\n\ncomparison[\"Validation_%\"] = (\n    comparison[\"Validation\"]\n    / comparison[\"Total\"]\n    * 100\n).round(1)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"LABEL DISTRIBUTION COMPARISON\")\nprint(\"=\" * 80)\n\ndisplay(comparison)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.44825Z","iopub.execute_input":"2026-08-19T04:44:52.448722Z","iopub.status.idle":"2026-08-19T04:44:52.496984Z","shell.execute_reply.started":"2026-08-19T04:44:52.448698Z","shell.execute_reply":"2026-08-19T04:44:52.496442Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 54 — SELECT PRIMARY MRI SERIES\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"PRIMARY MRI SERIES SELECTION\")\nprint(\"=\" * 80)\n\n# ------------------------------------------------------------\n# Select fluid-sensitive + fat-suppressed series\n# ------------------------------------------------------------\n\nprimary_series = series_manifest[\n    (series_manifest[\"Fluid_Sensitive\"] == 1) &\n    (series_manifest[\"Fat_Suppression\"] == 1)\n].copy()\n\n# ------------------------------------------------------------\n# Check distribution by plane\n# ------------------------------------------------------------\n\nprint(\"Total selected series:\", len(primary_series))\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"SERIES BY ANATOMICAL PLANE\")\nprint(\"=\" * 80)\n\ndisplay(\n    primary_series[\"Anatomical_Plane\"]\n    .value_counts()\n    .rename_axis(\"Plane\")\n    .reset_index(name=\"Series_Count\")\n)\n\n# ------------------------------------------------------------\n# Check studies represented\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STUDY COVERAGE\")\nprint(\"=\" * 80)\n\ntotal_primary_studies = (\n    primary_series[\"StudyInstanceUID\"]\n    .nunique()\n)\n\nprint(\n    \"Studies with primary series:\",\n    total_primary_studies\n)\n\nprint(\n    \"Total labeled studies:\",\n    labeled_studies[\"StudyInstanceUID\"].nunique()\n)\n\n# ------------------------------------------------------------\n# Merge with training / validation split\n# ------------------------------------------------------------\n\nprimary_train = primary_series[\n    primary_series[\"StudyInstanceUID\"].isin(\n        train_studies[\"StudyInstanceUID\"]\n    )\n].copy()\n\nprimary_val = primary_series[\n    primary_series[\"StudyInstanceUID\"].isin(\n        val_studies[\"StudyInstanceUID\"]\n    )\n].copy()\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAIN / VALIDATION SERIES\")\nprint(\"=\" * 80)\n\nprint(\n    \"Training series:\",\n    len(primary_train)\n)\n\nprint(\n    \"Validation series:\",\n    len(primary_val)\n)\n\nprint(\n    \"Training studies:\",\n    primary_train[\"StudyInstanceUID\"].nunique()\n)\n\nprint(\n    \"Validation studies:\",\n    primary_val[\"StudyInstanceUID\"].nunique()\n)\n\n# ------------------------------------------------------------\n# Series count per study\n# ------------------------------------------------------------\n\ntrain_series_per_study = (\n    primary_train\n    .groupby(\"StudyInstanceUID\")\n    .size()\n)\n\nval_series_per_study = (\n    primary_val\n    .groupby(\"StudyInstanceUID\")\n    .size()\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAINING SERIES COVERAGE PER STUDY\")\nprint(\"=\" * 80)\n\ndisplay(\n    train_series_per_study\n    .value_counts()\n    .sort_index()\n    .rename_axis(\"Number_of_Primary_Series\")\n    .reset_index(name=\"Number_of_Studies\")\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"VALIDATION SERIES COVERAGE PER STUDY\")\nprint(\"=\" * 80)\n\ndisplay(\n    val_series_per_study\n    .value_counts()\n    .sort_index()\n    .rename_axis(\"Number_of_Primary_Series\")\n    .reset_index(name=\"Number_of_Studies\")\n)\n\n# ------------------------------------------------------------\n# Plane coverage\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAINING PLANE DISTRIBUTION\")\nprint(\"=\" * 80)\n\ndisplay(\n    primary_train[\"Anatomical_Plane\"]\n    .value_counts()\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"VALIDATION PLANE DISTRIBUTION\")\nprint(\"=\" * 80)\n\ndisplay(\n    primary_val[\"Anatomical_Plane\"]\n    .value_counts()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.497754Z","iopub.execute_input":"2026-08-19T04:44:52.498026Z","iopub.status.idle":"2026-08-19T04:44:52.540839Z","shell.execute_reply.started":"2026-08-19T04:44:52.498003Z","shell.execute_reply":"2026-08-19T04:44:52.540067Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 55 — PRIMARY SERIES COVERAGE PER LABELED STUDY\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"PRIMARY SERIES COVERAGE PER LABELED STUDY\")\nprint(\"=\" * 80)\n\n# ------------------------------------------------------------\n# All labeled study IDs\n# ------------------------------------------------------------\n\nlabeled_ids = labeled_studies[\n    \"StudyInstanceUID\"\n].unique()\n\n# ------------------------------------------------------------\n# Primary series only\n# ------------------------------------------------------------\n\ncoverage = (\n    primary_series[\n        primary_series[\"StudyInstanceUID\"].isin(labeled_ids)\n    ]\n    .groupby(\"StudyInstanceUID\")\n    [\"Anatomical_Plane\"]\n    .agg(list)\n    .reset_index(name=\"Available_Planes\")\n)\n\n# ------------------------------------------------------------\n# Add all labeled studies\n# ------------------------------------------------------------\n\ncoverage_full = pd.DataFrame({\n    \"StudyInstanceUID\": labeled_ids\n})\n\ncoverage_full = coverage_full.merge(\n    coverage,\n    on=\"StudyInstanceUID\",\n    how=\"left\"\n)\n\ncoverage_full[\"Available_Planes\"] = (\n    coverage_full[\"Available_Planes\"]\n    .apply(\n        lambda x: x if isinstance(x, list) else []\n    )\n)\n\ncoverage_full[\"Number_of_Primary_Series\"] = (\n    coverage_full[\"Available_Planes\"]\n    .apply(len)\n)\n\n# ------------------------------------------------------------\n# Sort planes for consistent display\n# ------------------------------------------------------------\n\ncoverage_full[\"Available_Planes\"] = (\n    coverage_full[\"Available_Planes\"]\n    .apply(\n        lambda x: sorted(x)\n    )\n)\n\n# ------------------------------------------------------------\n# Distribution\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"NUMBER OF PRIMARY SERIES PER STUDY\")\nprint(\"=\" * 80)\n\ndisplay(\n    coverage_full[\n        \"Number_of_Primary_Series\"\n    ]\n    .value_counts()\n    .sort_index()\n    .rename_axis(\"Number_of_Primary_Series\")\n    .reset_index(name=\"Number_of_Studies\")\n)\n\n# ------------------------------------------------------------\n# Exact plane combinations\n# ------------------------------------------------------------\n\ncoverage_full[\"Plane_Combination\"] = (\n    coverage_full[\"Available_Planes\"]\n    .apply(\n        lambda x: \" + \".join(x)\n        if len(x) > 0\n        else \"NONE\"\n    )\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"PLANE COMBINATIONS\")\nprint(\"=\" * 80)\n\ndisplay(\n    coverage_full[\n        \"Plane_Combination\"\n    ]\n    .value_counts()\n    .rename_axis(\"Plane_Combination\")\n    .reset_index(name=\"Number_of_Studies\")\n)\n\n# ------------------------------------------------------------\n# Full coverage\n# ------------------------------------------------------------\n\ncomplete_3_plane = coverage_full[\n    coverage_full[\"Number_of_Primary_Series\"] == 3\n]\n\ntwo_plane = coverage_full[\n    coverage_full[\"Number_of_Primary_Series\"] == 2\n]\n\none_plane = coverage_full[\n    coverage_full[\"Number_of_Primary_Series\"] == 1\n]\n\nzero_plane = coverage_full[\n    coverage_full[\"Number_of_Primary_Series\"] == 0\n]\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"COVERAGE SUMMARY\")\nprint(\"=\" * 80)\n\nprint(\n    \"3-plane studies:\",\n    len(complete_3_plane)\n)\n\nprint(\n    \"2-plane studies:\",\n    len(two_plane)\n)\n\nprint(\n    \"1-plane studies:\",\n    len(one_plane)\n)\n\nprint(\n    \"0-plane studies:\",\n    len(zero_plane)\n)\n\n# ------------------------------------------------------------\n# Check train / validation coverage\n# ------------------------------------------------------------\n\ncoverage_full[\"Split\"] = \"Unknown\"\n\ncoverage_full.loc[\n    coverage_full[\"StudyInstanceUID\"].isin(\n        train_studies[\"StudyInstanceUID\"]\n    ),\n    \"Split\"\n] = \"Train\"\n\ncoverage_full.loc[\n    coverage_full[\"StudyInstanceUID\"].isin(\n        val_studies[\"StudyInstanceUID\"]\n    ),\n    \"Split\"\n] = \"Validation\"\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"COVERAGE BY TRAIN / VALIDATION\")\nprint(\"=\" * 80)\n\ndisplay(\n    pd.crosstab(\n        coverage_full[\"Number_of_Primary_Series\"],\n        coverage_full[\"Split\"]\n    )\n)\n\n# ------------------------------------------------------------\n# Show incomplete studies\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"INCOMPLETE PRIMARY SERIES STUDIES\")\nprint(\"=\" * 80)\n\ndisplay(\n    coverage_full[\n        coverage_full[\"Number_of_Primary_Series\"] < 3\n    ][\n        [\n            \"StudyInstanceUID\",\n            \"Available_Planes\",\n            \"Number_of_Primary_Series\",\n            \"Split\"\n        ]\n    ]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.54253Z","iopub.execute_input":"2026-08-19T04:44:52.543135Z","iopub.status.idle":"2026-08-19T04:44:52.587565Z","shell.execute_reply.started":"2026-08-19T04:44:52.543113Z","shell.execute_reply":"2026-08-19T04:44:52.587005Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 56 — FINAL PRIMARY MRI COVERAGE VERIFICATION\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"FINAL PRIMARY MRI COVERAGE\")\nprint(\"=\" * 80)\n\ncoverage_summary = (\n    coverage_full\n    .groupby(\n        [\"Split\", \"Number_of_Primary_Series\"]\n    )\n    .size()\n    .reset_index(name=\"Number_of_Studies\")\n)\n\ndisplay(coverage_summary)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"PLANE COVERAGE\")\nprint(\"=\" * 80)\n\nplane_presence = pd.DataFrame({\n    \"Sagittal\": coverage_full[\"Available_Planes\"]\n        .apply(lambda x: \"Sagittal\" in x),\n\n    \"Coronal\": coverage_full[\"Available_Planes\"]\n        .apply(lambda x: \"Coronal\" in x),\n\n    \"Axial\": coverage_full[\"Available_Planes\"]\n        .apply(lambda x: \"Axial\" in x)\n})\n\ndisplay(\n    plane_presence.sum()\n    .to_frame(\"Number_of_Studies\")\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"FINAL CHECK\")\nprint(\"=\" * 80)\n\nprint(\n    \"Labeled studies:\",\n    len(labeled_studies)\n)\n\nprint(\n    \"Train studies:\",\n    len(train_studies)\n)\n\nprint(\n    \"Validation studies:\",\n    len(val_studies)\n)\n\nprint(\n    \"3-plane studies:\",\n    (coverage_full[\"Number_of_Primary_Series\"] == 3).sum()\n)\n\nprint(\n    \"2-plane studies:\",\n    (coverage_full[\"Number_of_Primary_Series\"] == 2).sum()\n)\n\nprint(\n    \"1-plane studies:\",\n    (coverage_full[\"Number_of_Primary_Series\"] == 1).sum()\n)\n\nprint(\n    \"0-plane studies:\",\n    (coverage_full[\"Number_of_Primary_Series\"] == 0).sum()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.588537Z","iopub.execute_input":"2026-08-19T04:44:52.588767Z","iopub.status.idle":"2026-08-19T04:44:52.608676Z","shell.execute_reply.started":"2026-08-19T04:44:52.588746Z","shell.execute_reply":"2026-08-19T04:44:52.607816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 57 — MODEL-READY STUDY DATASET LOADER\n# ============================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport pydicom\nfrom PIL import Image\n\nPLANES = [\n    \"Sagittal\",\n    \"Coronal\",\n    \"Axial\"\n]\n\nTARGET_SLICES = 12\nIMAGE_SIZE = (256, 256)\n\n\n# ============================================================\n# NORMALIZATION\n# ============================================================\n\ndef normalize_mri(image):\n\n    image = image.astype(np.float32)\n\n    min_val = image.min()\n    max_val = image.max()\n\n    if max_val > min_val:\n\n        image = (\n            image - min_val\n        ) / (\n            max_val - min_val\n        )\n\n    else:\n\n        image = np.zeros_like(\n            image,\n            dtype=np.float32\n        )\n\n    return image\n\n\n# ============================================================\n# SLICE SAMPLING\n# ============================================================\n\ndef sample_slice_indices(\n    num_slices,\n    target_slices=12\n):\n\n    if num_slices <= 0:\n        return []\n\n    if num_slices <= target_slices:\n\n        return list(\n            range(num_slices)\n        )\n\n    # Use the central 80% of the series.\n    start = int(\n        round(\n            0.10 * (num_slices - 1)\n        )\n    )\n\n    end = int(\n        round(\n            0.90 * (num_slices - 1)\n        )\n    )\n\n    indices = np.linspace(\n        start,\n        end,\n        target_slices\n    )\n\n    indices = np.round(\n        indices\n    ).astype(int)\n\n    indices = np.unique(\n        indices\n    )\n\n    return indices.tolist()\n\n\n# ============================================================\n# GET SORTED DICOM FILES\n# ============================================================\n\ndef get_sorted_dicom_files(series_path):\n\n    files = []\n\n    for filename in os.listdir(series_path):\n\n        if not filename.lower().endswith(\".dcm\"):\n            continue\n\n        filepath = os.path.join(\n            series_path,\n            filename\n        )\n\n        try:\n\n            ds = pydicom.dcmread(\n                filepath,\n                stop_before_pixels=False\n            )\n\n            if not hasattr(\n                ds,\n                \"PixelData\"\n            ):\n                continue\n\n            position = getattr(\n                ds,\n                \"ImagePositionPatient\",\n                None\n            )\n\n            if position is not None:\n\n                coordinate = float(\n                    position[2]\n                )\n\n            else:\n\n                coordinate = float(\n                    getattr(\n                        ds,\n                        \"InstanceNumber\",\n                        0\n                    )\n                )\n\n            files.append({\n                \"file\": filepath,\n                \"coordinate\": coordinate,\n                \"instance\": int(\n                    getattr(\n                        ds,\n                        \"InstanceNumber\",\n                        0\n                    )\n                )\n            })\n\n        except Exception:\n            continue\n\n    files.sort(\n        key=lambda x: x[\"coordinate\"]\n    )\n\n    return files\n\n\n# ============================================================\n# LOAD ONE SERIES\n# ============================================================\n\ndef load_series_images(\n    series_path,\n    target_slices=12,\n    image_size=(256, 256)\n):\n\n    dicom_files = get_sorted_dicom_files(\n        series_path\n    )\n\n    if len(dicom_files) == 0:\n\n        return np.zeros(\n            (\n                target_slices,\n                image_size[1],\n                image_size[0]\n            ),\n            dtype=np.float32\n        )\n\n    selected_indices = sample_slice_indices(\n        len(dicom_files),\n        target_slices\n    )\n\n    selected_files = [\n        dicom_files[i]\n        for i in selected_indices\n    ]\n\n    images = []\n\n    for item in selected_files:\n\n        try:\n\n            ds = pydicom.dcmread(\n                item[\"file\"],\n                stop_before_pixels=False\n            )\n\n            image = ds.pixel_array\n\n            image = normalize_mri(\n                image\n            )\n\n            image = Image.fromarray(\n                image\n            )\n\n            image = image.resize(\n                image_size,\n                Image.Resampling.BILINEAR\n            )\n\n            image = np.asarray(\n                image,\n                dtype=np.float32\n            )\n\n            images.append(\n                image\n            )\n\n        except Exception:\n\n            continue\n\n    # --------------------------------------------------------\n    # Guarantee exactly TARGET_SLICES\n    # --------------------------------------------------------\n\n    if len(images) == 0:\n\n        return np.zeros(\n            (\n                target_slices,\n                image_size[1],\n                image_size[0]\n            ),\n            dtype=np.float32\n        )\n\n    while len(images) < target_slices:\n\n        images.append(\n            images[-1].copy()\n        )\n\n    images = images[\n        :target_slices\n    ]\n\n    return np.stack(\n        images,\n        axis=0\n    ).astype(\n        np.float32\n    )\n\n\n# ============================================================\n# LOAD ONE STUDY\n# ============================================================\n\ndef load_model_study(\n    study_id,\n    series_manifest,\n    target_slices=12,\n    image_size=(256, 256)\n):\n\n    study_rows = series_manifest[\n        series_manifest[\n            \"StudyInstanceUID\"\n        ].astype(str)\n        == str(study_id)\n    ].copy()\n\n    # --------------------------------------------------------\n    # Output containers\n    # --------------------------------------------------------\n\n    study_images = np.zeros(\n        (\n            3,\n            target_slices,\n            image_size[1],\n            image_size[0]\n        ),\n        dtype=np.float32\n    )\n\n    plane_mask = np.zeros(\n        3,\n        dtype=np.float32\n    )\n\n    # --------------------------------------------------------\n    # Load each plane\n    # --------------------------------------------------------\n\n    for plane_index, plane in enumerate(\n        PLANES\n    ):\n\n        # Prefer fluid-sensitive + fat-suppressed\n        candidates = study_rows[\n            (study_rows[\n                \"Anatomical_Plane\"\n            ] == plane)\n            &\n            (study_rows[\n                \"Fluid_Sensitive\"\n            ] == 1)\n            &\n            (study_rows[\n                \"Fat_Suppression\"\n            ] == 1)\n        ]\n\n        if len(candidates) == 0:\n\n            continue\n\n        # If multiple candidates exist,\n        # choose the series with the largest\n        # number of slices.\n\n        candidates = candidates.sort_values(\n            \"Number_of_Slices\",\n            ascending=False\n        )\n\n        row = candidates.iloc[0]\n\n        images = load_series_images(\n            row[\"SeriesPath\"],\n            target_slices=target_slices,\n            image_size=image_size\n        )\n\n        study_images[\n            plane_index\n        ] = images\n\n        plane_mask[\n            plane_index\n        ] = 1.0\n\n    return study_images, plane_mask\n\n\n# ============================================================\n# TEST ONE LABELED STUDY\n# ============================================================\n\ntest_study_id = (\n    labeled_studies[\n        \"StudyInstanceUID\"\n    ].iloc[0]\n)\n\ntest_images, test_mask = load_model_study(\n    test_study_id,\n    series_manifest,\n    target_slices=TARGET_SLICES,\n    image_size=IMAGE_SIZE\n)\n\nprint(\"=\" * 80)\nprint(\"MODEL-READY STUDY TEST\")\nprint(\"=\" * 80)\n\nprint(\n    \"Study ID:\",\n    test_study_id\n)\n\nprint(\n    \"Tensor shape:\",\n    test_images.shape\n)\n\nprint(\n    \"Plane mask:\",\n    test_mask\n)\n\nprint(\n    \"Tensor dtype:\",\n    test_images.dtype\n)\n\nprint(\n    \"Minimum:\",\n    test_images.min()\n)\n\nprint(\n    \"Maximum:\",\n    test_images.max()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.609684Z","iopub.execute_input":"2026-08-19T04:44:52.61019Z","iopub.status.idle":"2026-08-19T04:44:52.930942Z","shell.execute_reply.started":"2026-08-19T04:44:52.610152Z","shell.execute_reply":"2026-08-19T04:44:52.930311Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 58 — VISUALIZE MODEL-READY 3-PLANE INPUT\n# ============================================================\n\nimport matplotlib.pyplot as plt\n\nprint(\"=\" * 80)\nprint(\"MODEL INPUT VISUALIZATION\")\nprint(\"=\" * 80)\n\nplane_names = [\n    \"Sagittal\",\n    \"Coronal\",\n    \"Axial\"\n]\n\n# Select representative middle slice\nmiddle_slice = TARGET_SLICES // 2\n\nfig, axes = plt.subplots(\n    1,\n    3,\n    figsize=(15, 5)\n)\n\nfor i, plane in enumerate(plane_names):\n\n    image = test_images[\n        i,\n        middle_slice\n    ]\n\n    axes[i].imshow(\n        image,\n        cmap=\"gray\"\n    )\n\n    axes[i].set_title(\n        f\"{plane}\\nSlice {middle_slice + 1}/{TARGET_SLICES}\"\n    )\n\n    axes[i].axis(\"off\")\n\nplt.suptitle(\n    \"RSNA Knee MRI — Model Input\",\n    fontsize=16\n)\n\nplt.tight_layout()\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:52.931931Z","iopub.execute_input":"2026-08-19T04:44:52.93227Z","iopub.status.idle":"2026-08-19T04:44:53.261771Z","shell.execute_reply.started":"2026-08-19T04:44:52.932215Z","shell.execute_reply":"2026-08-19T04:44:53.260974Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 58B — VISUALIZE ALL SAMPLED SAGITTAL SLICES\n# ============================================================\n\nfig, axes = plt.subplots(\n    3,\n    4,\n    figsize=(12, 9)\n)\n\nfor i, ax in enumerate(\n    axes.ravel()\n):\n\n    ax.imshow(\n        test_images[0, i],\n        cmap=\"gray\"\n    )\n\n    ax.set_title(\n        f\"Slice {i + 1}\"\n    )\n\n    ax.axis(\"off\")\n\nplt.suptitle(\n    \"Sagittal — 12 Sampled Slices\",\n    fontsize=16\n)\n\nplt.tight_layout()\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T04:44:53.262871Z","iopub.execute_input":"2026-08-19T04:44:53.263184Z","iopub.status.idle":"2026-08-19T04:44:54.038123Z","shell.execute_reply.started":"2026-08-19T04:44:53.263152Z","shell.execute_reply":"2026-08-19T04:44:54.037311Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def __getitem__(self, idx):\n\n    # ============================================================\n    # GET STUDY ROW\n    # ============================================================\n\n    row = self.df.iloc[idx]\n\n    study_id = row[\"StudyInstanceUID\"]\n\n    # ============================================================\n    # LOAD STUDY\n    # ============================================================\n\n    loaded_study = load_study(\n        study_id,\n        self.series_manifest,\n        target_slices=self.target_slices,\n        image_size=self.image_size\n    )\n\n    # ============================================================\n    # INITIALIZE 3-PLANE TENSOR\n    # ============================================================\n\n    images = np.zeros(\n        (\n            3,\n            self.target_slices,\n            self.image_size,\n            self.image_size\n        ),\n        dtype=np.float32\n    )\n\n    # Plane order:\n    # 0 = Sagittal\n    # 1 = Axial\n    # 2 = Coronal\n\n    plane_to_index = {\n        \"Sagittal\": 0,\n        \"Axial\": 1,\n        \"Coronal\": 2\n    }\n\n    plane_mask = np.zeros(\n        3,\n        dtype=np.float32\n    )\n\n    # ============================================================\n    # INSERT LOADED SERIES\n    # ============================================================\n\n    for series in loaded_study:\n\n        plane = series[\"Anatomical_Plane\"]\n\n        if plane not in plane_to_index:\n            continue\n\n        plane_idx = plane_to_index[plane]\n\n        series_images = series[\"Images\"]\n\n        # Convert to numpy\n        series_images = np.asarray(\n            series_images,\n            dtype=np.float32\n        )\n\n        # Safety check\n        if series_images.ndim != 3:\n            raise ValueError(\n                f\"Unexpected image shape for {plane}: \"\n                f\"{series_images.shape}\"\n            )\n\n        # Number of selected slices\n        n = min(\n            series_images.shape[0],\n            self.target_slices\n        )\n\n        # Insert images\n        images[\n            plane_idx,\n            :n\n        ] = series_images[:n]\n\n        # Mark plane as available\n        plane_mask[plane_idx] = 1.0\n\n    # ============================================================\n    # LABELS\n    # ============================================================\n\n    labels = (\n        pd.to_numeric(\n            row[self.label_columns],\n            errors=\"coerce\"\n        )\n        .fillna(0)\n        .astype(np.float32)\n        .to_numpy()\n    )\n\n    # ============================================================\n    # CONVERT TO TORCH\n    # ============================================================\n\n    images = torch.from_numpy(\n        images\n    ).float()\n\n    plane_mask = torch.from_numpy(\n        plane_mask\n    ).float()\n\n    labels = torch.from_numpy(\n        labels\n    ).float()\n\n    # ============================================================\n    # RETURN\n    # ============================================================\n\n    return {\n        \"images\": images,\n        \"plane_mask\": plane_mask,\n        \"labels\": labels,\n        \"study_id\": study_id\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:03:05.939751Z","iopub.execute_input":"2026-08-19T05:03:05.940584Z","iopub.status.idle":"2026-08-19T05:03:05.9491Z","shell.execute_reply.started":"2026-08-19T05:03:05.940551Z","shell.execute_reply":"2026-08-19T05:03:05.948318Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# RSNA KNEE ABNORMALITY DETECTION\n# COMPLETE STUDY-LEVEL DATA PIPELINE\n# ============================================================\n\nimport os\nimport glob\nimport warnings\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\n\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as F\n\nfrom tqdm.auto import tqdm\n\n\n# ============================================================\n# CONFIGURATION\n# ============================================================\n\nTRAIN_DICOM_DIR = (\n    \"/kaggle/input/competitions/\"\n    \"rsna-knee-abnormality-detection/train_series\"\n)\n\nTARGET_SLICES = 12\nIMAGE_SIZE = 256\n\n# ------------------------------------------------------------\n# Plane order used throughout the project\n# ------------------------------------------------------------\n\nPLANE_ORDER = [\n    \"Sagittal\",\n    \"Axial\",\n    \"Coronal\"\n]\n\nPLANE_TO_INDEX = {\n    \"Sagittal\": 0,\n    \"Axial\": 1,\n    \"Coronal\": 2\n}\n\n\n# ============================================================\n# LABEL COLUMNS\n# ============================================================\n\nLABEL_COLUMNS = [\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_LABELS = len(LABEL_COLUMNS)\n\nprint(\"=\" * 80)\nprint(\"RSNA KNEE DATASET CONFIGURATION\")\nprint(\"=\" * 80)\n\nprint(\"Target slices :\", TARGET_SLICES)\nprint(\"Image size    :\", IMAGE_SIZE)\nprint(\"Planes        :\", PLANE_ORDER)\nprint(\"Labels        :\", NUM_LABELS)\n\nprint(\"\\nLabel columns:\")\nfor i, label in enumerate(LABEL_COLUMNS):\n    print(f\"{i:2d}. {label}\")\n\n\n# ============================================================\n# STEP 1 — BASIC DICOM READER\n# ============================================================\n\ndef read_dicom_image(filepath):\n    \"\"\"\n    Read one DICOM file and return a normalized 2D float32 image.\n    \"\"\"\n\n    ds = pydicom.dcmread(\n        filepath,\n        force=True\n    )\n\n    # --------------------------------------------------------\n    # Check Pixel Data\n    # --------------------------------------------------------\n\n    if not hasattr(ds, \"PixelData\"):\n        return None, ds\n\n    try:\n        image = ds.pixel_array.astype(\n            np.float32\n        )\n    except Exception:\n        return None, ds\n\n    # --------------------------------------------------------\n    # Handle multi-frame DICOM\n    # --------------------------------------------------------\n\n    if image.ndim > 2:\n\n        # Use first frame if unexpectedly multi-frame\n        image = image[0]\n\n    # --------------------------------------------------------\n    # Rescale if present\n    # --------------------------------------------------------\n\n    slope = float(\n        getattr(ds, \"RescaleSlope\", 1.0)\n    )\n\n    intercept = float(\n        getattr(ds, \"RescaleIntercept\", 0.0)\n    )\n\n    image = (\n        image * slope\n        + intercept\n    )\n\n    # --------------------------------------------------------\n    # MONOCHROME1 inversion\n    # --------------------------------------------------------\n\n    photometric = getattr(\n        ds,\n        \"PhotometricInterpretation\",\n        \"\"\n    )\n\n    if photometric == \"MONOCHROME1\":\n\n        image = image.max() - image\n\n    return image.astype(np.float32), ds\n\n\n# ============================================================\n# STEP 2 — ROBUST MRI NORMALIZATION\n# ============================================================\n\ndef normalize_mri(image):\n    \"\"\"\n    Robust MRI normalization to approximately [0, 1].\n\n    Uses non-zero percentile clipping to reduce the\n    influence of extreme intensity values.\n    \"\"\"\n\n    image = np.asarray(\n        image,\n        dtype=np.float32\n    )\n\n    # --------------------------------------------------------\n    # Remove NaN / Inf\n    # --------------------------------------------------------\n\n    image = np.nan_to_num(\n        image,\n        nan=0.0,\n        posinf=0.0,\n        neginf=0.0\n    )\n\n    # --------------------------------------------------------\n    # If completely empty\n    # --------------------------------------------------------\n\n    if image.size == 0:\n        return image.astype(np.float32)\n\n    # --------------------------------------------------------\n    # Use non-zero pixels\n    # --------------------------------------------------------\n\n    nonzero = image[\n        np.isfinite(image) &\n        (image > 0)\n    ]\n\n    if nonzero.size < 10:\n\n        min_val = image.min()\n        max_val = image.max()\n\n        if max_val > min_val:\n\n            image = (\n                image - min_val\n            ) / (\n                max_val - min_val\n            )\n\n        else:\n\n            image = np.zeros_like(\n                image,\n                dtype=np.float32\n            )\n\n        return image.astype(np.float32)\n\n    # --------------------------------------------------------\n    # Robust percentile clipping\n    # --------------------------------------------------------\n\n    low = np.percentile(\n        nonzero,\n        1\n    )\n\n    high = np.percentile(\n        nonzero,\n        99\n    )\n\n    if high <= low:\n\n        low = nonzero.min()\n        high = nonzero.max()\n\n    image = np.clip(\n        image,\n        low,\n        high\n    )\n\n    # --------------------------------------------------------\n    # Normalize\n    # --------------------------------------------------------\n\n    image = (\n        image - low\n    ) / (\n        high - low + 1e-8\n    )\n\n    # --------------------------------------------------------\n    # Background remains zero\n    # --------------------------------------------------------\n\n    image[\n        ~np.isfinite(image)\n    ] = 0\n\n    image = np.clip(\n        image,\n        0.0,\n        1.0\n    )\n\n    return image.astype(\n        np.float32\n    )\n\n\n# ============================================================\n# STEP 3 — RESIZE IMAGE\n# ============================================================\n\ndef resize_image(\n    image,\n    image_size=256\n):\n    \"\"\"\n    Resize 2D image to image_size × image_size\n    using PyTorch bilinear interpolation.\n    \"\"\"\n\n    tensor = torch.from_numpy(\n        image.astype(np.float32)\n    )\n\n    tensor = tensor.unsqueeze(\n        0\n    ).unsqueeze(\n        0\n    )\n\n    tensor = F.interpolate(\n        tensor,\n        size=(\n            image_size,\n            image_size\n        ),\n        mode=\"bilinear\",\n        align_corners=False\n    )\n\n    resized = tensor[\n        0,\n        0\n    ].numpy()\n\n    return resized.astype(\n        np.float32\n    )\n\n\n# ============================================================\n# STEP 4 — PHYSICAL SLICE SORTING\n# ============================================================\n\ndef get_slice_coordinate(ds):\n    \"\"\"\n    Calculate physical slice coordinate using\n    ImageOrientationPatient and ImagePositionPatient.\n    \"\"\"\n\n    position = getattr(\n        ds,\n        \"ImagePositionPatient\",\n        None\n    )\n\n    orientation = getattr(\n        ds,\n        \"ImageOrientationPatient\",\n        None\n    )\n\n    if position is None:\n        return None\n\n    position = np.asarray(\n        position,\n        dtype=np.float64\n    )\n\n    # --------------------------------------------------------\n    # If orientation exists\n    # --------------------------------------------------------\n\n    if orientation is not None:\n\n        orientation = np.asarray(\n            orientation,\n            dtype=np.float64\n        )\n\n        if len(orientation) >= 6:\n\n            row_cosines = orientation[\n                0:3\n            ]\n\n            col_cosines = orientation[\n                3:6\n            ]\n\n            normal = np.cross(\n                row_cosines,\n                col_cosines\n            )\n\n            normal_norm = np.linalg.norm(\n                normal\n            )\n\n            if normal_norm > 0:\n\n                normal = (\n                    normal / normal_norm\n                )\n\n                return float(\n                    np.dot(\n                        position,\n                        normal\n                    )\n                )\n\n    # --------------------------------------------------------\n    # Fallback\n    # --------------------------------------------------------\n\n    return float(\n        position[2]\n    )\n\n\n# ============================================================\n# STEP 5 — GET SORTED DICOM FILES\n# ============================================================\n\ndef get_sorted_dicom_files(\n    series_path\n):\n\n    files = glob.glob(\n        os.path.join(\n            series_path,\n            \"*.dcm\"\n        )\n    )\n\n    # --------------------------------------------------------\n    # If no .dcm files found, try all files\n    # --------------------------------------------------------\n\n    if len(files) == 0:\n\n        files = glob.glob(\n            os.path.join(\n                series_path,\n                \"*\"\n            )\n        )\n\n    records = []\n\n    for filepath in files:\n\n        try:\n\n            ds = pydicom.dcmread(\n                filepath,\n                stop_before_pixels=True,\n                force=True\n            )\n\n            # ------------------------------------------------\n            # Ignore files without image geometry\n            # ------------------------------------------------\n\n            if not hasattr(\n                ds,\n                \"ImagePositionPatient\"\n            ):\n\n                continue\n\n            coordinate = get_slice_coordinate(\n                ds\n            )\n\n            instance_number = int(\n                getattr(\n                    ds,\n                    \"InstanceNumber\",\n                    0\n                )\n            )\n\n            records.append(\n                {\n                    \"file\": filepath,\n                    \"coordinate\": coordinate,\n                    \"InstanceNumber\":\n                        instance_number\n                }\n            )\n\n        except Exception:\n            continue\n\n    # --------------------------------------------------------\n    # Physical sorting\n    # --------------------------------------------------------\n\n    if len(records) == 0:\n\n        return []\n\n    records = sorted(\n        records,\n        key=lambda x: (\n            x[\"coordinate\"]\n            if x[\"coordinate\"] is not None\n            else x[\"InstanceNumber\"]\n        )\n    )\n\n    return records\n\n\n# ============================================================\n# STEP 6 — ADAPTIVE CENTRAL 80% SAMPLING\n# ============================================================\n\ndef sample_central_80_indices(\n    num_slices,\n    target_slices=12\n):\n    \"\"\"\n    Select approximately evenly distributed slices\n    from the central 80% of the volume.\n\n    This avoids excessive peripheral slices while\n    maintaining broad anatomical coverage.\n    \"\"\"\n\n    if num_slices <= 0:\n\n        return np.array(\n            [],\n            dtype=int\n        )\n\n    # --------------------------------------------------------\n    # If fewer slices than target\n    # --------------------------------------------------------\n\n    if num_slices <= target_slices:\n\n        return np.arange(\n            num_slices,\n            dtype=int\n        )\n\n    # --------------------------------------------------------\n    # Central 80%\n    # --------------------------------------------------------\n\n    start = int(\n        round(\n            0.10 *\n            (num_slices - 1)\n        )\n    )\n\n    end = int(\n        round(\n            0.90 *\n            (num_slices - 1)\n        )\n    )\n\n    if end <= start:\n\n        return np.arange(\n            num_slices,\n            dtype=int\n        )\n\n    indices = np.linspace(\n        start,\n        end,\n        target_slices\n    )\n\n    indices = np.round(\n        indices\n    ).astype(int)\n\n    indices = np.unique(\n        indices\n    )\n\n    return indices\n\n\n# ============================================================\n# STEP 7 — LOAD ONE SERIES\n# ============================================================\n\ndef load_series_images(\n    series_path,\n    target_slices=12,\n    image_size=256\n):\n    \"\"\"\n    Load, physically sort, sample, normalize and resize\n    one DICOM series.\n    \"\"\"\n\n    sorted_records = get_sorted_dicom_files(\n        series_path\n    )\n\n    if len(sorted_records) == 0:\n\n        return np.zeros(\n            (\n                target_slices,\n                image_size,\n                image_size\n            ),\n            dtype=np.float32\n        )\n\n    # --------------------------------------------------------\n    # Sampling\n    # --------------------------------------------------------\n\n    selected_indices = sample_central_80_indices(\n        len(sorted_records),\n        target_slices\n    )\n\n    selected_records = [\n        sorted_records[i]\n        for i in selected_indices\n    ]\n\n    processed_images = []\n\n    # --------------------------------------------------------\n    # Read selected slices\n    # --------------------------------------------------------\n\n    for record in selected_records:\n\n        image, ds = read_dicom_image(\n            record[\"file\"]\n        )\n\n        if image is None:\n\n            continue\n\n        image = normalize_mri(\n            image\n        )\n\n        image = resize_image(\n            image,\n            image_size\n        )\n\n        processed_images.append(\n            image\n        )\n\n    # --------------------------------------------------------\n    # Empty result\n    # --------------------------------------------------------\n\n    if len(processed_images) == 0:\n\n        return np.zeros(\n            (\n                target_slices,\n                image_size,\n                image_size\n            ),\n            dtype=np.float32\n        )\n\n    # --------------------------------------------------------\n    # Stack\n    # --------------------------------------------------------\n\n    images = np.stack(\n        processed_images,\n        axis=0\n    ).astype(\n        np.float32\n    )\n\n    # --------------------------------------------------------\n    # Pad if necessary\n    # --------------------------------------------------------\n\n    if images.shape[0] < target_slices:\n\n        padding = np.zeros(\n            (\n                target_slices -\n                images.shape[0],\n                image_size,\n                image_size\n            ),\n            dtype=np.float32\n        )\n\n        images = np.concatenate(\n            [\n                images,\n                padding\n            ],\n            axis=0\n        )\n\n    # --------------------------------------------------------\n    # Trim if necessary\n    # --------------------------------------------------------\n\n    images = images[\n        :target_slices\n    ]\n\n    return images\n\n\n# ============================================================\n# STEP 8 — SELECT PRIMARY SERIES FOR EACH PLANE\n# ============================================================\n\ndef select_primary_series(\n    study_manifest\n):\n    \"\"\"\n    Select one primary series for each anatomical plane.\n\n    Preference:\n        1. Fluid-sensitive + Fat-suppressed\n        2. Fluid-sensitive\n        3. Fat-suppressed\n        4. Other\n\n    If multiple candidates have the same priority,\n    choose the series with the largest number of slices.\n    \"\"\"\n\n    selected = {}\n\n    for plane in PLANE_ORDER:\n\n        candidates = study_manifest[\n            study_manifest[\n                \"Anatomical_Plane\"\n            ].astype(str).str.lower()\n            ==\n            plane.lower()\n        ].copy()\n\n        if len(candidates) == 0:\n\n            continue\n\n        # ----------------------------------------------------\n        # Priority score\n        # ----------------------------------------------------\n\n        candidates[\"Priority\"] = (\n            candidates[\n                \"Fluid_Sensitive\"\n            ].astype(int) * 2\n            +\n            candidates[\n                \"Fat_Suppression\"\n            ].astype(int)\n        )\n\n        candidates = candidates.sort_values(\n            [\n                \"Priority\",\n                \"Number_of_Slices\"\n            ],\n            ascending=[\n                False,\n                False\n            ]\n        )\n\n        selected[\n            plane\n        ] = candidates.iloc[0]\n\n    return selected\n\n\n# ============================================================\n# STEP 9 — COMPLETE STUDY LOADER\n# ============================================================\n\ndef load_study(\n    study_id,\n    series_manifest,\n    target_slices=12,\n    image_size=256\n):\n    \"\"\"\n    Load one complete study.\n\n    Returns a list of dictionaries, one per available plane.\n    \"\"\"\n\n    study_manifest = (\n        series_manifest[\n            series_manifest[\n                \"StudyInstanceUID\"\n            ].astype(str)\n            ==\n            str(study_id)\n        ]\n        .copy()\n    )\n\n    if len(study_manifest) == 0:\n\n        raise ValueError(\n            f\"Study not found: {study_id}\"\n        )\n\n    # --------------------------------------------------------\n    # Select primary series\n    # --------------------------------------------------------\n\n    selected_series = select_primary_series(\n        study_manifest\n    )\n\n    loaded_study = []\n\n    # --------------------------------------------------------\n    # Load each plane\n    # --------------------------------------------------------\n\n    for plane in PLANE_ORDER:\n\n        if plane not in selected_series:\n\n            continue\n\n        row = selected_series[\n            plane\n        ]\n\n        series_path = row[\n            \"SeriesPath\"\n        ]\n\n        images = load_series_images(\n            series_path,\n            target_slices=target_slices,\n            image_size=image_size\n        )\n\n        loaded_study.append(\n            {\n                \"StudyInstanceUID\":\n                    study_id,\n\n                \"SeriesInstanceUID\":\n                    row[\n                        \"SeriesInstanceUID\"\n                    ],\n\n                \"Anatomical_Plane\":\n                    plane,\n\n                \"Fluid_Sensitive\":\n                    int(\n                        row[\n                            \"Fluid_Sensitive\"\n                        ]\n                    ),\n\n                \"Fat_Suppression\":\n                    int(\n                        row[\n                            \"Fat_Suppression\"\n                        ]\n                    ),\n\n                \"Number_of_Slices\":\n                    int(\n                        row[\n                            \"Number_of_Slices\"\n                        ]\n                    ),\n\n                \"Images\":\n                    images,\n\n                \"SeriesPath\":\n                    series_path\n            }\n        )\n\n    return loaded_study\n\n\n# ============================================================\n# STEP 10 — BUILD CORRECT SERIES PATHS\n# ============================================================\n\nseries_manifest = series_manifest.copy()\n\nseries_manifest[\n    \"SeriesPath\"\n] = (\n    TRAIN_DICOM_DIR\n    + \"/\"\n    + series_manifest[\n        \"StudyInstanceUID\"\n    ].astype(str)\n    + \"/\"\n    + series_manifest[\n        \"SeriesInstanceUID\"\n    ].astype(str)\n)\n\nprint(\"=\" * 80)\nprint(\"SERIES MANIFEST READY\")\nprint(\"=\" * 80)\n\nprint(\n    \"Studies:\",\n    series_manifest[\n        \"StudyInstanceUID\"\n    ].nunique()\n)\n\nprint(\n    \"Series:\",\n    len(series_manifest)\n)\n\nprint(\n    \"Missing paths:\",\n    series_manifest[\n        \"SeriesPath\"\n    ].isna().sum()\n)\n\n\n# ============================================================\n# STEP 11 — FIND THE LABEL DATAFRAME\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CHECKING LABEL DATA\")\nprint(\"=\" * 80)\n\n# ------------------------------------------------------------\n# The code expects a dataframe containing:\n#\n# StudyInstanceUID\n# ACL\n# MCL\n# Medial Meniscus\n# ...\n# Fracture\n#\n# If your dataframe is named labeled_studies, this works.\n# ------------------------------------------------------------\n\nif \"labeled_studies\" in globals():\n\n    label_df = labeled_studies.copy()\n\nelif \"labeled_study_df\" in globals():\n\n    label_df = labeled_study_df.copy()\n\nelif \"study_labels\" in globals():\n\n    label_df = study_labels.copy()\n\nelse:\n\n    raise NameError(\n        \"\\nNo labeled study dataframe found.\\n\"\n        \"Expected one of:\\n\"\n        \"  labeled_studies\\n\"\n        \"  labeled_study_df\\n\"\n        \"  study_labels\\n\"\n    )\n\n\n# ============================================================\n# STEP 12 — VERIFY LABEL COLUMNS\n# ============================================================\n\nmissing_labels = [\n    col\n    for col in LABEL_COLUMNS\n    if col not in label_df.columns\n]\n\nif \"StudyInstanceUID\" not in label_df.columns:\n\n    raise ValueError(\n        \"label_df must contain StudyInstanceUID\"\n    )\n\nif len(missing_labels) > 0:\n\n    raise ValueError(\n        \"Missing label columns:\\n\"\n        + \"\\n\".join(missing_labels)\n    )\n\nprint(\n    \"Labeled studies:\",\n    label_df[\n        \"StudyInstanceUID\"\n    ].nunique()\n)\n\nprint(\n    \"Label columns:\",\n    len(LABEL_COLUMNS)\n)\n\n\n# ============================================================\n# STEP 13 — CLEAN LABELS\n# ============================================================\n\nfor col in LABEL_COLUMNS:\n\n    label_df[col] = pd.to_numeric(\n        label_df[col],\n        errors=\"coerce\"\n    )\n\n    label_df[col] = (\n        label_df[col]\n        .fillna(0)\n        .astype(np.float32)\n    )\n\n    label_df[col] = np.clip(\n        label_df[col],\n        0,\n        1\n    )\n\n\n# ============================================================\n# STEP 14 — STUDY-LEVEL DATASET CLASS\n# ============================================================\n\nclass RSNAKneeStudyDataset(\n    Dataset\n):\n\n    def __init__(\n        self,\n        df,\n        series_manifest,\n        label_columns,\n        target_slices=12,\n        image_size=256\n    ):\n\n        self.df = (\n            df\n            .copy()\n            .reset_index(drop=True)\n        )\n\n        self.series_manifest = (\n            series_manifest\n        )\n\n        self.label_columns = (\n            label_columns\n        )\n\n        self.target_slices = (\n            target_slices\n        )\n\n        self.image_size = (\n            image_size\n        )\n\n    # --------------------------------------------------------\n    # Length\n    # --------------------------------------------------------\n\n    def __len__(self):\n\n        return len(\n            self.df\n        )\n\n    # --------------------------------------------------------\n    # Get one study\n    # --------------------------------------------------------\n\n    def __getitem__(\n        self,\n        idx\n    ):\n\n        # ====================================================\n        # GET STUDY\n        # ====================================================\n\n        row = self.df.iloc[\n            idx\n        ]\n\n        study_id = str(\n            row[\n                \"StudyInstanceUID\"\n            ]\n        )\n\n        # ====================================================\n        # LOAD STUDY\n        # ====================================================\n\n        loaded_study = load_study(\n            study_id,\n            self.series_manifest,\n            target_slices=\n                self.target_slices,\n            image_size=\n                self.image_size\n        )\n\n        # ====================================================\n        # INITIALIZE 3-PLANE TENSOR\n        # ====================================================\n\n        images = np.zeros(\n            (\n                3,\n                self.target_slices,\n                self.image_size,\n                self.image_size\n            ),\n            dtype=np.float32\n        )\n\n        # ====================================================\n        # PLANE MASK\n        # ====================================================\n\n        plane_mask = np.zeros(\n            3,\n            dtype=np.float32\n        )\n\n        # ====================================================\n        # INSERT PLANES\n        # ====================================================\n\n        for series in loaded_study:\n\n            plane = series[\n                \"Anatomical_Plane\"\n            ]\n\n            if plane not in PLANE_TO_INDEX:\n\n                continue\n\n            plane_idx = (\n                PLANE_TO_INDEX[\n                    plane\n                ]\n            )\n\n            series_images = np.asarray(\n                series[\"Images\"],\n                dtype=np.float32\n            )\n\n            # -----------------------------------------------\n            # Safety\n            # -----------------------------------------------\n\n            if series_images.ndim != 3:\n\n                raise ValueError(\n                    f\"Unexpected image shape \"\n                    f\"for {plane}: \"\n                    f\"{series_images.shape}\"\n                )\n\n            n = min(\n                series_images.shape[0],\n                self.target_slices\n            )\n\n            images[\n                plane_idx,\n                :n,\n                :,\n                :\n            ] = series_images[\n                :n\n            ]\n\n            plane_mask[\n                plane_idx\n            ] = 1.0\n\n        # ====================================================\n        # LABELS\n        # ====================================================\n\n        # Avoid pandas FutureWarning by converting\n        # explicitly to numeric before filling.\n        label_values = pd.to_numeric(\n            row[\n                self.label_columns\n            ],\n            errors=\"coerce\"\n        )\n\n        label_values = (\n            label_values\n            .fillna(0.0)\n            .astype(np.float32)\n        )\n\n        label_values = np.clip(\n            label_values.to_numpy(\n                dtype=np.float32\n            ),\n            0.0,\n            1.0\n        )\n\n        # ====================================================\n        # TORCH CONVERSION\n        # ====================================================\n\n        images = torch.from_numpy(\n            images\n        ).float()\n\n        plane_mask = torch.from_numpy(\n            plane_mask\n        ).float()\n\n        labels = torch.from_numpy(\n            label_values\n        ).float()\n\n        # ====================================================\n        # RETURN\n        # ====================================================\n\n        return {\n            \"images\": images,\n            \"plane_mask\": plane_mask,\n            \"labels\": labels,\n            \"study_id\": study_id\n        }\n\n\n# ============================================================\n# STEP 15 — CREATE TRAIN / VALIDATION DATA\n# ============================================================\n\n# ------------------------------------------------------------\n# If train_studies and validation_studies already exist,\n# use them.\n# Otherwise perform a study-level split.\n# ------------------------------------------------------------\n\nif (\n    \"train_studies\" in globals()\n    and\n    \"validation_studies\" in globals()\n):\n\n    train_df = train_studies.copy()\n\n    val_df = validation_studies.copy()\n\nelse:\n\n    from sklearn.model_selection import train_test_split\n\n    unique_studies = (\n        label_df[\n            \"StudyInstanceUID\"\n        ]\n        .astype(str)\n        .unique()\n    )\n\n    train_ids, val_ids = train_test_split(\n        unique_studies,\n        test_size=0.20,\n        random_state=42,\n        shuffle=True\n    )\n\n    train_df = label_df[\n        label_df[\n            \"StudyInstanceUID\"\n        ].astype(str).isin(\n            train_ids\n        )\n    ].copy()\n\n    val_df = label_df[\n        label_df[\n            \"StudyInstanceUID\"\n        ].astype(str).isin(\n            val_ids\n        )\n    ].copy()\n\n\n# ============================================================\n# STEP 16 — REMOVE STUDIES WITHOUT MRI SERIES\n# ============================================================\n\navailable_studies = set(\n    series_manifest[\n        \"StudyInstanceUID\"\n    ]\n    .astype(str)\n    .unique()\n)\n\ntrain_df = train_df[\n    train_df[\n        \"StudyInstanceUID\"\n    ].astype(str).isin(\n        available_studies\n    )\n].reset_index(\n    drop=True\n)\n\nval_df = val_df[\n    val_df[\n        \"StudyInstanceUID\"\n    ].astype(str).isin(\n        available_studies\n    )\n].reset_index(\n    drop=True\n)\n\n\n# ============================================================\n# STEP 17 — CREATE DATASETS\n# ============================================================\n\ntrain_dataset = RSNAKneeStudyDataset(\n    df=train_df,\n    series_manifest=series_manifest,\n    label_columns=LABEL_COLUMNS,\n    target_slices=TARGET_SLICES,\n    image_size=IMAGE_SIZE\n)\n\nval_dataset = RSNAKneeStudyDataset(\n    df=val_df,\n    series_manifest=series_manifest,\n    label_columns=LABEL_COLUMNS,\n    target_slices=TARGET_SLICES,\n    image_size=IMAGE_SIZE\n)\n\n\n# ============================================================\n# STEP 18 — DATASET SUMMARY\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"DATASET SUMMARY\")\nprint(\"=\" * 80)\n\nprint(\n    \"Train studies      :\",\n    len(train_dataset)\n)\n\nprint(\n    \"Validation studies :\",\n    len(val_dataset)\n)\n\nprint(\n    \"Number of labels   :\",\n    NUM_LABELS\n)\n\nprint(\n    \"Input shape        :\",\n    (\n        3,\n        TARGET_SLICES,\n        IMAGE_SIZE,\n        IMAGE_SIZE\n    )\n)\n\n\n# ============================================================\n# STEP 19 — TEST ONE TRAIN STUDY\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"MODEL-READY STUDY TEST\")\nprint(\"=\" * 80)\n\nsample = train_dataset[0]\n\nprint(\n    \"Study ID:\",\n    sample[\n        \"study_id\"\n    ]\n)\n\nprint(\n    \"Tensor shape:\",\n    tuple(\n        sample[\n            \"images\"\n        ].shape\n    )\n)\n\nprint(\n    \"Plane mask:\",\n    sample[\n        \"plane_mask\"\n    ].numpy()\n)\n\nprint(\n    \"Tensor dtype:\",\n    sample[\n        \"images\"\n    ].dtype\n)\n\nprint(\n    \"Minimum:\",\n    sample[\n        \"images\"\n    ].min().item()\n)\n\nprint(\n    \"Maximum:\",\n    sample[\n        \"images\"\n    ].max().item()\n)\n\nprint(\n    \"Labels shape:\",\n    tuple(\n        sample[\n            \"labels\"\n        ].shape\n    )\n)\n\nprint(\n    \"Labels:\",\n    sample[\n        \"labels\"\n    ].numpy()\n)\n\n\n# ============================================================\n# STEP 20 — VERIFY EXPECTED SHAPE\n# ============================================================\n\nexpected_shape = (\n    3,\n    TARGET_SLICES,\n    IMAGE_SIZE,\n    IMAGE_SIZE\n)\n\nactual_shape = tuple(\n    sample[\n        \"images\"\n    ].shape\n)\n\nif actual_shape == expected_shape:\n\n    print(\n        \"\\n✓ IMAGE TENSOR SHAPE CORRECT\"\n    )\n\nelse:\n\n    print(\n        \"\\n✗ IMAGE TENSOR SHAPE ERROR\"\n    )\n\n    print(\n        \"Expected:\",\n        expected_shape\n    )\n\n    print(\n        \"Actual:\",\n        actual_shape\n    )\n\n\n# ============================================================\n# STEP 21 — CREATE DATALOADERS\n# ============================================================\n\nBATCH_SIZE = 2\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=0,\n    pin_memory=torch.cuda.is_available()\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=0,\n    pin_memory=torch.cuda.is_available()\n)\n\n\n# ============================================================\n# STEP 22 — TEST ONE BATCH\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"DATALOADER TEST\")\nprint(\"=\" * 80)\n\nbatch = next(\n    iter(train_loader)\n)\n\nprint(\n    \"Images batch shape :\",\n    tuple(\n        batch[\n            \"images\"\n        ].shape\n    )\n)\n\nprint(\n    \"Plane mask shape   :\",\n    tuple(\n        batch[\n            \"plane_mask\"\n        ].shape\n    )\n)\n\nprint(\n    \"Labels batch shape :\",\n    tuple(\n        batch[\n            \"labels\"\n        ].shape\n    )\n)\n\nprint(\n    \"Study IDs          :\",\n    batch[\n        \"study_id\"\n    ]\n)\n\n\n# ============================================================\n# STEP 23 — FINAL VALIDATION\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"FINAL PIPELINE CHECK\")\nprint(\"=\" * 80)\n\nprint(\n    \"✓ DICOM loading\"\n)\n\nprint(\n    \"✓ Physical slice sorting\"\n)\n\nprint(\n    \"✓ Central 80% adaptive sampling\"\n)\n\nprint(\n    \"✓ MRI normalization\"\n)\n\nprint(\n    \"✓ 256 × 256 resizing\"\n)\n\nprint(\n    \"✓ Primary series selection\"\n)\n\nprint(\n    \"✓ Sagittal / Axial / Coronal arrangement\"\n)\n\nprint(\n    \"✓ Plane mask\"\n)\n\nprint(\n    \"✓ 12 abnormality labels\"\n)\n\nprint(\n    \"✓ PyTorch Dataset\"\n)\n\nprint(\n    \"✓ PyTorch DataLoader\"\n)\n\nprint(\n    \"\\nPipeline is MODEL READY.\"\n)\n\nprint(\n    \"\\nExpected study tensor:\"\n)\n\nprint(\n    \"(3, 12, 256, 256)\"\n)\n\nprint(\n    \"\\nExpected batch tensor:\"\n)\n\nprint(\n    f\"(BATCH_SIZE, 3, 12, 256, 256)\"\n)\n\nprint(\n    \"\\nPlane order:\"\n)\n\nprint(\n    \"[Sagittal, Axial, Coronal]\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:03:14.74984Z","iopub.execute_input":"2026-08-19T05:03:14.750587Z","iopub.status.idle":"2026-08-19T05:03:16.124323Z","shell.execute_reply.started":"2026-08-19T05:03:14.750557Z","shell.execute_reply":"2026-08-19T05:03:16.123629Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 60 — MULTI-PLANE MRI MODEL\n# ============================================================\n\nimport torch\nimport torch.nn as nn\nimport torchvision.models as models\n\n\n# ============================================================\n# CONFIGURATION\n# ============================================================\n\nDEVICE = torch.device(\n    \"cuda\" if torch.cuda.is_available()\n    else \"cpu\"\n)\n\nNUM_CLASSES = 12\nNUM_PLANES = 3\nNUM_SLICES = 12\n\nprint(\"=\" * 80)\nprint(\"MODEL CONFIGURATION\")\nprint(\"=\" * 80)\n\nprint(\"Device       :\", DEVICE)\nprint(\"Planes       :\", NUM_PLANES)\nprint(\"Slices/plane :\", NUM_SLICES)\nprint(\"Classes      :\", NUM_CLASSES)\n\n\n# ============================================================\n# STEP 60A — PRETRAINED RESNET18\n# ============================================================\n\nweights = models.ResNet18_Weights.DEFAULT\n\nbackbone = models.resnet18(\n    weights=weights\n)\n\n# Remove original classification layer\nfeature_dim = backbone.fc.in_features\n\nbackbone.fc = nn.Identity()\n\nprint(\"\\nBackbone feature dimension:\", feature_dim)\n\n\n# ============================================================\n# STEP 60B — MULTI-PLANE MODEL\n# ============================================================\n\nclass RSNAKneeMultiPlaneModel(\n    nn.Module\n):\n\n    def __init__(\n        self,\n        num_classes=12\n    ):\n\n        super().__init__()\n\n        # ----------------------------------------------------\n        # Shared CNN backbone\n        # ----------------------------------------------------\n\n        self.backbone = backbone\n\n        self.feature_dim = feature_dim\n\n        # ----------------------------------------------------\n        # Slice aggregation\n        # ----------------------------------------------------\n\n        self.slice_attention = nn.Sequential(\n\n            nn.Linear(\n                self.feature_dim,\n                128\n            ),\n\n            nn.ReLU(),\n\n            nn.Linear(\n                128,\n                1\n            )\n        )\n\n        # ----------------------------------------------------\n        # Plane projection\n        # ----------------------------------------------------\n\n        self.plane_projection = nn.Sequential(\n\n            nn.Linear(\n                self.feature_dim,\n                256\n            ),\n\n            nn.ReLU(),\n\n            nn.Dropout(\n                0.30\n            )\n        )\n\n        # ----------------------------------------------------\n        # Plane fusion\n        # ----------------------------------------------------\n\n        self.fusion = nn.Sequential(\n\n            nn.Linear(\n                256 * 3 + 3,\n                256\n            ),\n\n            nn.ReLU(),\n\n            nn.Dropout(\n                0.30\n            ),\n\n            nn.Linear(\n                256,\n                128\n            ),\n\n            nn.ReLU(),\n\n            nn.Dropout(\n                0.20\n            )\n        )\n\n        # ----------------------------------------------------\n        # Final 12 abnormality outputs\n        # ----------------------------------------------------\n\n        self.classifier = nn.Linear(\n            128,\n            num_classes\n        )\n\n    # ========================================================\n    # FORWARD\n    # ========================================================\n\n    def forward(\n        self,\n        x,\n        plane_mask\n    ):\n        \"\"\"\n        x:\n            [B, 3, 12, 256, 256]\n\n        plane_mask:\n            [B, 3]\n        \"\"\"\n\n        B, P, S, H, W = x.shape\n\n        # ----------------------------------------------------\n        # Reshape planes and slices into batch\n        # ----------------------------------------------------\n\n        x = x.reshape(\n            B * P * S,\n            1,\n            H,\n            W\n        )\n\n        # ----------------------------------------------------\n        # Convert grayscale → 3 channels\n        #\n        # ResNet18 expects 3-channel input.\n        # MRI intensity is repeated across RGB channels.\n        # ----------------------------------------------------\n\n        x = x.repeat(\n            1,\n            3,\n            1,\n            1\n        )\n\n        # ----------------------------------------------------\n        # CNN feature extraction\n        # ----------------------------------------------------\n\n        features = self.backbone(\n            x\n        )\n\n        # Shape:\n        # [B*P*S, 512]\n\n        features = features.reshape(\n            B,\n            P,\n            S,\n            self.feature_dim\n        )\n\n        # ----------------------------------------------------\n        # Slice attention\n        # ----------------------------------------------------\n\n        attention_scores = (\n            self.slice_attention(\n                features\n            )\n        )\n\n        # [B, P, S, 1]\n\n        attention_weights = torch.softmax(\n            attention_scores,\n            dim=2\n        )\n\n        # ----------------------------------------------------\n        # Weighted slice aggregation\n        # ----------------------------------------------------\n\n        plane_features = (\n            features *\n            attention_weights\n        ).sum(\n            dim=2\n        )\n\n        # [B, P, 512]\n\n        # ----------------------------------------------------\n        # Project each plane\n        # ----------------------------------------------------\n\n        plane_features = self.plane_projection(\n            plane_features\n        )\n\n        # [B, P, 256]\n\n        # ----------------------------------------------------\n        # Handle missing planes\n        # ----------------------------------------------------\n\n        mask = plane_mask.unsqueeze(\n            -1\n        )\n\n        plane_features = (\n            plane_features *\n            mask\n        )\n\n        # ----------------------------------------------------\n        # Flatten plane features\n        # ----------------------------------------------------\n\n        plane_features = plane_features.reshape(\n            B,\n            P * 256\n        )\n\n        # ----------------------------------------------------\n        # Add plane mask explicitly\n        # ----------------------------------------------------\n\n        fusion_input = torch.cat(\n            [\n                plane_features,\n                plane_mask\n            ],\n            dim=1\n        )\n\n        # ----------------------------------------------------\n        # Fusion\n        # ----------------------------------------------------\n\n        fused = self.fusion(\n            fusion_input\n        )\n\n        # ----------------------------------------------------\n        # Classification logits\n        # ----------------------------------------------------\n\n        logits = self.classifier(\n            fused\n        )\n\n        return logits\n\n\n# ============================================================\n# STEP 60C — CREATE MODEL\n# ============================================================\n\nmodel = RSNAKneeMultiPlaneModel(\n    num_classes=NUM_CLASSES\n)\n\nmodel = model.to(\n    DEVICE\n)\n\n\n# ============================================================\n# MODEL SUMMARY\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"MODEL CREATED\")\nprint(\"=\" * 80)\n\nprint(model)\n\n\n# ============================================================\n# PARAMETER COUNT\n# ============================================================\n\ntotal_params = sum(\n    p.numel()\n    for p in model.parameters()\n)\n\ntrainable_params = sum(\n    p.numel()\n    for p in model.parameters()\n    if p.requires_grad\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"PARAMETERS\")\nprint(\"=\" * 80)\n\nprint(\n    \"Total parameters    :\",\n    f\"{total_params:,}\"\n)\n\nprint(\n    \"Trainable parameters:\",\n    f\"{trainable_params:,}\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:04:07.82613Z","iopub.execute_input":"2026-08-19T05:04:07.826658Z","iopub.status.idle":"2026-08-19T05:04:08.017094Z","shell.execute_reply.started":"2026-08-19T05:04:07.826628Z","shell.execute_reply":"2026-08-19T05:04:08.01627Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 61 — MODEL FORWARD PASS TEST\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"STEP 61 — MODEL FORWARD PASS TEST\")\nprint(\"=\" * 80)\n\n# ------------------------------------------------------------\n# Get one batch\n# ------------------------------------------------------------\n\nbatch = next(iter(train_loader))\n\nimages = batch[\"images\"].to(DEVICE)\nplane_mask = batch[\"plane_mask\"].to(DEVICE)\nlabels = batch[\"labels\"].to(DEVICE)\n\nprint(\"\\nINPUT\")\nprint(\"-\" * 60)\n\nprint(\"Images shape      :\", tuple(images.shape))\nprint(\"Plane mask shape  :\", tuple(plane_mask.shape))\nprint(\"Labels shape      :\", tuple(labels.shape))\n\n# ------------------------------------------------------------\n# Forward pass\n# ------------------------------------------------------------\n\nmodel.eval()\n\nwith torch.no_grad():\n\n    logits = model(\n        images,\n        plane_mask\n    )\n\n# ------------------------------------------------------------\n# Check output\n# ------------------------------------------------------------\n\nprint(\"\\nOUTPUT\")\nprint(\"-\" * 60)\n\nprint(\"Logits shape      :\", tuple(logits.shape))\nprint(\"Logits dtype      :\", logits.dtype)\n\n# ------------------------------------------------------------\n# Convert logits → probabilities\n# ------------------------------------------------------------\n\nprobabilities = torch.sigmoid(logits)\n\nprint(\n    \"Probability shape :\",\n    tuple(probabilities.shape)\n)\n\n# ------------------------------------------------------------\n# Display first study predictions\n# ------------------------------------------------------------\n\nprint(\"\\nFIRST STUDY PREDICTIONS\")\nprint(\"-\" * 60)\n\nfor i, label in enumerate(LABEL_COLUMNS):\n\n    print(\n        f\"{i:2d}. \"\n        f\"{label:20s} : \"\n        f\"{probabilities[0, i].item():.4f}\"\n    )\n\n# ------------------------------------------------------------\n# Basic validation\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"FORWARD PASS CHECK\")\nprint(\"=\" * 80)\n\nexpected_shape = (\n    images.shape[0],\n    NUM_CLASSES\n)\n\nactual_shape = tuple(\n    logits.shape\n)\n\nprint(\"Expected:\", expected_shape)\nprint(\"Actual  :\", actual_shape)\n\nif actual_shape == expected_shape:\n\n    print(\"\\n✓ MODEL FORWARD PASS SUCCESSFUL\")\n    print(\"✓ Input shape is correct\")\n    print(\"✓ Output shape is correct\")\n    print(\"✓ 12 abnormality predictions generated\")\n\nelse:\n\n    print(\"\\n✗ OUTPUT SHAPE ERROR\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:04:08.018672Z","iopub.execute_input":"2026-08-19T05:04:08.019002Z","iopub.status.idle":"2026-08-19T05:04:09.151604Z","shell.execute_reply.started":"2026-08-19T05:04:08.018975Z","shell.execute_reply":"2026-08-19T05:04:09.150818Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 62 — CLASS-WEIGHTED LOSS\n# FULL INDEPENDENT VERSION\n# ============================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\n\nfrom sklearn.model_selection import train_test_split\n\n\nprint(\"=\" * 80)\nprint(\"STEP 62 — CLASS-WEIGHTED LOSS\")\nprint(\"=\" * 80)\n\n\n# ============================================================\n# 1. CONFIGURATION\n# ============================================================\n\nLABEL_COLUMNS = [\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(LABEL_COLUMNS)\n\nprint(\"\\nNumber of classes:\", NUM_CLASSES)\n\nprint(\"\\nLabels:\")\nfor i, label in enumerate(LABEL_COLUMNS):\n    print(f\"{i:2d}. {label}\")\n\n\n# ============================================================\n# 2. DATASET PATH\n# ============================================================\n\nINPUT_DIR = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n\nTRAIN_CSV = os.path.join(\n    INPUT_DIR,\n    \"train.csv\"\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"LOADING TRAIN.CSV\")\nprint(\"=\" * 80)\n\nprint(\"Path:\")\nprint(TRAIN_CSV)\n\nif not os.path.exists(TRAIN_CSV):\n\n    raise FileNotFoundError(\n        f\"train.csv not found:\\n{TRAIN_CSV}\"\n    )\n\n\n# ============================================================\n# 3. LOAD TRAIN CSV\n# ============================================================\n\ntrain_df = pd.read_csv(TRAIN_CSV)\n\nprint(\"\\nTrain shape:\")\nprint(train_df.shape)\n\nprint(\"\\n✓ train.csv loaded\")\n\n\n# ============================================================\n# 4. VERIFY LABEL COLUMNS\n# ============================================================\n\nmissing_labels = [\n    col\n    for col in LABEL_COLUMNS\n    if col not in train_df.columns\n]\n\nif len(missing_labels) > 0:\n\n    raise ValueError(\n        \"Missing label columns:\\n\"\n        + \"\\n\".join(missing_labels)\n    )\n\nprint(\"\\n✓ All 12 label columns found\")\n\n\n# ============================================================\n# 5. GET ONLY LABELED STUDIES\n# ============================================================\n\n# A study is labeled if at least one of the 12 labels\n# is not NaN.\n\nlabeled_mask = (\n    train_df[LABEL_COLUMNS]\n    .notna()\n    .any(axis=1)\n)\n\nlabeled_df = (\n    train_df[\n        labeled_mask\n    ]\n    .copy()\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"LABELED DATA\")\nprint(\"=\" * 80)\n\nprint(\n    \"Total training studies :\",\n    len(train_df)\n)\n\nprint(\n    \"Labeled studies        :\",\n    len(labeled_df)\n)\n\nprint(\n    \"Unlabeled studies      :\",\n    len(train_df) - len(labeled_df)\n)\n\n\n# ============================================================\n# 6. REMOVE INCOMPLETE LABEL ROWS\n# ============================================================\n\n# For supervised training, all 12 labels are required.\n\ncomplete_mask = (\n    labeled_df[LABEL_COLUMNS]\n    .notna()\n    .all(axis=1)\n)\n\ncomplete_labeled_df = (\n    labeled_df[\n        complete_mask\n    ]\n    .copy()\n)\n\nprint(\n    \"\\nStudies with all 12 labels:\",\n    len(complete_labeled_df)\n)\n\n\n# ============================================================\n# 7. SAFETY CHECK\n# ============================================================\n\nif len(complete_labeled_df) == 0:\n\n    raise ValueError(\n        \"No studies contain complete 12-label targets.\"\n    )\n\n\n# ============================================================\n# 8. CREATE LABEL MATRIX\n# ============================================================\n\nlabel_matrix = (\n    complete_labeled_df[\n        LABEL_COLUMNS\n    ]\n    .astype(np.float32)\n    .values\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"LABEL MATRIX\")\nprint(\"=\" * 80)\n\nprint(\n    \"Shape:\",\n    label_matrix.shape\n)\n\nprint(\n    \"Expected second dimension:\",\n    NUM_CLASSES\n)\n\n\n# ============================================================\n# 9. CREATE TRAIN / VALIDATION SPLIT\n# ============================================================\n\nall_labeled_ids = (\n    complete_labeled_df[\n        \"StudyInstanceUID\"\n    ]\n    .astype(str)\n    .tolist()\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAIN / VALIDATION SPLIT\")\nprint(\"=\" * 80)\n\n\ntrain_study_ids, validation_study_ids = train_test_split(\n    all_labeled_ids,\n    test_size=0.20,\n    random_state=42,\n    shuffle=True\n)\n\nprint(\n    \"Training studies   :\",\n    len(train_study_ids)\n)\n\nprint(\n    \"Validation studies :\",\n    len(validation_study_ids)\n)\n\n\n# ============================================================\n# 10. GET TRAINING LABEL MATRIX\n# ============================================================\n\ntrain_label_df = complete_labeled_df[\n    complete_labeled_df[\n        \"StudyInstanceUID\"\n    ]\n    .astype(str)\n    .isin(train_study_ids)\n].copy()\n\n\ntrain_labels = (\n    train_label_df[\n        LABEL_COLUMNS\n    ]\n    .astype(np.float32)\n    .values\n)\n\n\nprint(\"\\nTraining label matrix:\")\nprint(\n    \"Shape:\",\n    train_labels.shape\n)\n\n\n# ============================================================\n# 11. COUNT POSITIVE LABELS\n# ============================================================\n\npositive_counts = (\n    train_labels\n    .sum(axis=0)\n)\n\n\n# ============================================================\n# 12. COUNT NEGATIVE LABELS\n# ============================================================\n\nnegative_counts = (\n    train_labels.shape[0]\n    - positive_counts\n)\n\n\n# ============================================================\n# 13. CALCULATE POSITIVE WEIGHTS\n# ============================================================\n\npos_weights = np.ones(\n    NUM_CLASSES,\n    dtype=np.float32\n)\n\n\nfor i in range(NUM_CLASSES):\n\n    positive = positive_counts[i]\n    negative = negative_counts[i]\n\n    if positive > 0:\n\n        pos_weights[i] = (\n            negative / positive\n        )\n\n    else:\n\n        # No positive examples\n        # → keep weight at 1\n        pos_weights[i] = 1.0\n\n\n# ============================================================\n# 14. DISPLAY CLASS BALANCE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAINING CLASS BALANCE\")\nprint(\"=\" * 80)\n\nprint(\n    f\"{'No.':<5}\"\n    f\"{'Abnormality':<22}\"\n    f\"{'Positive':>10}\"\n    f\"{'Negative':>10}\"\n    f\"{'PosWeight':>12}\"\n)\n\nprint(\"-\" * 80)\n\n\nfor i, label in enumerate(LABEL_COLUMNS):\n\n    print(\n        f\"{i:<5}\"\n        f\"{label:<22}\"\n        f\"{int(positive_counts[i]):>10}\"\n        f\"{int(negative_counts[i]):>10}\"\n        f\"{pos_weights[i]:>12.3f}\"\n    )\n\n\n# ============================================================\n# 15. DEVICE\n# ============================================================\n\nif \"DEVICE\" not in globals():\n\n    DEVICE = torch.device(\n        \"cuda\"\n        if torch.cuda.is_available()\n        else \"cpu\"\n    )\n\nprint(\"\\nDevice:\")\nprint(DEVICE)\n\n\n# ============================================================\n# 16. CREATE POS_WEIGHT TENSOR\n# ============================================================\n\npos_weight_tensor = torch.tensor(\n    pos_weights,\n    dtype=torch.float32,\n    device=DEVICE\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"POSITIVE WEIGHT TENSOR\")\nprint(\"=\" * 80)\n\nprint(\n    \"Shape:\",\n    tuple(pos_weight_tensor.shape)\n)\n\nprint(\n    \"Device:\",\n    pos_weight_tensor.device\n)\n\nprint(\n    pos_weight_tensor\n)\n\n\n# ============================================================\n# 17. CREATE LOSS FUNCTION\n# ============================================================\n\ncriterion = nn.BCEWithLogitsLoss(\n    pos_weight=pos_weight_tensor\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"LOSS FUNCTION\")\nprint(\"=\" * 80)\n\nprint(criterion)\n\n\n# ============================================================\n# 18. FINAL CHECK\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 62 CHECK\")\nprint(\"=\" * 80)\n\nprint(\n    \"✓ train.csv loaded\"\n)\n\nprint(\n    \"✓ Labeled studies identified\"\n)\n\nprint(\n    \"✓ Complete 12-label studies identified\"\n)\n\nprint(\n    \"✓ Train / validation split created\"\n)\n\nprint(\n    \"✓ Positive / negative counts calculated\"\n)\n\nprint(\n    \"✓ Class weights calculated\"\n)\n\nprint(\n    \"✓ BCEWithLogitsLoss created\"\n)\n\nprint(\"\\nREADY FOR STEP 63\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:04:09.152645Z","iopub.execute_input":"2026-08-19T05:04:09.152849Z","iopub.status.idle":"2026-08-19T05:04:09.261391Z","shell.execute_reply.started":"2026-08-19T05:04:09.152828Z","shell.execute_reply":"2026-08-19T05:04:09.260702Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 63 — OPTIMIZER + LEARNING RATE SCHEDULER\n# ============================================================\n\nimport torch\nimport torch.optim as optim\n\n\nprint(\"=\" * 80)\nprint(\"STEP 63 — OPTIMIZER + LEARNING RATE SCHEDULER\")\nprint(\"=\" * 80)\n\n\n# ============================================================\n# 1. VERIFY MODEL\n# ============================================================\n\nif \"model\" not in globals():\n\n    raise NameError(\n        \"Model is not defined. Run the model creation cell first.\"\n    )\n\nprint(\"\\n✓ Model found\")\n\n\n# ============================================================\n# 2. VERIFY LOSS\n# ============================================================\n\nif \"criterion\" not in globals():\n\n    raise NameError(\n        \"criterion is not defined. Run STEP 62 first.\"\n    )\n\nprint(\"✓ Loss function found\")\n\n\n# ============================================================\n# 3. MODEL → DEVICE\n# ============================================================\n\nmodel = model.to(DEVICE)\n\nprint(\"✓ Model moved to:\", DEVICE)\n\n\n# ============================================================\n# 4. SEPARATE BACKBONE AND NEW LAYERS\n# ============================================================\n\nbackbone_parameters = []\nhead_parameters = []\n\n\nfor name, parameter in model.named_parameters():\n\n    if not parameter.requires_grad:\n        continue\n\n    if name.startswith(\"backbone.\"):\n\n        backbone_parameters.append(parameter)\n\n    else:\n\n        head_parameters.append(parameter)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"PARAMETER GROUPS\")\nprint(\"=\" * 80)\n\nprint(\n    \"Backbone parameters:\",\n    sum(p.numel() for p in backbone_parameters)\n)\n\nprint(\n    \"Head parameters    :\",\n    sum(p.numel() for p in head_parameters)\n)\n\nprint(\n    \"Total trainable    :\",\n    sum(\n        p.numel()\n        for p in model.parameters()\n        if p.requires_grad\n    )\n)\n\n\n# ============================================================\n# 5. LEARNING RATES\n# ============================================================\n\nBACKBONE_LR = 1e-5\nHEAD_LR = 3e-4\n\nWEIGHT_DECAY = 1e-4\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"LEARNING RATE CONFIGURATION\")\nprint(\"=\" * 80)\n\nprint(\n    \"Backbone LR :\",\n    BACKBONE_LR\n)\n\nprint(\n    \"Head LR     :\",\n    HEAD_LR\n)\n\nprint(\n    \"Weight decay:\",\n    WEIGHT_DECAY\n)\n\n\n# ============================================================\n# 6. ADAMW OPTIMIZER\n# ============================================================\n\noptimizer = optim.AdamW(\n    [\n        {\n            \"params\": backbone_parameters,\n            \"lr\": BACKBONE_LR\n        },\n        {\n            \"params\": head_parameters,\n            \"lr\": HEAD_LR\n        }\n    ],\n    weight_decay=WEIGHT_DECAY\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"OPTIMIZER\")\nprint(\"=\" * 80)\n\nprint(optimizer)\n\n\n# ============================================================\n# 7. COSINE ANNEALING SCHEDULER\n# ============================================================\n\nNUM_EPOCHS = 30\n\nscheduler = optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=NUM_EPOCHS,\n    eta_min=1e-6\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"SCHEDULER\")\nprint(\"=\" * 80)\n\nprint(\n    \"Scheduler:\",\n    scheduler.__class__.__name__\n)\n\nprint(\n    \"Epochs:\",\n    NUM_EPOCHS\n)\n\nprint(\n    \"Minimum LR:\",\n    1e-6\n)\n\n\n# ============================================================\n# 8. GRADIENT CLIPPING\n# ============================================================\n\nGRADIENT_CLIP = 1.0\n\nprint(\n    \"\\nGradient clipping:\",\n    GRADIENT_CLIP\n)\n\n\n# ============================================================\n# 9. CHECK INITIAL LEARNING RATES\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"INITIAL LEARNING RATES\")\nprint(\"=\" * 80)\n\nfor i, group in enumerate(optimizer.param_groups):\n\n    print(\n        f\"Group {i}: \"\n        f\"{group['lr']:.8f}\"\n    )\n\n\n# ============================================================\n# 10. FINAL CHECK\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 63 CHECK\")\nprint(\"=\" * 80)\n\nprint(\"✓ AdamW optimizer created\")\nprint(\"✓ Backbone uses low learning rate\")\nprint(\"✓ Classification/fusion layers use higher learning rate\")\nprint(\"✓ Weight decay enabled\")\nprint(\"✓ Cosine learning-rate scheduler created\")\nprint(\"✓ Gradient clipping configured\")\nprint(\"✓ Training epochs:\", NUM_EPOCHS)\n\nprint(\"\\nREADY FOR STEP 64\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:04:09.263127Z","iopub.execute_input":"2026-08-19T05:04:09.26382Z","iopub.status.idle":"2026-08-19T05:04:09.280038Z","shell.execute_reply.started":"2026-08-19T05:04:09.263795Z","shell.execute_reply":"2026-08-19T05:04:09.279449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 64 — TRAINING + VALIDATION\n# ============================================================\n\nimport os\nimport copy\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\n\nfrom sklearn.metrics import (\n    roc_auc_score,\n    f1_score,\n    precision_score,\n    recall_score,\n    accuracy_score\n)\n\n\nprint(\"=\" * 80)\nprint(\"STEP 64 — TRAINING + VALIDATION\")\nprint(\"=\" * 80)\n\n\n# ============================================================\n# 1. VERIFY REQUIRED OBJECTS\n# ============================================================\n\nrequired_objects = [\n    \"model\",\n    \"train_loader\",\n    \"val_loader\",\n    \"criterion\",\n    \"optimizer\",\n    \"scheduler\",\n    \"DEVICE\",\n    \"LABEL_COLUMNS\"\n]\n\nfor name in required_objects:\n\n    if name not in globals():\n\n        raise NameError(\n            f\"Required object '{name}' is missing.\"\n        )\n\n    print(f\"✓ {name}\")\n\n\n# ============================================================\n# 2. TRAINING CONFIGURATION\n# ============================================================\n\nNUM_EPOCHS = 30\n\nEARLY_STOPPING_PATIENCE = 7\n\nGRADIENT_CLIP = 1.0\n\nCHECKPOINT_PATH = (\n    \"/kaggle/working/\"\n    \"best_rsna_knee_model.pth\"\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAINING CONFIGURATION\")\nprint(\"=\" * 80)\n\nprint(\"Epochs              :\", NUM_EPOCHS)\nprint(\"Early stopping      :\", EARLY_STOPPING_PATIENCE)\nprint(\"Gradient clipping   :\", GRADIENT_CLIP)\nprint(\"Checkpoint          :\", CHECKPOINT_PATH)\n\n\n# ============================================================\n# 3. METRIC FUNCTION\n# ============================================================\n\ndef calculate_metrics(\n    targets,\n    probabilities\n):\n\n    targets = np.asarray(targets)\n    probabilities = np.asarray(probabilities)\n\n    predictions = (\n        probabilities >= 0.5\n    ).astype(int)\n\n    auc_scores = []\n    f1_scores = []\n    precision_scores = []\n    recall_scores = []\n\n    for i in range(NUM_CLASSES):\n\n        y_true = targets[:, i]\n        y_prob = probabilities[:, i]\n        y_pred = predictions[:, i]\n\n        # ----------------------------------------------------\n        # AUROC\n        # ----------------------------------------------------\n\n        if len(np.unique(y_true)) >= 2:\n\n            try:\n\n                auc = roc_auc_score(\n                    y_true,\n                    y_prob\n                )\n\n            except Exception:\n\n                auc = np.nan\n\n        else:\n\n            auc = np.nan\n\n\n        # ----------------------------------------------------\n        # F1\n        # ----------------------------------------------------\n\n        f1 = f1_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        )\n\n\n        # ----------------------------------------------------\n        # Precision\n        # ----------------------------------------------------\n\n        precision = precision_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        )\n\n\n        # ----------------------------------------------------\n        # Recall\n        # ----------------------------------------------------\n\n        recall = recall_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        )\n\n\n        auc_scores.append(auc)\n        f1_scores.append(f1)\n        precision_scores.append(precision)\n        recall_scores.append(recall)\n\n\n    valid_auc = [\n        x for x in auc_scores\n        if not np.isnan(x)\n    ]\n\n\n    macro_auc = (\n        np.mean(valid_auc)\n        if len(valid_auc) > 0\n        else np.nan\n    )\n\n    macro_f1 = np.mean(f1_scores)\n\n    macro_precision = np.mean(\n        precision_scores\n    )\n\n    macro_recall = np.mean(\n        recall_scores\n    )\n\n\n    return {\n        \"macro_auc\": macro_auc,\n        \"macro_f1\": macro_f1,\n        \"macro_precision\": macro_precision,\n        \"macro_recall\": macro_recall,\n        \"auc_per_class\": auc_scores,\n        \"f1_per_class\": f1_scores,\n        \"precision_per_class\": precision_scores,\n        \"recall_per_class\": recall_scores\n    }\n\n\n# ============================================================\n# 4. ONE TRAINING EPOCH\n# ============================================================\n\ndef train_one_epoch(\n    model,\n    loader,\n    criterion,\n    optimizer,\n    device\n):\n\n    model.train()\n\n    running_loss = 0.0\n\n    all_targets = []\n    all_probabilities = []\n\n    num_batches = 0\n\n\n    for batch in loader:\n\n        # ----------------------------------------------------\n        # Load batch\n        # ----------------------------------------------------\n\n        images = batch[\"images\"].to(\n            device,\n            non_blocking=True\n        )\n\n        plane_mask = batch[\"plane_mask\"].to(\n            device,\n            non_blocking=True\n        )\n\n        labels = batch[\"labels\"].to(\n            device,\n            non_blocking=True\n        )\n\n\n        # ----------------------------------------------------\n        # Clear gradients\n        # ----------------------------------------------------\n\n        optimizer.zero_grad(\n            set_to_none=True\n        )\n\n\n        # ----------------------------------------------------\n        # Forward\n        # ----------------------------------------------------\n\n        logits = model(\n            images,\n            plane_mask\n        )\n\n\n        # ----------------------------------------------------\n        # Loss\n        # ----------------------------------------------------\n\n        loss = criterion(\n            logits,\n            labels\n        )\n\n\n        # ----------------------------------------------------\n        # Backward\n        # ----------------------------------------------------\n\n        loss.backward()\n\n\n        # ----------------------------------------------------\n        # Gradient clipping\n        # ----------------------------------------------------\n\n        torch.nn.utils.clip_grad_norm_(\n            model.parameters(),\n            GRADIENT_CLIP\n        )\n\n\n        # ----------------------------------------------------\n        # Optimizer\n        # ----------------------------------------------------\n\n        optimizer.step()\n\n\n        # ----------------------------------------------------\n        # Statistics\n        # ----------------------------------------------------\n\n        running_loss += loss.item()\n\n        num_batches += 1\n\n\n        # ----------------------------------------------------\n        # Probabilities\n        # ----------------------------------------------------\n\n        probabilities = torch.sigmoid(\n            logits\n        )\n\n\n        all_targets.append(\n            labels.detach()\n            .cpu()\n            .numpy()\n        )\n\n        all_probabilities.append(\n            probabilities.detach()\n            .cpu()\n            .numpy()\n        )\n\n\n    # ========================================================\n    # Combine predictions\n    # ========================================================\n\n    all_targets = np.concatenate(\n        all_targets,\n        axis=0\n    )\n\n    all_probabilities = np.concatenate(\n        all_probabilities,\n        axis=0\n    )\n\n\n    average_loss = (\n        running_loss /\n        max(num_batches, 1)\n    )\n\n\n    metrics = calculate_metrics(\n        all_targets,\n        all_probabilities\n    )\n\n\n    return (\n        average_loss,\n        metrics\n    )\n\n\n# ============================================================\n# 5. VALIDATION\n# ============================================================\n\ndef validate_one_epoch(\n    model,\n    loader,\n    criterion,\n    device\n):\n\n    model.eval()\n\n    running_loss = 0.0\n\n    all_targets = []\n    all_probabilities = []\n\n    num_batches = 0\n\n\n    with torch.no_grad():\n\n        for batch in loader:\n\n            # ------------------------------------------------\n            # Load batch\n            # ------------------------------------------------\n\n            images = batch[\"images\"].to(\n                device,\n                non_blocking=True\n            )\n\n            plane_mask = batch[\"plane_mask\"].to(\n                device,\n                non_blocking=True\n            )\n\n            labels = batch[\"labels\"].to(\n                device,\n                non_blocking=True\n            )\n\n\n            # ------------------------------------------------\n            # Forward\n            # ------------------------------------------------\n\n            logits = model(\n                images,\n                plane_mask\n            )\n\n\n            # ------------------------------------------------\n            # Loss\n            # ------------------------------------------------\n\n            loss = criterion(\n                logits,\n                labels\n            )\n\n\n            running_loss += loss.item()\n\n            num_batches += 1\n\n\n            # ------------------------------------------------\n            # Probability\n            # ------------------------------------------------\n\n            probabilities = torch.sigmoid(\n                logits\n            )\n\n\n            all_targets.append(\n                labels.cpu().numpy()\n            )\n\n            all_probabilities.append(\n                probabilities.cpu().numpy()\n            )\n\n\n    # ========================================================\n    # Combine\n    # ========================================================\n\n    all_targets = np.concatenate(\n        all_targets,\n        axis=0\n    )\n\n    all_probabilities = np.concatenate(\n        all_probabilities,\n        axis=0\n    )\n\n\n    average_loss = (\n        running_loss /\n        max(num_batches, 1)\n    )\n\n\n    metrics = calculate_metrics(\n        all_targets,\n        all_probabilities\n    )\n\n\n    return (\n        average_loss,\n        metrics,\n        all_targets,\n        all_probabilities\n    )\n\n\n# ============================================================\n# 6. TRAINING HISTORY\n# ============================================================\n\nhistory = {\n\n    \"epoch\": [],\n\n    \"train_loss\": [],\n    \"val_loss\": [],\n\n    \"train_auc\": [],\n    \"val_auc\": [],\n\n    \"train_f1\": [],\n    \"val_f1\": [],\n\n    \"train_precision\": [],\n    \"val_precision\": [],\n\n    \"train_recall\": [],\n    \"val_recall\": [],\n\n    \"learning_rate_backbone\": [],\n    \"learning_rate_head\": []\n}\n\n\n# ============================================================\n# 7. BEST MODEL TRACKING\n# ============================================================\n\nbest_val_auc = -np.inf\n\nbest_epoch = 0\n\nepochs_without_improvement = 0\n\nbest_state = None\n\n\n# ============================================================\n# 8. TRAINING LOOP\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STARTING TRAINING\")\nprint(\"=\" * 80)\n\n\nfor epoch in range(1, NUM_EPOCHS + 1):\n\n\n    # ========================================================\n    # TRAIN\n    # ========================================================\n\n    train_loss, train_metrics = train_one_epoch(\n        model,\n        train_loader,\n        criterion,\n        optimizer,\n        DEVICE\n    )\n\n\n    # ========================================================\n    # VALIDATION\n    # ========================================================\n\n    (\n        val_loss,\n        val_metrics,\n        val_targets,\n        val_probabilities\n    ) = validate_one_epoch(\n        model,\n        val_loader,\n        criterion,\n        DEVICE\n    )\n\n\n    # ========================================================\n    # CURRENT LR\n    # ========================================================\n\n    current_lr_backbone = (\n        optimizer.param_groups[0][\"lr\"]\n    )\n\n    current_lr_head = (\n        optimizer.param_groups[1][\"lr\"]\n    )\n\n\n    # ========================================================\n    # SAVE HISTORY\n    # ========================================================\n\n    history[\"epoch\"].append(epoch)\n\n    history[\"train_loss\"].append(\n        train_loss\n    )\n\n    history[\"val_loss\"].append(\n        val_loss\n    )\n\n    history[\"train_auc\"].append(\n        train_metrics[\"macro_auc\"]\n    )\n\n    history[\"val_auc\"].append(\n        val_metrics[\"macro_auc\"]\n    )\n\n    history[\"train_f1\"].append(\n        train_metrics[\"macro_f1\"]\n    )\n\n    history[\"val_f1\"].append(\n        val_metrics[\"macro_f1\"]\n    )\n\n    history[\"train_precision\"].append(\n        train_metrics[\"macro_precision\"]\n    )\n\n    history[\"val_precision\"].append(\n        val_metrics[\"macro_precision\"]\n    )\n\n    history[\"train_recall\"].append(\n        train_metrics[\"macro_recall\"]\n    )\n\n    history[\"val_recall\"].append(\n        val_metrics[\"macro_recall\"]\n    )\n\n    history[\n        \"learning_rate_backbone\"\n    ].append(\n        current_lr_backbone\n    )\n\n    history[\n        \"learning_rate_head\"\n    ].append(\n        current_lr_head\n    )\n\n\n    # ========================================================\n    # PRINT EPOCH\n    # ========================================================\n\n    print(\n        f\"\\nEpoch {epoch:02d}/{NUM_EPOCHS}\"\n    )\n\n    print(\n        \"-\" * 80\n    )\n\n    print(\n        f\"Train Loss : {train_loss:.4f}\"\n    )\n\n    print(\n        f\"Val Loss   : {val_loss:.4f}\"\n    )\n\n    print(\n        f\"Train AUC  : {train_metrics['macro_auc']:.4f}\"\n    )\n\n    print(\n        f\"Val AUC    : {val_metrics['macro_auc']:.4f}\"\n    )\n\n    print(\n        f\"Train F1   : {train_metrics['macro_f1']:.4f}\"\n    )\n\n    print(\n        f\"Val F1     : {val_metrics['macro_f1']:.4f}\"\n    )\n\n    print(\n        f\"Val Recall : {val_metrics['macro_recall']:.4f}\"\n    )\n\n    print(\n        f\"LR Backbone: {current_lr_backbone:.8f}\"\n    )\n\n    print(\n        f\"LR Head    : {current_lr_head:.8f}\"\n    )\n\n\n    # ========================================================\n    # BEST MODEL\n    # ========================================================\n\n    current_val_auc = val_metrics[\"macro_auc\"]\n\n\n    if not np.isnan(current_val_auc):\n\n        if current_val_auc > best_val_auc:\n\n            best_val_auc = current_val_auc\n\n            best_epoch = epoch\n\n            epochs_without_improvement = 0\n\n            best_state = copy.deepcopy(\n                model.state_dict()\n            )\n\n\n            # ------------------------------------------------\n            # Save checkpoint\n            # ------------------------------------------------\n\n            torch.save(\n                {\n                    \"epoch\": epoch,\n\n                    \"model_state_dict\":\n                        model.state_dict(),\n\n                    \"optimizer_state_dict\":\n                        optimizer.state_dict(),\n\n                    \"scheduler_state_dict\":\n                        scheduler.state_dict(),\n\n                    \"best_val_auc\":\n                        best_val_auc,\n\n                    \"label_columns\":\n                        LABEL_COLUMNS\n                },\n                CHECKPOINT_PATH\n            )\n\n\n            print(\n                \"✓ NEW BEST MODEL SAVED\"\n            )\n\n        else:\n\n            epochs_without_improvement += 1\n\n\n    # ========================================================\n    # LR SCHEDULER\n    # ========================================================\n\n    scheduler.step()\n\n\n    # ========================================================\n    # EARLY STOPPING\n    # ========================================================\n\n    if (\n        epochs_without_improvement\n        >= EARLY_STOPPING_PATIENCE\n    ):\n\n        print(\n            \"\\nEarly stopping triggered.\"\n        )\n\n        print(\n            \"Best epoch:\",\n            best_epoch\n        )\n\n        break\n\n\n# ============================================================\n# 9. RESTORE BEST MODEL\n# ============================================================\n\nif best_state is not None:\n\n    model.load_state_dict(\n        best_state\n    )\n\n    print(\n        \"\\n✓ Best model weights restored\"\n    )\n\n\n# ============================================================\n# 10. SAVE HISTORY\n# ============================================================\n\nhistory_df = pd.DataFrame(\n    history\n)\n\nhistory_path = (\n    \"/kaggle/working/\"\n    \"rsna_knee_training_history.csv\"\n)\n\nhistory_df.to_csv(\n    history_path,\n    index=False\n)\n\n\n# ============================================================\n# 11. FINAL SUMMARY\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 64 COMPLETE\")\nprint(\"=\" * 80)\n\nprint(\n    \"Best epoch        :\",\n    best_epoch\n)\n\nprint(\n    \"Best validation AUC:\",\n    f\"{best_val_auc:.4f}\"\n)\n\nprint(\n    \"Checkpoint:\",\n    CHECKPOINT_PATH\n)\n\nprint(\n    \"History:\",\n    history_path\n)\n\nprint(\n    \"\\n✓ Training completed\"\n)\n\nprint(\n    \"✓ Validation completed\"\n)\n\nprint(\n    \"✓ Best model selected\"\n)\n\nprint(\n    \"✓ Training history saved\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:40:50.285698Z","iopub.execute_input":"2026-08-19T05:40:50.286161Z","iopub.status.idle":"2026-08-19T05:48:01.918176Z","shell.execute_reply.started":"2026-08-19T05:40:50.286127Z","shell.execute_reply":"2026-08-19T05:48:01.917426Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 65A — LOAD OFFICIAL LABELED REPORTS\n# ============================================================\n\nimport pandas as pd\nimport numpy as np\nimport re\n\nprint(\"=\" * 80)\nprint(\"STEP 65A — OFFICIAL LABELED REPORTS\")\nprint(\"=\" * 80)\n\nINPUT_DIR = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n\nTRAIN_CSV = f\"{INPUT_DIR}/train.csv\"\n\nLABEL_COLUMNS = [\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\n# Load train.csv\ntrain_df = pd.read_csv(TRAIN_CSV)\n\n# Keep studies with complete official labels\nlabeled_df = train_df.dropna(\n    subset=LABEL_COLUMNS\n).copy()\n\nprint(\"Total studies       :\", len(train_df))\nprint(\"Labeled studies     :\", len(labeled_df))\nprint(\"Unlabeled studies   :\", len(train_df) - len(labeled_df))\n\nprint(\"\\nReports available   :\", labeled_df[\"Report\"].notna().sum())\nprint(\n    \"Empty reports       :\",\n    labeled_df[\"Report\"].fillna(\"\").str.strip().eq(\"\").sum()\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"LABEL COUNTS\")\nprint(\"=\" * 80)\n\nfor label in LABEL_COLUMNS:\n    print(\n        f\"{label:20s} : \"\n        f\"{int(labeled_df[label].sum())} positive / \"\n        f\"{len(labeled_df)} total\"\n    )\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"SAMPLE REPORTS\")\nprint(\"=\" * 80)\n\ndisplay(\n    labeled_df[\n        [\"StudyInstanceUID\", \"Report\"] + LABEL_COLUMNS\n    ].head(5)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:01.91989Z","iopub.execute_input":"2026-08-19T05:48:01.920466Z","iopub.status.idle":"2026-08-19T05:48:02.030038Z","shell.execute_reply.started":"2026-08-19T05:48:01.920437Z","shell.execute_reply":"2026-08-19T05:48:02.029194Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 65B — REPORT NORMALIZATION\n# ============================================================\n\ndef normalize_report(text):\n\n    if pd.isna(text):\n        return \"\"\n\n    text = str(text)\n\n    # Lowercase\n    text = text.lower()\n\n    # Normalize line breaks\n    text = re.sub(r\"[\\r\\n\\t]+\", \" \", text)\n\n    # Remove repeated spaces\n    text = re.sub(r\"\\s+\", \" \", text)\n\n    return text.strip()\n\n\nlabeled_df[\"Report_Normalized\"] = (\n    labeled_df[\"Report\"]\n    .apply(normalize_report)\n)\n\nprint(\"=\" * 80)\nprint(\"REPORT NORMALIZATION\")\nprint(\"=\" * 80)\n\nprint(\n    \"Reports processed:\",\n    len(labeled_df)\n)\n\nprint(\n    \"Empty normalized reports:\",\n    labeled_df[\"Report_Normalized\"]\n    .eq(\"\")\n    .sum()\n)\n\nprint(\"\\nExample:\")\nprint(labeled_df[\"Report_Normalized\"].iloc[0][:1000])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:02.031158Z","iopub.execute_input":"2026-08-19T05:48:02.031695Z","iopub.status.idle":"2026-08-19T05:48:02.046969Z","shell.execute_reply.started":"2026-08-19T05:48:02.031668Z","shell.execute_reply":"2026-08-19T05:48:02.046142Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 65C — INITIAL REPORT TERMINOLOGY\n# ============================================================\n\nREPORT_TERMS = {\n\n    \"ACL\": [\n        \"acl tear\",\n        \"acl rupture\",\n        \"acl injury\",\n        \"acl insufficiency\",\n        \"acl deficient\",\n        \"anterior cruciate ligament tear\",\n        \"anterior cruciate ligament rupture\",\n        \"anterior cruciate ligament injury\"\n    ],\n\n    \"MCL\": [\n        \"mcl tear\",\n        \"mcl rupture\",\n        \"mcl injury\",\n        \"medial collateral ligament tear\",\n        \"medial collateral ligament rupture\",\n        \"medial collateral ligament injury\"\n    ],\n\n    \"Medial Meniscus\": [\n        \"medial meniscus tear\",\n        \"medial meniscal tear\",\n        \"medial meniscus rupture\",\n        \"medial meniscus lesion\",\n        \"medial meniscus injury\",\n        \"tear of the medial meniscus\"\n    ],\n\n    \"Lateral Meniscus\": [\n        \"lateral meniscus tear\",\n        \"lateral meniscal tear\",\n        \"lateral meniscus rupture\",\n        \"lateral meniscus lesion\",\n        \"lateral meniscus injury\",\n        \"tear of the lateral meniscus\"\n    ],\n\n    \"Medial OA\": [\n        \"medial compartment osteoarthritis\",\n        \"medial compartment oa\",\n        \"medial femorotibial osteoarthritis\",\n        \"medial compartment degenerative\",\n        \"medial joint space narrowing\",\n        \"medial cartilage loss\"\n    ],\n\n    \"Lateral OA\": [\n        \"lateral compartment osteoarthritis\",\n        \"lateral compartment oa\",\n        \"lateral femorotibial osteoarthritis\",\n        \"lateral compartment degenerative\",\n        \"lateral joint space narrowing\",\n        \"lateral cartilage loss\"\n    ],\n\n    \"PF OA\": [\n        \"patellofemoral osteoarthritis\",\n        \"patellofemoral oa\",\n        \"patellofemoral degenerative\",\n        \"patellofemoral cartilage loss\",\n        \"patellofemoral joint space narrowing\",\n        \"patellar osteoarthritis\"\n    ],\n\n    \"Effusion\": [\n        \"joint effusion\",\n        \"knee effusion\",\n        \"joint fluid\",\n        \"increased joint fluid\",\n        \"increased amount of fluid\",\n        \"joint recess effusion\"\n    ],\n\n    \"Synovitis\": [\n        \"synovitis\",\n        \"synovial inflammation\",\n        \"synovial thickening\",\n        \"synovial hypertrophy\"\n    ],\n\n    \"Baker's\": [\n        \"baker cyst\",\n        \"baker's cyst\",\n        \"baker cystic\",\n        \"popliteal cyst\",\n        \"popliteal cystic lesion\"\n    ],\n\n    \"Contusion\": [\n        \"bone contusion\",\n        \"bone bruise\",\n        \"bone marrow contusion\",\n        \"bone marrow edema\",\n        \"bone marrow oedema\",\n        \"osseous contusion\"\n    ],\n\n    \"Fracture\": [\n        \"fracture\",\n        \"fractured\",\n        \"fracture line\",\n        \"bone fracture\"\n    ]\n}\n\nprint(\"=\" * 80)\nprint(\"REPORT TERMINOLOGY CREATED\")\nprint(\"=\" * 80)\n\nfor label, terms in REPORT_TERMS.items():\n    print(f\"{label:20s}: {len(terms)} terms\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:02.047993Z","iopub.execute_input":"2026-08-19T05:48:02.048311Z","iopub.status.idle":"2026-08-19T05:48:02.067621Z","shell.execute_reply.started":"2026-08-19T05:48:02.048277Z","shell.execute_reply":"2026-08-19T05:48:02.066594Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 65D — GENERATE PRELIMINARY REPORT LABELS\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"STEP 65D — PRELIMINARY REPORT LABEL GENERATION\")\nprint(\"=\" * 80)\n\n\ndef keyword_positive(report, terms):\n\n    if not isinstance(report, str):\n        return 0\n\n    report = report.lower()\n\n    for term in terms:\n\n        if term.lower() in report:\n            return 1\n\n    return 0\n\n\n# ------------------------------------------------------------\n# Generate report-based prediction for each abnormality\n# ------------------------------------------------------------\n\nfor label in LABEL_COLUMNS:\n\n    prediction_column = f\"Report_{label}\"\n\n    labeled_df[prediction_column] = (\n        labeled_df[\"Report_Normalized\"]\n        .apply(\n            lambda report:\n                keyword_positive(\n                    report,\n                    REPORT_TERMS[label]\n                )\n        )\n    )\n\n\n# ------------------------------------------------------------\n# Display comparison\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"REPORT VS OFFICIAL POSITIVE COUNTS\")\nprint(\"=\" * 80)\n\nfor label in LABEL_COLUMNS:\n\n    official_count = int(\n        labeled_df[label].sum()\n    )\n\n    report_count = int(\n        labeled_df[f\"Report_{label}\"].sum()\n    )\n\n    print(\n        f\"{label:20s} | \"\n        f\"Official = {official_count:2d} | \"\n        f\"Report = {report_count:2d}\"\n    )\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"PRELIMINARY REPORT LABEL MATRIX\")\nprint(\"=\" * 80)\n\nreport_prediction_columns = [\n    f\"Report_{label}\"\n    for label in LABEL_COLUMNS\n]\n\ndisplay(\n    labeled_df[\n        [\"StudyInstanceUID\"]\n        + report_prediction_columns\n    ].head(20)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:02.069556Z","iopub.execute_input":"2026-08-19T05:48:02.069838Z","iopub.status.idle":"2026-08-19T05:48:02.106943Z","shell.execute_reply.started":"2026-08-19T05:48:02.069816Z","shell.execute_reply":"2026-08-19T05:48:02.106335Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 65E — REPORT LABEL PERFORMANCE\n# ============================================================\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    precision_score,\n    recall_score,\n    f1_score\n)\n\nprint(\"=\" * 80)\nprint(\"STEP 65E — REPORT VS OFFICIAL LABEL PERFORMANCE\")\nprint(\"=\" * 80)\n\n\nresults = []\n\n\nfor label in LABEL_COLUMNS:\n\n    # Official ground truth\n    y_true = (\n        labeled_df[label]\n        .astype(int)\n        .values\n    )\n\n    # Report-derived prediction\n    y_pred = (\n        labeled_df[f\"Report_{label}\"]\n        .astype(int)\n        .values\n    )\n\n    results.append({\n\n        \"Abnormality\": label,\n\n        \"Official_Positive\": int(\n            y_true.sum()\n        ),\n\n        \"Report_Positive\": int(\n            y_pred.sum()\n        ),\n\n        \"Accuracy\": accuracy_score(\n            y_true,\n            y_pred\n        ),\n\n        \"Precision\": precision_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        ),\n\n        \"Recall\": recall_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        ),\n\n        \"F1\": f1_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        )\n    })\n\n\nreport_eval = pd.DataFrame(results)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"FINAL REPORT LABEL PERFORMANCE\")\nprint(\"=\" * 80)\n\ndisplay(\n    report_eval.round(4)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:02.107844Z","iopub.execute_input":"2026-08-19T05:48:02.108065Z","iopub.status.idle":"2026-08-19T05:48:02.210581Z","shell.execute_reply.started":"2026-08-19T05:48:02.108045Z","shell.execute_reply":"2026-08-19T05:48:02.209967Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 66A — MULTILINGUAL + NEGATION-AWARE REPORT PARSER\n# ============================================================\n\nimport re\nimport pandas as pd\nimport numpy as np\n\nprint(\"=\" * 80)\nprint(\"STEP 66A — MULTILINGUAL REPORT PARSER\")\nprint(\"=\" * 80)\n\n\n# ------------------------------------------------------------\n# Helper\n# ------------------------------------------------------------\n\ndef clean_report(text):\n\n    if pd.isna(text):\n        return \"\"\n\n    text = str(text).lower()\n\n    # Normalize line breaks\n    text = re.sub(r\"[\\r\\n\\t]+\", \" \", text)\n\n    # Normalize punctuation\n    text = re.sub(r\"[,;:()\\[\\]{}]+\", \" \", text)\n\n    # Remove repeated spaces\n    text = re.sub(r\"\\s+\", \" \", text)\n\n    return text.strip()\n\n\nlabeled_df[\"Report_Clean\"] = (\n    labeled_df[\"Report\"]\n    .apply(clean_report)\n)\n\n\n# ------------------------------------------------------------\n# Negative / normal terminology\n# ------------------------------------------------------------\n\nNEGATION_TERMS = [\n\n    \"no\",\n    \"not\",\n    \"without\",\n    \"normal\",\n    \"intact\",\n    \"unremarkable\",\n    \"no evidence\",\n    \"negative for\",\n    \"absence of\",\n    \"without evidence\",\n    \"no signs\",\n    \"no sign\",\n    \"no abnormality\",\n    \"no significant abnormality\",\n\n    # Spanish\n    \"sin\",\n    \"no hay\",\n    \"normal\",\n    \"sin signos\",\n    \"sin evidencia\",\n    \"sin alteraciones\",\n\n    # German\n    \"kein\",\n    \"keine\",\n    \"keinen\",\n    \"keiner\",\n    \"ohne\",\n    \"unauffällig\",\n    \"intakt\",\n\n    # Turkish\n    \"yok\",\n    \"izlenmedi\",\n    \"saptanmadı\",\n    \"normal\",\n    \"intakt\",\n\n    # Croatian / related\n    \"bez\",\n    \"nema\",\n    \"uredan\",\n    \"uredno\",\n\n    # Greek\n    \"χωρίς\",\n    \"δεν\",\n    \"φυσιολογικό\",\n\n    # Bulgarian\n    \"няма\",\n    \"без\",\n    \"нормален\"\n]\n\n\n# ------------------------------------------------------------\n# Positive terminology\n# ------------------------------------------------------------\n\nMULTILINGUAL_TERMS = {\n\n    \"ACL\": [\n\n        # English\n        \"acl tear\",\n        \"acl rupture\",\n        \"acl injury\",\n        \"acl insufficiency\",\n        \"acl deficient\",\n        \"acl disruption\",\n        \"anterior cruciate ligament tear\",\n        \"anterior cruciate ligament rupture\",\n        \"anterior cruciate ligament injury\",\n\n        # Spanish\n        \"rotura del ligamento cruzado anterior\",\n        \"lesión del ligamento cruzado anterior\",\n\n        # German\n        \"vkb rupture\",\n        \"vkb-riss\",\n        \"vorderes kreuzband rupture\",\n        \"vorderes kreuzbandriss\",\n\n        # Turkish\n        \"ön çapraz bağ yırtığı\",\n        \"ön çapraz bağ rüptürü\"\n    ],\n\n    \"MCL\": [\n\n        \"mcl tear\",\n        \"mcl rupture\",\n        \"mcl injury\",\n        \"medial collateral ligament tear\",\n        \"medial collateral ligament rupture\",\n        \"medial collateral ligament injury\",\n\n        \"rotura del ligamento colateral medial\",\n        \"lesión del ligamento colateral medial\",\n\n        \"mediales kollateralband rupture\",\n\n        \"medial kollateral bağ yırtığı\",\n        \"medial kollateral bağ rüptürü\"\n    ],\n\n    \"Medial Meniscus\": [\n\n        \"medial meniscus tear\",\n        \"medial meniscal tear\",\n        \"medial meniscus rupture\",\n        \"medial meniscus lesion\",\n        \"tear of the medial meniscus\",\n\n        \"rotura del menisco medial\",\n        \"lesión del menisco medial\",\n\n        \"medialer meniskusriss\",\n\n        \"medial menisküs yırtığı\",\n        \"medial menisküs rüptürü\"\n    ],\n\n    \"Lateral Meniscus\": [\n\n        \"lateral meniscus tear\",\n        \"lateral meniscal tear\",\n        \"lateral meniscus rupture\",\n        \"lateral meniscus lesion\",\n        \"tear of the lateral meniscus\",\n\n        \"rotura del menisco lateral\",\n        \"lesión del menisco lateral\",\n\n        \"lateraler meniskusriss\",\n\n        \"lateral menisküs yırtığı\",\n        \"lateral menisküs rüptürü\"\n    ],\n\n    \"Medial OA\": [\n\n        \"medial compartment osteoarthritis\",\n        \"medial compartment oa\",\n        \"medial femorotibial osteoarthritis\",\n        \"medial compartment degenerative\",\n        \"medial joint space narrowing\",\n        \"medial cartilage loss\",\n\n        \"osteoarthritis of the medial compartment\",\n        \"medial compartment arthrosis\",\n        \"medial compartment chondrosis\",\n\n        \"artrosis del compartimento medial\",\n        \"artrosis femorotibial medial\",\n        \"condropatía medial\",\n\n        \"mediale gonarthrose\",\n        \"mediales kompartiment arthrose\",\n\n        \"medial kompartman osteoartriti\",\n        \"medial kompartman artrozu\"\n    ],\n\n    \"Lateral OA\": [\n\n        \"lateral compartment osteoarthritis\",\n        \"lateral compartment oa\",\n        \"lateral femorotibial osteoarthritis\",\n        \"lateral compartment degenerative\",\n        \"lateral joint space narrowing\",\n        \"lateral cartilage loss\",\n\n        \"osteoarthritis of the lateral compartment\",\n        \"lateral compartment arthrosis\",\n        \"lateral compartment chondrosis\",\n\n        \"artrosis del compartimento lateral\",\n        \"artrosis femorotibial lateral\",\n\n        \"laterale gonarthrose\",\n        \"laterales kompartiment arthrose\",\n\n        \"lateral kompartman osteoartriti\",\n        \"lateral kompartman artrozu\"\n    ],\n\n    \"PF OA\": [\n\n        \"patellofemoral osteoarthritis\",\n        \"patellofemoral oa\",\n        \"patellofemoral degenerative\",\n        \"patellofemoral cartilage loss\",\n        \"patellofemoral joint space narrowing\",\n        \"patellar osteoarthritis\",\n\n        \"patellofemoral arthrosis\",\n        \"patellofemoral chondrosis\",\n\n        \"artrosis patelofemoral\",\n        \"condropatía patelofemoral\",\n\n        \"patellofemorale arthrose\",\n        \"patellofemorale chondrose\",\n\n        \"patellofemoral osteoartrit\",\n        \"patellofemoral artroz\"\n    ],\n\n    \"Effusion\": [\n\n        \"joint effusion\",\n        \"knee effusion\",\n        \"joint fluid\",\n        \"joint fluid collection\",\n        \"increased joint fluid\",\n\n        \"derrames articulares\",\n        \"derrame articular\",\n        \"leve derrame articular\",\n\n        \"gelenkerguss\",\n        \"kniegelenkerguss\",\n\n        \"eklem sıvısı\",\n        \"eklem içi sıvı\",\n        \"eklem efüzyonu\",\n\n        \"ставен излив\"\n    ],\n\n    \"Synovitis\": [\n\n        \"synovitis\",\n        \"synovial inflammation\",\n        \"synovial thickening\",\n        \"synovial hypertrophy\",\n\n        \"sinovitis\",\n        \"synovial inflammation\",\n\n        \"synovialitis\",\n        \"synoviale verdickung\",\n\n        \"sinovit\",\n        \"sinovyal kalınlaşma\",\n\n        \"синовит\"\n    ],\n\n    \"Baker's\": [\n\n        \"baker cyst\",\n        \"baker's cyst\",\n        \"popliteal cyst\",\n        \"popliteal cystic lesion\",\n\n        \"quiste de baker\",\n        \"quiste poplíteo\",\n\n        \"baker zyste\",\n        \"poplitealzyste\",\n\n        \"baker kisti\",\n        \"popliteal kist\",\n\n        \"киста на бейкър\"\n    ],\n\n    \"Contusion\": [\n\n        \"bone contusion\",\n        \"bone bruise\",\n        \"bone marrow contusion\",\n        \"bone marrow edema\",\n        \"bone marrow oedema\",\n        \"osseous contusion\",\n\n        \"contusión ósea\",\n        \"contusion ósea\",\n\n        \"knochenkontusion\",\n        \"knochenmarködem\",\n\n        \"kemik kontüzyon\",\n        \"kemik iliği ödemi\",\n\n        \"костна контузия\"\n    ],\n\n    \"Fracture\": [\n\n        \"fracture\",\n        \"fractured\",\n        \"fracture line\",\n        \"bone fracture\",\n\n        \"fractura\",\n        \"fractura ósea\",\n\n        \"fraktur\",\n        \"knochenfraktur\",\n\n        \"kırık\",\n        \"kemik kırığı\",\n\n        \"фрактура\"\n    ]\n}\n\n\nprint(\"Languages / terminology expanded\")\nprint()\n\nfor label in LABEL_COLUMNS:\n\n    print(\n        f\"{label:20s}: \"\n        f\"{len(MULTILINGUAL_TERMS[label])} terms\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:02.211596Z","iopub.execute_input":"2026-08-19T05:48:02.211931Z","iopub.status.idle":"2026-08-19T05:48:02.234613Z","shell.execute_reply.started":"2026-08-19T05:48:02.211907Z","shell.execute_reply":"2026-08-19T05:48:02.233964Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 66B — NEGATION-AWARE DETECTOR\n# ============================================================\n\ndef is_negated(text, start_position):\n\n    # Look at the approximately previous 80 characters\n    # around the finding.\n\n    context_start = max(\n        0,\n        start_position - 80\n    )\n\n    context = text[\n        context_start:start_position\n    ]\n\n    context = context.lower()\n\n    for neg in NEGATION_TERMS:\n\n        if neg in context:\n            return True\n\n    return False\n\n\ndef detect_abnormality(text, terms):\n\n    if not isinstance(text, str):\n        return 0\n\n    text = text.lower()\n\n    for term in terms:\n\n        term = term.lower()\n\n        start = 0\n\n        while True:\n\n            position = text.find(\n                term,\n                start\n            )\n\n            if position == -1:\n                break\n\n            # Positive if the finding is not\n            # preceded by a negation phrase.\n            if not is_negated(\n                text,\n                position\n            ):\n                return 1\n\n            start = position + len(term)\n\n    return 0\n\n\n# ------------------------------------------------------------\n# Generate improved predictions\n# ------------------------------------------------------------\n\nfor label in LABEL_COLUMNS:\n\n    labeled_df[\n        f\"Improved_Report_{label}\"\n    ] = (\n        labeled_df[\"Report_Clean\"]\n        .apply(\n            lambda x:\n                detect_abnormality(\n                    x,\n                    MULTILINGUAL_TERMS[label]\n                )\n        )\n    )\n\n\nprint(\"=\" * 80)\nprint(\"IMPROVED REPORT LABEL COUNTS\")\nprint(\"=\" * 80)\n\nfor label in LABEL_COLUMNS:\n\n    count = int(\n        labeled_df[\n            f\"Improved_Report_{label}\"\n        ].sum()\n    )\n\n    official = int(\n        labeled_df[label].sum()\n    )\n\n    print(\n        f\"{label:20s} | \"\n        f\"Official = {official:2d} | \"\n        f\"Detected = {count:2d}\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:02.235489Z","iopub.execute_input":"2026-08-19T05:48:02.235753Z","iopub.status.idle":"2026-08-19T05:48:02.271314Z","shell.execute_reply.started":"2026-08-19T05:48:02.235724Z","shell.execute_reply":"2026-08-19T05:48:02.270701Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 66C — EVALUATE IMPROVED REPORT SYSTEM\n# ============================================================\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    precision_score,\n    recall_score,\n    f1_score\n)\n\nimproved_results = []\n\n\nfor label in LABEL_COLUMNS:\n\n    y_true = (\n        labeled_df[label]\n        .astype(int)\n        .values\n    )\n\n    y_pred = (\n        labeled_df[\n            f\"Improved_Report_{label}\"\n        ]\n        .astype(int)\n        .values\n    )\n\n    improved_results.append({\n\n        \"Abnormality\": label,\n\n        \"Official_Positive\": int(\n            y_true.sum()\n        ),\n\n        \"Detected_Positive\": int(\n            y_pred.sum()\n        ),\n\n        \"Accuracy\": accuracy_score(\n            y_true,\n            y_pred\n        ),\n\n        \"Precision\": precision_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        ),\n\n        \"Recall\": recall_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        ),\n\n        \"F1\": f1_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        )\n    })\n\n\nimproved_report_eval = pd.DataFrame(\n    improved_results\n)\n\n\nprint(\"=\" * 80)\nprint(\"STEP 66 — IMPROVED REPORT PERFORMANCE\")\nprint(\"=\" * 80)\n\ndisplay(\n    improved_report_eval.round(4)\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"AVERAGE PERFORMANCE\")\nprint(\"=\" * 80)\n\nprint(\n    \"Mean Accuracy :\",\n    round(\n        improved_report_eval[\"Accuracy\"].mean(),\n        4\n    )\n)\n\nprint(\n    \"Mean Precision:\",\n    round(\n        improved_report_eval[\"Precision\"].mean(),\n        4\n    )\n)\n\nprint(\n    \"Mean Recall   :\",\n    round(\n        improved_report_eval[\"Recall\"].mean(),\n        4\n    )\n)\n\nprint(\n    \"Mean F1       :\",\n    round(\n        improved_report_eval[\"F1\"].mean(),\n        4\n    )\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:02.272393Z","iopub.execute_input":"2026-08-19T05:48:02.272668Z","iopub.status.idle":"2026-08-19T05:48:02.366634Z","shell.execute_reply.started":"2026-08-19T05:48:02.272646Z","shell.execute_reply":"2026-08-19T05:48:02.365815Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 67 — TF-IDF MULTI-LABEL REPORT CLASSIFIER\n# ============================================================\n\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn.feature_extraction.text import TfidfVectorizer\nfrom sklearn.multiclass import OneVsRestClassifier\nfrom sklearn.linear_model import LogisticRegression\n\nfrom sklearn.metrics import (\n    precision_score,\n    recall_score,\n    f1_score,\n    roc_auc_score\n)\n\nfrom sklearn.model_selection import train_test_split\n\n\nprint(\"=\" * 80)\nprint(\"STEP 67 — TF-IDF MULTI-LABEL REPORT CLASSIFIER\")\nprint(\"=\" * 80)\n\n\n# ------------------------------------------------------------\n# CHECK DATA\n# ------------------------------------------------------------\n\nrequired_columns = [\n    \"StudyInstanceUID\",\n    \"Report\"\n] + LABEL_COLUMNS\n\n\nmissing = [\n    c for c in required_columns\n    if c not in labeled_df.columns\n]\n\nif missing:\n\n    raise ValueError(\n        f\"Missing columns: {missing}\"\n    )\n\n\nprint(\"Studies:\", len(labeled_df))\nprint(\"Labels :\", len(LABEL_COLUMNS))\n\n\n# ------------------------------------------------------------\n# CLEAN REPORTS\n# ------------------------------------------------------------\n\ndef prepare_text(text):\n\n    if pd.isna(text):\n        return \"\"\n\n    text = str(text).lower()\n\n    text = text.replace(\n        \"\\n\",\n        \" \"\n    )\n\n    text = text.replace(\n        \"\\r\",\n        \" \"\n    )\n\n    text = \" \".join(\n        text.split()\n    )\n\n    return text\n\n\nX_text = (\n    labeled_df[\"Report\"]\n    .apply(prepare_text)\n    .values\n)\n\n\n# ------------------------------------------------------------\n# LABEL MATRIX\n# ------------------------------------------------------------\n\nY = (\n    labeled_df[LABEL_COLUMNS]\n    .astype(int)\n    .values\n)\n\n\nprint()\nprint(\"Text samples:\", len(X_text))\nprint(\"Label matrix:\", Y.shape)\n\n\n# ------------------------------------------------------------\n# TRAIN / VALIDATION SPLIT\n# ------------------------------------------------------------\n\ntrain_idx, val_idx = train_test_split(\n    np.arange(len(labeled_df)),\n    test_size=0.20,\n    random_state=42\n)\n\n\nX_train_text = X_text[train_idx]\nX_val_text   = X_text[val_idx]\n\nY_train = Y[train_idx]\nY_val   = Y[val_idx]\n\n\nprint()\nprint(\"=\" * 80)\nprint(\"SPLIT\")\nprint(\"=\" * 80)\n\nprint(\n    \"Training studies   :\",\n    len(train_idx)\n)\n\nprint(\n    \"Validation studies :\",\n    len(val_idx)\n)\n\n\n# ------------------------------------------------------------\n# TF-IDF\n# ------------------------------------------------------------\n\nvectorizer = TfidfVectorizer(\n\n    lowercase=True,\n\n    analyzer=\"word\",\n\n    ngram_range=(1, 2),\n\n    min_df=1,\n\n    max_df=0.98,\n\n    sublinear_tf=True,\n\n    max_features=10000\n)\n\n\nX_train = vectorizer.fit_transform(\n    X_train_text\n)\n\nX_val = vectorizer.transform(\n    X_val_text\n)\n\n\nprint()\nprint(\"=\" * 80)\nprint(\"TF-IDF\")\nprint(\"=\" * 80)\n\nprint(\n    \"Training matrix:\",\n    X_train.shape\n)\n\nprint(\n    \"Validation matrix:\",\n    X_val.shape\n)\n\n\n# ------------------------------------------------------------\n# MULTI-LABEL CLASSIFIER\n# ------------------------------------------------------------\n\nclassifier = OneVsRestClassifier(\n\n    LogisticRegression(\n\n        C=1.0,\n\n        max_iter=2000,\n\n        class_weight=\"balanced\",\n\n        solver=\"liblinear\"\n    )\n)\n\n\nprint()\nprint(\"=\" * 80)\nprint(\"TRAINING TEXT CLASSIFIER\")\nprint(\"=\" * 80)\n\n\nclassifier.fit(\n    X_train,\n    Y_train\n)\n\n\nprint(\"✓ Text classifier trained\")\n\n\n# ------------------------------------------------------------\n# PREDICTION\n# ------------------------------------------------------------\n\nY_prob = classifier.predict_proba(\n    X_val\n)\n\n\nY_pred = (\n    Y_prob >= 0.50\n).astype(int)\n\n\n# ------------------------------------------------------------\n# RESULTS\n# ------------------------------------------------------------\n\nresults = []\n\n\nfor i, label in enumerate(\n    LABEL_COLUMNS\n):\n\n    y_true = Y_val[:, i]\n\n    y_probability = Y_prob[:, i]\n\n    y_pred = Y_pred[:, i]\n\n    try:\n\n        auc = roc_auc_score(\n            y_true,\n            y_probability\n        )\n\n    except:\n\n        auc = np.nan\n\n\n    results.append({\n\n        \"Abnormality\": label,\n\n        \"Positive\": int(\n            y_true.sum()\n        ),\n\n        \"Predicted\": int(\n            y_pred.sum()\n        ),\n\n        \"AUC\": auc,\n\n        \"Precision\": precision_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        ),\n\n        \"Recall\": recall_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        ),\n\n        \"F1\": f1_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        )\n    })\n\n\ntext_results = pd.DataFrame(\n    results\n)\n\n\nprint()\nprint(\"=\" * 80)\nprint(\"STEP 67 — TEXT CLASSIFIER RESULTS\")\nprint(\"=\" * 80)\n\ndisplay(\n    text_results.round(4)\n)\n\n\n# ------------------------------------------------------------\n# MACRO RESULTS\n# ------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"MACRO PERFORMANCE\")\nprint(\"=\" * 80)\n\nprint(\n    \"Macro AUC       :\",\n    round(\n        text_results[\"AUC\"].mean(),\n        4\n    )\n)\n\nprint(\n    \"Macro Precision :\",\n    round(\n        text_results[\"Precision\"].mean(),\n        4\n    )\n)\n\nprint(\n    \"Macro Recall    :\",\n    round(\n        text_results[\"Recall\"].mean(),\n        4\n    )\n)\n\nprint(\n    \"Macro F1        :\",\n    round(\n        text_results[\"F1\"].mean(),\n        4\n    )\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:02.367568Z","iopub.execute_input":"2026-08-19T05:48:02.367893Z","iopub.status.idle":"2026-08-19T05:48:02.52064Z","shell.execute_reply.started":"2026-08-19T05:48:02.367838Z","shell.execute_reply":"2026-08-19T05:48:02.519939Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 68 — COMPREHENSIVE VALIDATION EVALUATION\n# FIXED FOR PYTORCH 2.6+\n# ============================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\n\nfrom sklearn.metrics import (\n    roc_auc_score,\n    average_precision_score,\n    precision_score,\n    recall_score,\n    f1_score\n)\n\n\nprint(\"=\" * 80)\nprint(\"STEP 68 — COMPREHENSIVE VALIDATION EVALUATION\")\nprint(\"=\" * 80)\n\n\n# ============================================================\n# 1. CHECK REQUIRED VARIABLES\n# ============================================================\n\nrequired_variables = [\n    \"model\",\n    \"val_loader\",\n    \"LABEL_COLUMNS\",\n    \"DEVICE\"\n]\n\nprint(\"\\nChecking required variables...\")\n\nfor variable in required_variables:\n\n    if variable not in globals():\n\n        raise NameError(\n            f\"Required variable '{variable}' is not defined.\"\n        )\n\n    print(f\"✓ {variable}\")\n\n\n# ============================================================\n# 2. CHECK CHECKPOINT\n# ============================================================\n\nCHECKPOINT = (\n    \"/kaggle/working/\"\n    \"best_rsna_knee_model.pth\"\n)\n\nprint(\"\\nCheckpoint:\")\nprint(CHECKPOINT)\n\n\nif not os.path.exists(CHECKPOINT):\n\n    raise FileNotFoundError(\n        f\"\\nCheckpoint not found:\\n{CHECKPOINT}\"\n    )\n\n\nprint(\"✓ Checkpoint exists\")\n\n\n# ============================================================\n# 3. LOAD CHECKPOINT\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"LOADING BEST MODEL CHECKPOINT\")\nprint(\"=\" * 80)\n\nprint(\n    \"Using weights_only=False because this is \"\n    \"the trusted checkpoint created during training.\"\n)\n\n\ncheckpoint = torch.load(\n    CHECKPOINT,\n    map_location=DEVICE,\n    weights_only=False\n)\n\n\nprint(\"✓ Checkpoint loaded successfully\")\n\n\n# ============================================================\n# 4. IDENTIFY CHECKPOINT FORMAT\n# ============================================================\n\nprint(\"\\nCheckpoint type:\")\n\nprint(\n    type(checkpoint)\n)\n\n\nif isinstance(checkpoint, dict):\n\n    print(\"\\nCheckpoint keys:\")\n\n    print(\n        list(checkpoint.keys())\n    )\n\n\n# ============================================================\n# 5. EXTRACT MODEL STATE DICTIONARY\n# ============================================================\n\nif isinstance(checkpoint, dict):\n\n    if \"model_state_dict\" in checkpoint:\n\n        state_dict = (\n            checkpoint[\"model_state_dict\"]\n        )\n\n        print(\n            \"\\n✓ Found 'model_state_dict'\"\n        )\n\n    elif \"state_dict\" in checkpoint:\n\n        state_dict = (\n            checkpoint[\"state_dict\"]\n        )\n\n        print(\n            \"\\n✓ Found 'state_dict'\"\n        )\n\n    else:\n\n        # Check whether the dictionary itself\n        # looks like a state dictionary.\n\n        tensor_values = [\n            isinstance(v, torch.Tensor)\n            for v in checkpoint.values()\n        ]\n\n        if (\n            len(tensor_values) > 0\n            and all(tensor_values)\n        ):\n\n            state_dict = checkpoint\n\n            print(\n                \"\\n✓ Checkpoint itself is a state_dict\"\n            )\n\n        else:\n\n            raise RuntimeError(\n                \"Could not find model state dictionary \"\n                \"inside checkpoint.\"\n            )\n\nelse:\n\n    raise RuntimeError(\n        \"Unexpected checkpoint format.\"\n    )\n\n\n# ============================================================\n# 6. LOAD MODEL WEIGHTS\n# ============================================================\n\ntry:\n\n    model.load_state_dict(\n        state_dict,\n        strict=True\n    )\n\n    print(\n        \"✓ Model weights loaded with strict=True\"\n    )\n\nexcept RuntimeError as e:\n\n    print(\n        \"\\nStrict loading failed.\"\n    )\n\n    print(\n        \"Trying strict=False...\"\n    )\n\n    result = model.load_state_dict(\n        state_dict,\n        strict=False\n    )\n\n    print(\n        \"Missing keys:\",\n        result.missing_keys\n    )\n\n    print(\n        \"Unexpected keys:\",\n        result.unexpected_keys\n    )\n\n\n# ============================================================\n# 7. MOVE MODEL TO DEVICE\n# ============================================================\n\nmodel = model.to(DEVICE)\n\nmodel.eval()\n\n\nprint(\n    \"✓ Model moved to:\",\n    DEVICE\n)\n\nprint(\n    \"✓ Model set to evaluation mode\"\n)\n\n\n# ============================================================\n# 8. VALIDATION PREDICTIONS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"GENERATING VALIDATION PREDICTIONS\")\nprint(\"=\" * 80)\n\n\nall_probabilities = []\n\nall_labels = []\n\nall_study_ids = []\n\n\nwith torch.no_grad():\n\n    for batch_index, batch in enumerate(\n        val_loader\n    ):\n\n        # ----------------------------------------------------\n        # Images\n        # ----------------------------------------------------\n\n        images = batch[\"images\"].to(\n            DEVICE,\n            non_blocking=True\n        )\n\n\n        # ----------------------------------------------------\n        # Plane mask\n        # ----------------------------------------------------\n\n        plane_mask = batch[\n            \"plane_mask\"\n        ].to(\n            DEVICE,\n            non_blocking=True\n        )\n\n\n        # ----------------------------------------------------\n        # Labels\n        # ----------------------------------------------------\n\n        labels = batch[\"labels\"].to(\n            DEVICE,\n            non_blocking=True\n        )\n\n\n        # ----------------------------------------------------\n        # Forward pass\n        # ----------------------------------------------------\n\n        logits = model(\n            images,\n            plane_mask\n        )\n\n\n        # ----------------------------------------------------\n        # Sigmoid probabilities\n        # ----------------------------------------------------\n\n        probabilities = torch.sigmoid(\n            logits\n        )\n\n\n        all_probabilities.append(\n            probabilities.detach()\n            .cpu()\n            .numpy()\n        )\n\n\n        all_labels.append(\n            labels.detach()\n            .cpu()\n            .numpy()\n        )\n\n\n        # ----------------------------------------------------\n        # Study IDs if available\n        # ----------------------------------------------------\n\n        if \"study_ids\" in batch:\n\n            all_study_ids.extend(\n                batch[\"study_ids\"]\n            )\n\n\nprint(\n    \"\\n✓ Validation inference completed\"\n)\n\n\n# ============================================================\n# 9. COMBINE PREDICTIONS\n# ============================================================\n\ny_prob = np.concatenate(\n    all_probabilities,\n    axis=0\n)\n\n\ny_true = np.concatenate(\n    all_labels,\n    axis=0\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"PREDICTION MATRIX\")\nprint(\"=\" * 80)\n\n\nprint(\n    \"Validation samples :\",\n    len(y_true)\n)\n\n\nprint(\n    \"Probability shape  :\",\n    y_prob.shape\n)\n\n\nprint(\n    \"Label shape        :\",\n    y_true.shape\n)\n\n\n# ============================================================\n# 10. CHECK SHAPES\n# ============================================================\n\nexpected_classes = len(\n    LABEL_COLUMNS\n)\n\n\nif y_prob.shape[1] != expected_classes:\n\n    raise ValueError(\n        f\"Probability matrix has \"\n        f\"{y_prob.shape[1]} classes, \"\n        f\"expected {expected_classes}.\"\n    )\n\n\nif y_true.shape[1] != expected_classes:\n\n    raise ValueError(\n        f\"Label matrix has \"\n        f\"{y_true.shape[1]} classes, \"\n        f\"expected {expected_classes}.\"\n    )\n\n\nprint(\n    \"\\n✓ Prediction dimensions correct\"\n)\n\n\n# ============================================================\n# 11. DEFAULT 0.50 PERFORMANCE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"PERFORMANCE AT THRESHOLD = 0.50\")\nprint(\"=\" * 80)\n\n\ndefault_threshold = 0.50\n\n\ny_pred_default = (\n    y_prob >= default_threshold\n).astype(int)\n\n\ndefault_results = []\n\n\nfor class_index, label in enumerate(\n    LABEL_COLUMNS\n):\n\n    true_class = y_true[\n        :,\n        class_index\n    ]\n\n\n    prob_class = y_prob[\n        :,\n        class_index\n    ]\n\n\n    pred_class = y_pred_default[\n        :,\n        class_index\n    ]\n\n\n    try:\n\n        auc = roc_auc_score(\n            true_class,\n            prob_class\n        )\n\n    except ValueError:\n\n        auc = np.nan\n\n\n    try:\n\n        ap = average_precision_score(\n            true_class,\n            prob_class\n        )\n\n    except ValueError:\n\n        ap = np.nan\n\n\n    precision = precision_score(\n        true_class,\n        pred_class,\n        zero_division=0\n    )\n\n\n    recall = recall_score(\n        true_class,\n        pred_class,\n        zero_division=0\n    )\n\n\n    f1 = f1_score(\n        true_class,\n        pred_class,\n        zero_division=0\n    )\n\n\n    default_results.append({\n\n        \"Abnormality\": label,\n\n        \"Positive\": int(\n            true_class.sum()\n        ),\n\n        \"Predicted\": int(\n            pred_class.sum()\n        ),\n\n        \"AUC\": auc,\n\n        \"AP\": ap,\n\n        \"Precision\": precision,\n\n        \"Recall\": recall,\n\n        \"F1\": f1\n    })\n\n\ndefault_results_df = pd.DataFrame(\n    default_results\n)\n\n\ndisplay(\n    default_results_df.round(4)\n)\n\n\n# ============================================================\n# 12. OPTIMIZE THRESHOLD FOR EACH CLASS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"OPTIMIZING PER-CLASS THRESHOLDS\")\nprint(\"=\" * 80)\n\n\ncandidate_thresholds = np.arange(\n    0.10,\n    0.91,\n    0.05\n)\n\n\noptimized_thresholds = []\n\noptimized_results = []\n\n\nfor class_index, label in enumerate(\n    LABEL_COLUMNS\n):\n\n    true_class = y_true[\n        :,\n        class_index\n    ]\n\n\n    prob_class = y_prob[\n        :,\n        class_index\n    ]\n\n\n    best_threshold = 0.50\n\n    best_f1 = -1.0\n\n\n    # --------------------------------------------------------\n    # Search thresholds\n    # --------------------------------------------------------\n\n    for threshold in candidate_thresholds:\n\n        pred_class = (\n            prob_class >= threshold\n        ).astype(int)\n\n\n        score = f1_score(\n            true_class,\n            pred_class,\n            zero_division=0\n        )\n\n\n        if score > best_f1:\n\n            best_f1 = score\n\n            best_threshold = threshold\n\n\n    optimized_thresholds.append(\n        best_threshold\n    )\n\n\n    # --------------------------------------------------------\n    # Predictions using best threshold\n    # --------------------------------------------------------\n\n    pred_class = (\n        prob_class >= best_threshold\n    ).astype(int)\n\n\n    # --------------------------------------------------------\n    # Metrics\n    # --------------------------------------------------------\n\n    try:\n\n        auc = roc_auc_score(\n            true_class,\n            prob_class\n        )\n\n    except ValueError:\n\n        auc = np.nan\n\n\n    try:\n\n        ap = average_precision_score(\n            true_class,\n            prob_class\n        )\n\n    except ValueError:\n\n        ap = np.nan\n\n\n    precision = precision_score(\n        true_class,\n        pred_class,\n        zero_division=0\n    )\n\n\n    recall = recall_score(\n        true_class,\n        pred_class,\n        zero_division=0\n    )\n\n\n    f1 = f1_score(\n        true_class,\n        pred_class,\n        zero_division=0\n    )\n\n\n    optimized_results.append({\n\n        \"Abnormality\": label,\n\n        \"Positive\": int(\n            true_class.sum()\n        ),\n\n        \"Threshold\": best_threshold,\n\n        \"Predicted\": int(\n            pred_class.sum()\n        ),\n\n        \"AUC\": auc,\n\n        \"AP\": ap,\n\n        \"Precision\": precision,\n\n        \"Recall\": recall,\n\n        \"F1\": f1\n    })\n\n\noptimized_results_df = pd.DataFrame(\n    optimized_results\n)\n\n\nprint(\n    \"\\n✓ Per-class thresholds optimized\"\n)\n\n\n# ============================================================\n# 13. DISPLAY OPTIMIZED RESULTS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"PER-CLASS VALIDATION RESULTS\")\nprint(\"=\" * 80)\n\n\ndisplay(\n    optimized_results_df.round(4)\n)\n\n\n# ============================================================\n# 14. MACRO PERFORMANCE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"MACRO PERFORMANCE\")\nprint(\"=\" * 80)\n\n\nmacro_auc = optimized_results_df[\n    \"AUC\"\n].mean()\n\n\nmacro_ap = optimized_results_df[\n    \"AP\"\n].mean()\n\n\nmacro_precision = optimized_results_df[\n    \"Precision\"\n].mean()\n\n\nmacro_recall = optimized_results_df[\n    \"Recall\"\n].mean()\n\n\nmacro_f1 = optimized_results_df[\n    \"F1\"\n].mean()\n\n\nprint(\n    \"Macro AUC       :\",\n    round(macro_auc, 4)\n)\n\n\nprint(\n    \"Macro AP        :\",\n    round(macro_ap, 4)\n)\n\n\nprint(\n    \"Macro Precision :\",\n    round(macro_precision, 4)\n)\n\n\nprint(\n    \"Macro Recall    :\",\n    round(macro_recall, 4)\n)\n\n\nprint(\n    \"Macro F1        :\",\n    round(macro_f1, 4)\n)\n\n\n# ============================================================\n# 15. MICRO PERFORMANCE\n# ============================================================\n\nthreshold_array = np.array(\n    optimized_thresholds\n)\n\n\ny_pred_optimized = (\n    y_prob\n    >= threshold_array.reshape(\n        1,\n        -1\n    )\n).astype(int)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"MICRO PERFORMANCE\")\nprint(\"=\" * 80)\n\n\nmicro_precision = precision_score(\n    y_true.flatten(),\n    y_pred_optimized.flatten(),\n    zero_division=0\n)\n\n\nmicro_recall = recall_score(\n    y_true.flatten(),\n    y_pred_optimized.flatten(),\n    zero_division=0\n)\n\n\nmicro_f1 = f1_score(\n    y_true.flatten(),\n    y_pred_optimized.flatten(),\n    zero_division=0\n)\n\n\nprint(\n    \"Micro Precision :\",\n    round(micro_precision, 4)\n)\n\n\nprint(\n    \"Micro Recall    :\",\n    round(micro_recall, 4)\n)\n\n\nprint(\n    \"Micro F1        :\",\n    round(micro_f1, 4)\n)\n\n\n# ============================================================\n# 16. THRESHOLD TABLE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"OPTIMIZED THRESHOLDS\")\nprint(\"=\" * 80)\n\n\nthreshold_df = pd.DataFrame({\n\n    \"Abnormality\":\n        LABEL_COLUMNS,\n\n    \"Optimal_Threshold\":\n        optimized_thresholds\n})\n\n\ndisplay(\n    threshold_df.round(3)\n)\n\n\n# ============================================================\n# 17. SAVE RESULTS\n# ============================================================\n\nRESULT_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_validation_results.csv\"\n)\n\n\nTHRESHOLD_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_optimal_thresholds.csv\"\n)\n\n\noptimized_results_df.to_csv(\n    RESULT_PATH,\n    index=False\n)\n\n\nthreshold_df.to_csv(\n    THRESHOLD_PATH,\n    index=False\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 68 COMPLETE\")\nprint(\"=\" * 80)\n\n\nprint(\n    \"✓ Best checkpoint loaded\"\n)\n\nprint(\n    \"✓ Validation predictions generated\"\n)\n\nprint(\n    \"✓ Default threshold evaluation completed\"\n)\n\nprint(\n    \"✓ Per-class thresholds optimized\"\n)\n\nprint(\n    \"✓ AUC calculated\"\n)\n\nprint(\n    \"✓ Average Precision calculated\"\n)\n\nprint(\n    \"✓ Precision calculated\"\n)\n\nprint(\n    \"✓ Recall calculated\"\n)\n\nprint(\n    \"✓ F1 calculated\"\n)\n\nprint(\n    \"\\nResults:\"\n)\n\nprint(\n    RESULT_PATH\n)\n\nprint(\n    \"\\nThresholds:\"\n)\n\nprint(\n    THRESHOLD_PATH\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:02.521989Z","iopub.execute_input":"2026-08-19T05:48:02.522382Z","iopub.status.idle":"2026-08-19T05:48:09.189533Z","shell.execute_reply.started":"2026-08-19T05:48:02.522326Z","shell.execute_reply":"2026-08-19T05:48:09.188703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 69 — VALIDATION RESULT ANALYSIS\n# ============================================================\n\nimport pandas as pd\nimport numpy as np\n\nRESULT_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_validation_results.csv\"\n)\n\nTHRESHOLD_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_optimal_thresholds.csv\"\n)\n\nresults = pd.read_csv(RESULT_PATH)\nthresholds = pd.read_csv(THRESHOLD_PATH)\n\nprint(\"=\" * 80)\nprint(\"STEP 69 — VALIDATION RESULT ANALYSIS\")\nprint(\"=\" * 80)\n\nprint(\"\\nValidation results:\")\ndisplay(results.round(4))\n\nprint(\"\\nOptimized thresholds:\")\ndisplay(thresholds.round(4))\n\n\n# ============================================================\n# RANK CLASSES BY F1\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CLASSES RANKED BY F1\")\nprint(\"=\" * 80)\n\nf1_ranked = (\n    results[\n        [\n            \"Abnormality\",\n            \"AUC\",\n            \"AP\",\n            \"Precision\",\n            \"Recall\",\n            \"F1\"\n        ]\n    ]\n    .sort_values(\n        \"F1\",\n        ascending=False\n    )\n)\n\ndisplay(\n    f1_ranked.round(4)\n)\n\n\n# ============================================================\n# RANK CLASSES BY AUC\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CLASSES RANKED BY AUC\")\nprint(\"=\" * 80)\n\nauc_ranked = (\n    results[\n        [\n            \"Abnormality\",\n            \"AUC\",\n            \"AP\",\n            \"F1\"\n        ]\n    ]\n    .sort_values(\n        \"AUC\",\n        ascending=False\n    )\n)\n\ndisplay(\n    auc_ranked.round(4)\n)\n\n\n# ============================================================\n# STRONG / WEAK CLASSES\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CLASS GROUPING\")\nprint(\"=\" * 80)\n\nstrong = results[\n    results[\"F1\"] >= 0.40\n]\n\nmoderate = results[\n    (results[\"F1\"] >= 0.20) &\n    (results[\"F1\"] < 0.40)\n]\n\nweak = results[\n    results[\"F1\"] < 0.20\n]\n\nprint(\"\\nStrong classes (F1 >= 0.40):\")\n\ndisplay(\n    strong.round(4)\n)\n\nprint(\"\\nModerate classes (0.20 <= F1 < 0.40):\")\n\ndisplay(\n    moderate.round(4)\n)\n\nprint(\"\\nWeak classes (F1 < 0.20):\")\n\ndisplay(\n    weak.round(4)\n)\n\n\n# ============================================================\n# OVERALL SUMMARY\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"OVERALL SUMMARY\")\nprint(\"=\" * 80)\n\nprint(\n    \"Mean AUC       :\",\n    round(results[\"AUC\"].mean(), 4)\n)\n\nprint(\n    \"Mean AP        :\",\n    round(results[\"AP\"].mean(), 4)\n)\n\nprint(\n    \"Mean Precision :\",\n    round(results[\"Precision\"].mean(), 4)\n)\n\nprint(\n    \"Mean Recall    :\",\n    round(results[\"Recall\"].mean(), 4)\n)\n\nprint(\n    \"Mean F1        :\",\n    round(results[\"F1\"].mean(), 4)\n)\n\nprint(\n    \"\\nBest F1 class  :\",\n    results.loc[\n        results[\"F1\"].idxmax(),\n        \"Abnormality\"\n    ]\n)\n\nprint(\n    \"Best F1        :\",\n    round(\n        results[\"F1\"].max(),\n        4\n    )\n)\n\nprint(\n    \"\\nBest AUC class :\",\n    results.loc[\n        results[\"AUC\"].idxmax(),\n        \"Abnormality\"\n    ]\n)\n\nprint(\n    \"Best AUC        :\",\n    round(\n        results[\"AUC\"].max(),\n        4\n    )\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 69 COMPLETE\")\nprint(\"=\" * 80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:09.19064Z","iopub.execute_input":"2026-08-19T05:48:09.190963Z","iopub.status.idle":"2026-08-19T05:48:09.264692Z","shell.execute_reply.started":"2026-08-19T05:48:09.190937Z","shell.execute_reply":"2026-08-19T05:48:09.263971Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 69 — VALIDATION RESULT ANALYSIS\n# ============================================================\n\nimport pandas as pd\nimport numpy as np\n\nRESULT_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_validation_results.csv\"\n)\n\nTHRESHOLD_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_optimal_thresholds.csv\"\n)\n\nresults = pd.read_csv(RESULT_PATH)\nthresholds = pd.read_csv(THRESHOLD_PATH)\n\nprint(\"=\" * 80)\nprint(\"STEP 69 — VALIDATION RESULT ANALYSIS\")\nprint(\"=\" * 80)\n\nprint(\"\\nValidation results:\")\ndisplay(results.round(4))\n\nprint(\"\\nOptimized thresholds:\")\ndisplay(thresholds.round(4))\n\n\n# ============================================================\n# RANK CLASSES BY F1\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CLASSES RANKED BY F1\")\nprint(\"=\" * 80)\n\nf1_ranked = (\n    results[\n        [\n            \"Abnormality\",\n            \"AUC\",\n            \"AP\",\n            \"Precision\",\n            \"Recall\",\n            \"F1\"\n        ]\n    ]\n    .sort_values(\n        \"F1\",\n        ascending=False\n    )\n)\n\ndisplay(\n    f1_ranked.round(4)\n)\n\n\n# ============================================================\n# RANK CLASSES BY AUC\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CLASSES RANKED BY AUC\")\nprint(\"=\" * 80)\n\nauc_ranked = (\n    results[\n        [\n            \"Abnormality\",\n            \"AUC\",\n            \"AP\",\n            \"F1\"\n        ]\n    ]\n    .sort_values(\n        \"AUC\",\n        ascending=False\n    )\n)\n\ndisplay(\n    auc_ranked.round(4)\n)\n\n\n# ============================================================\n# STRONG / WEAK CLASSES\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CLASS GROUPING\")\nprint(\"=\" * 80)\n\nstrong = results[\n    results[\"F1\"] >= 0.40\n]\n\nmoderate = results[\n    (results[\"F1\"] >= 0.20) &\n    (results[\"F1\"] < 0.40)\n]\n\nweak = results[\n    results[\"F1\"] < 0.20\n]\n\nprint(\"\\nStrong classes (F1 >= 0.40):\")\n\ndisplay(\n    strong.round(4)\n)\n\nprint(\"\\nModerate classes (0.20 <= F1 < 0.40):\")\n\ndisplay(\n    moderate.round(4)\n)\n\nprint(\"\\nWeak classes (F1 < 0.20):\")\n\ndisplay(\n    weak.round(4)\n)\n\n\n# ============================================================\n# OVERALL SUMMARY\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"OVERALL SUMMARY\")\nprint(\"=\" * 80)\n\nprint(\n    \"Mean AUC       :\",\n    round(results[\"AUC\"].mean(), 4)\n)\n\nprint(\n    \"Mean AP        :\",\n    round(results[\"AP\"].mean(), 4)\n)\n\nprint(\n    \"Mean Precision :\",\n    round(results[\"Precision\"].mean(), 4)\n)\n\nprint(\n    \"Mean Recall    :\",\n    round(results[\"Recall\"].mean(), 4)\n)\n\nprint(\n    \"Mean F1        :\",\n    round(results[\"F1\"].mean(), 4)\n)\n\nprint(\n    \"\\nBest F1 class  :\",\n    results.loc[\n        results[\"F1\"].idxmax(),\n        \"Abnormality\"\n    ]\n)\n\nprint(\n    \"Best F1        :\",\n    round(\n        results[\"F1\"].max(),\n        4\n    )\n)\n\nprint(\n    \"\\nBest AUC class :\",\n    results.loc[\n        results[\"AUC\"].idxmax(),\n        \"Abnormality\"\n    ]\n)\n\nprint(\n    \"Best AUC        :\",\n    round(\n        results[\"AUC\"].max(),\n        4\n    )\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 69 COMPLETE\")\nprint(\"=\" * 80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:09.265664Z","iopub.execute_input":"2026-08-19T05:48:09.266014Z","iopub.status.idle":"2026-08-19T05:48:09.338564Z","shell.execute_reply.started":"2026-08-19T05:48:09.265989Z","shell.execute_reply":"2026-08-19T05:48:09.337827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 70 — DETAILED PER-CLASS ERROR ANALYSIS\n# ============================================================\n\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn.metrics import (\n    confusion_matrix,\n    precision_score,\n    recall_score,\n    f1_score,\n    roc_auc_score,\n    average_precision_score\n)\n\n\nprint(\"=\" * 80)\nprint(\"STEP 70 — DETAILED PER-CLASS ERROR ANALYSIS\")\nprint(\"=\" * 80)\n\n\n# ============================================================\n# CHECK VARIABLES FROM STEP 68\n# ============================================================\n\nrequired = [\n    \"y_true\",\n    \"y_prob\",\n    \"LABEL_COLUMNS\",\n    \"optimized_thresholds\"\n]\n\nfor variable in required:\n\n    if variable not in globals():\n\n        raise NameError(\n            f\"Required variable '{variable}' is missing. \"\n            f\"Run STEP 68 first.\"\n        )\n\n    print(\n        f\"✓ {variable}\"\n    )\n\n\n# ============================================================\n# ARRAYS\n# ============================================================\n\nthresholds = np.asarray(\n    optimized_thresholds,\n    dtype=float\n)\n\n\ny_pred = (\n    y_prob >= thresholds.reshape(1, -1)\n).astype(int)\n\n\nprint(\"\\nProbability shape:\")\nprint(y_prob.shape)\n\nprint(\"\\nGround-truth shape:\")\nprint(y_true.shape)\n\nprint(\"\\nPrediction shape:\")\nprint(y_pred.shape)\n\n\n# ============================================================\n# PER-CLASS ANALYSIS\n# ============================================================\n\nanalysis_rows = []\n\n\nfor i, label in enumerate(\n    LABEL_COLUMNS\n):\n\n    true = y_true[:, i]\n\n    pred = y_pred[:, i]\n\n    prob = y_prob[:, i]\n\n\n    # --------------------------------------------------------\n    # Confusion matrix\n    # --------------------------------------------------------\n\n    tn, fp, fn, tp = confusion_matrix(\n        true,\n        pred,\n        labels=[0, 1]\n    ).ravel()\n\n\n    # --------------------------------------------------------\n    # Metrics\n    # --------------------------------------------------------\n\n    precision = precision_score(\n        true,\n        pred,\n        zero_division=0\n    )\n\n\n    recall = recall_score(\n        true,\n        pred,\n        zero_division=0\n    )\n\n\n    f1 = f1_score(\n        true,\n        pred,\n        zero_division=0\n    )\n\n\n    try:\n\n        auc = roc_auc_score(\n            true,\n            prob\n        )\n\n    except:\n\n        auc = np.nan\n\n\n    try:\n\n        ap = average_precision_score(\n            true,\n            prob\n        )\n\n    except:\n\n        ap = np.nan\n\n\n    # --------------------------------------------------------\n    # Specificity\n    # --------------------------------------------------------\n\n    specificity = (\n        tn / (tn + fp)\n        if (tn + fp) > 0\n        else 0\n    )\n\n\n    # --------------------------------------------------------\n    # Positive rate\n    # --------------------------------------------------------\n\n    actual_positive_rate = (\n        true.mean()\n    )\n\n\n    predicted_positive_rate = (\n        pred.mean()\n    )\n\n\n    analysis_rows.append({\n\n        \"Abnormality\": label,\n\n        \"Threshold\": thresholds[i],\n\n        \"Actual_Positive\": int(\n            true.sum()\n        ),\n\n        \"Predicted_Positive\": int(\n            pred.sum()\n        ),\n\n        \"TP\": int(tp),\n\n        \"FP\": int(fp),\n\n        \"FN\": int(fn),\n\n        \"TN\": int(tn),\n\n        \"AUC\": auc,\n\n        \"AP\": ap,\n\n        \"Precision\": precision,\n\n        \"Recall\": recall,\n\n        \"Specificity\": specificity,\n\n        \"F1\": f1,\n\n        \"Actual_Positive_Rate\":\n            actual_positive_rate,\n\n        \"Predicted_Positive_Rate\":\n            predicted_positive_rate\n    })\n\n\nanalysis_df = pd.DataFrame(\n    analysis_rows\n)\n\n\n# ============================================================\n# DISPLAY\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"PER-CLASS ERROR ANALYSIS\")\nprint(\"=\" * 80)\n\ndisplay(\n    analysis_df.round(4)\n)\n\n\n# ============================================================\n# SORT BY F1\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CLASSES — BEST TO WORST F1\")\nprint(\"=\" * 80)\n\ndisplay(\n    analysis_df[\n        [\n            \"Abnormality\",\n            \"AUC\",\n            \"AP\",\n            \"Precision\",\n            \"Recall\",\n            \"Specificity\",\n            \"F1\"\n        ]\n    ]\n    .sort_values(\n        \"F1\",\n        ascending=False\n    )\n    .round(4)\n)\n\n\n# ============================================================\n# HIGH AUC BUT LOW F1\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"HIGH AUC / LOW F1 CLASSES\")\nprint(\"=\" * 80)\n\nhigh_auc_low_f1 = analysis_df[\n    (analysis_df[\"AUC\"] >= 0.60) &\n    (analysis_df[\"F1\"] < 0.40)\n]\n\ndisplay(\n    high_auc_low_f1.round(4)\n)\n\n\n# ============================================================\n# PRECISION PROBLEM\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CLASSES WITH LOW PRECISION\")\nprint(\"=\" * 80)\n\nlow_precision = analysis_df[\n    analysis_df[\"Precision\"] < 0.30\n]\n\ndisplay(\n    low_precision.round(4)\n)\n\n\n# ============================================================\n# RECALL PROBLEM\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CLASSES WITH LOW RECALL\")\nprint(\"=\" * 80)\n\nlow_recall = analysis_df[\n    analysis_df[\"Recall\"] < 0.50\n]\n\ndisplay(\n    low_recall.round(4)\n)\n\n\n# ============================================================\n# SAVE\n# ============================================================\n\nERROR_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_detailed_error_analysis.csv\"\n)\n\n\nanalysis_df.to_csv(\n    ERROR_PATH,\n    index=False\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 70 COMPLETE\")\nprint(\"=\" * 80)\n\nprint(\n    \"Saved:\",\n    ERROR_PATH\n)\n\nprint(\"\\n✓ TP / FP / FN / TN calculated\")\nprint(\"✓ Specificity calculated\")\nprint(\"✓ Precision calculated\")\nprint(\"✓ Recall calculated\")\nprint(\"✓ F1 calculated\")\nprint(\"✓ AUC calculated\")\nprint(\"✓ Average Precision calculated\")\nprint(\"✓ Per-class error patterns identified\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:09.341383Z","iopub.execute_input":"2026-08-19T05:48:09.341611Z","iopub.status.idle":"2026-08-19T05:48:09.517018Z","shell.execute_reply.started":"2026-08-19T05:48:09.341591Z","shell.execute_reply":"2026-08-19T05:48:09.51645Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 71 — AUTOMATIC ERROR ANALYSIS\n# ============================================================\n\nimport os\nimport pandas as pd\nimport numpy as np\n\nERROR_PATH = \"/kaggle/working/rsna_knee_detailed_error_analysis.csv\"\n\nif not os.path.exists(ERROR_PATH):\n    raise FileNotFoundError(\n        f\"Error analysis file not found:\\n{ERROR_PATH}\\n\"\n        \"Run STEP 70 first.\"\n    )\n\ndf = pd.read_csv(ERROR_PATH)\n\nprint(\"=\" * 80)\nprint(\"STEP 71 — AUTOMATIC MODEL ERROR ANALYSIS\")\nprint(\"=\" * 80)\n\nprint(\"\\nRows:\", len(df))\nprint(\"Columns:\", len(df.columns))\n\n\n# ============================================================\n# REQUIRED COLUMNS\n# ============================================================\n\nrequired_columns = [\n    \"Abnormality\",\n    \"Threshold\",\n    \"Actual_Positive\",\n    \"Predicted_Positive\",\n    \"TP\",\n    \"FP\",\n    \"FN\",\n    \"TN\",\n    \"AUC\",\n    \"AP\",\n    \"Precision\",\n    \"Recall\",\n    \"Specificity\",\n    \"F1\"\n]\n\nmissing = [\n    c for c in required_columns\n    if c not in df.columns\n]\n\nif missing:\n    raise ValueError(\n        f\"Missing columns: {missing}\"\n    )\n\nprint(\"\\n✓ Required columns found\")\n\n\n# ============================================================\n# CLEAN NUMERIC DATA\n# ============================================================\n\nnumeric_columns = [\n    \"Threshold\",\n    \"Actual_Positive\",\n    \"Predicted_Positive\",\n    \"TP\",\n    \"FP\",\n    \"FN\",\n    \"TN\",\n    \"AUC\",\n    \"AP\",\n    \"Precision\",\n    \"Recall\",\n    \"Specificity\",\n    \"F1\"\n]\n\nfor col in numeric_columns:\n    df[col] = pd.to_numeric(\n        df[col],\n        errors=\"coerce\"\n    )\n\n\n# ============================================================\n# MODEL QUALITY CATEGORY\n# ============================================================\n\ndef classify_model_behavior(row):\n\n    auc = row[\"AUC\"]\n    f1 = row[\"F1\"]\n    precision = row[\"Precision\"]\n    recall = row[\"Recall\"]\n\n    if pd.isna(auc):\n        return \"NO_AUC\"\n\n    # Strong ranking ability\n    if auc >= 0.75 and f1 >= 0.50:\n        return \"STRONG\"\n\n    # Good ranking but threshold/classification problem\n    if auc >= 0.70 and f1 < 0.50:\n        return \"GOOD_AUC_THRESHOLD_PROBLEM\"\n\n    # Reasonable ranking\n    if auc >= 0.60 and f1 < 0.40:\n        return \"LEARNING_BUT_WEAK_CLASSIFICATION\"\n\n    # High recall but poor precision\n    if recall >= 0.70 and precision < 0.40:\n        return \"HIGH_RECALL_LOW_PRECISION\"\n\n    # Weak ranking\n    if auc < 0.55:\n        return \"WEAK_LEARNING\"\n\n    return \"MODERATE\"\n\n\ndf[\"Model_Behavior\"] = df.apply(\n    classify_model_behavior,\n    axis=1\n)\n\n\n# ============================================================\n# ERROR BURDEN\n# ============================================================\n\ndf[\"Total_Errors\"] = (\n    df[\"FP\"] + df[\"FN\"]\n)\n\ndf[\"False_Positive_Rate\"] = np.where(\n    (df[\"FP\"] + df[\"TN\"]) > 0,\n    df[\"FP\"] / (\n        df[\"FP\"] + df[\"TN\"]\n    ),\n    0\n)\n\n\n# ============================================================\n# SORTED SUMMARY\n# ============================================================\n\nsummary = df[\n    [\n        \"Abnormality\",\n        \"Threshold\",\n        \"Actual_Positive\",\n        \"Predicted_Positive\",\n        \"TP\",\n        \"FP\",\n        \"FN\",\n        \"TN\",\n        \"AUC\",\n        \"AP\",\n        \"Precision\",\n        \"Recall\",\n        \"Specificity\",\n        \"F1\",\n        \"Model_Behavior\"\n    ]\n].sort_values(\n    \"F1\",\n    ascending=False\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"COMPLETE PER-CLASS SUMMARY\")\nprint(\"=\" * 80)\n\ndisplay(\n    summary.round(4)\n)\n\n\n# ============================================================\n# BEST CLASSES\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"BEST PERFORMING CLASSES\")\nprint(\"=\" * 80)\n\nbest = df.sort_values(\n    [\"F1\", \"AUC\"],\n    ascending=False\n).head(5)\n\ndisplay(\n    best[\n        [\n            \"Abnormality\",\n            \"AUC\",\n            \"AP\",\n            \"Precision\",\n            \"Recall\",\n            \"F1\"\n        ]\n    ].round(4)\n)\n\n\n# ============================================================\n# WORST CLASSES\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"WORST PERFORMING CLASSES\")\nprint(\"=\" * 80)\n\nworst = df.sort_values(\n    [\"F1\", \"AUC\"],\n    ascending=True\n).head(5)\n\ndisplay(\n    worst[\n        [\n            \"Abnormality\",\n            \"AUC\",\n            \"AP\",\n            \"Precision\",\n            \"Recall\",\n            \"F1\"\n        ]\n    ].round(4)\n)\n\n\n# ============================================================\n# HIGH AUC / LOW F1\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"HIGH AUC BUT LOW F1\")\nprint(\"=\" * 80)\n\nthreshold_problem = df[\n    (df[\"AUC\"] >= 0.65) &\n    (df[\"F1\"] < 0.50)\n].sort_values(\n    \"AUC\",\n    ascending=False\n)\n\nif len(threshold_problem) == 0:\n\n    print(\"None found.\")\n\nelse:\n\n    display(\n        threshold_problem[\n            [\n                \"Abnormality\",\n                \"Threshold\",\n                \"AUC\",\n                \"Precision\",\n                \"Recall\",\n                \"F1\",\n                \"FP\",\n                \"FN\"\n            ]\n        ].round(4)\n    )\n\n\n# ============================================================\n# HIGH RECALL / LOW PRECISION\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"HIGH RECALL / LOW PRECISION\")\nprint(\"=\" * 80)\n\nhigh_recall_low_precision = df[\n    (df[\"Recall\"] >= 0.70) &\n    (df[\"Precision\"] < 0.40)\n].sort_values(\n    \"Recall\",\n    ascending=False\n)\n\nif len(high_recall_low_precision) == 0:\n\n    print(\"None found.\")\n\nelse:\n\n    display(\n        high_recall_low_precision[\n            [\n                \"Abnormality\",\n                \"Threshold\",\n                \"Precision\",\n                \"Recall\",\n                \"F1\",\n                \"FP\",\n                \"FN\"\n            ]\n        ].round(4)\n    )\n\n\n# ============================================================\n# WEAKLY LEARNED CLASSES\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"WEAKLY LEARNED CLASSES\")\nprint(\"=\" * 80)\n\nweak = df[\n    df[\"AUC\"] < 0.55\n].sort_values(\n    \"AUC\"\n)\n\nif len(weak) == 0:\n\n    print(\"None.\")\n\nelse:\n\n    display(\n        weak[\n            [\n                \"Abnormality\",\n                \"AUC\",\n                \"AP\",\n                \"Precision\",\n                \"Recall\",\n                \"F1\",\n                \"TP\",\n                \"FP\",\n                \"FN\"\n            ]\n        ].round(4)\n    )\n\n\n# ============================================================\n# FALSE POSITIVE ANALYSIS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"FALSE POSITIVE BURDEN\")\nprint(\"=\" * 80)\n\nfp_table = df[\n    [\n        \"Abnormality\",\n        \"FP\",\n        \"TN\",\n        \"Precision\",\n        \"Specificity\"\n    ]\n].sort_values(\n    \"FP\",\n    ascending=False\n)\n\ndisplay(\n    fp_table.round(4)\n)\n\n\n# ============================================================\n# FALSE NEGATIVE ANALYSIS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"FALSE NEGATIVE BURDEN\")\nprint(\"=\" * 80)\n\nfn_table = df[\n    [\n        \"Abnormality\",\n        \"FN\",\n        \"TP\",\n        \"Recall\",\n        \"F1\"\n    ]\n].sort_values(\n    \"FN\",\n    ascending=False\n)\n\ndisplay(\n    fn_table.round(4)\n)\n\n\n# ============================================================\n# OVERALL STATISTICS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 71 OVERALL STATISTICS\")\nprint(\"=\" * 80)\n\nprint(\n    f\"Mean AUC       : {df['AUC'].mean():.4f}\"\n)\n\nprint(\n    f\"Mean AP        : {df['AP'].mean():.4f}\"\n)\n\nprint(\n    f\"Mean Precision : {df['Precision'].mean():.4f}\"\n)\n\nprint(\n    f\"Mean Recall    : {df['Recall'].mean():.4f}\"\n)\n\nprint(\n    f\"Mean Specificity : \"\n    f\"{df['Specificity'].mean():.4f}\"\n)\n\nprint(\n    f\"Mean F1        : {df['F1'].mean():.4f}\"\n)\n\n\n# ============================================================\n# BEHAVIOR COUNTS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"MODEL BEHAVIOR COUNTS\")\nprint(\"=\" * 80)\n\nprint(\n    df[\"Model_Behavior\"]\n    .value_counts()\n)\n\n\n# ============================================================\n# SAVE EXTENDED ANALYSIS\n# ============================================================\n\nOUTPUT_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_step71_error_analysis.csv\"\n)\n\ndf.to_csv(\n    OUTPUT_PATH,\n    index=False\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 71 COMPLETE\")\nprint(\"=\" * 80)\n\nprint(\n    \"Saved:\",\n    OUTPUT_PATH\n)\n\nprint(\"\\n✓ Best classes identified\")\nprint(\"✓ Worst classes identified\")\nprint(\"✓ Threshold problems identified\")\nprint(\"✓ High-recall / low-precision classes identified\")\nprint(\"✓ Weak-learning classes identified\")\nprint(\"✓ False positives analyzed\")\nprint(\"✓ False negatives analyzed\")\nprint(\"✓ Training strategy can now be selected\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:09.518214Z","iopub.execute_input":"2026-08-19T05:48:09.518599Z","iopub.status.idle":"2026-08-19T05:48:09.622917Z","shell.execute_reply.started":"2026-08-19T05:48:09.518576Z","shell.execute_reply":"2026-08-19T05:48:09.622308Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 72 — DIRECT VALIDATION PROBABILITY GENERATION\n# ============================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\n\nfrom sklearn.metrics import (\n    roc_auc_score,\n    average_precision_score,\n    precision_score,\n    recall_score,\n    f1_score,\n    confusion_matrix\n)\n\nprint(\"=\" * 80)\nprint(\"STEP 72 — DIRECT VALIDATION PROBABILITY GENERATION\")\nprint(\"=\" * 80)\n\n\n# ============================================================\n# CONFIGURATION\n# ============================================================\n\nCHECKPOINT = (\n    \"/kaggle/working/\"\n    \"best_rsna_knee_model.pth\"\n)\n\nLABEL_COLUMNS = [\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(LABEL_COLUMNS)\n\n\n# ============================================================\n# CHECK REQUIRED VARIABLES\n# ============================================================\n\nrequired_variables = [\n    \"model\",\n    \"val_loader\",\n    \"DEVICE\"\n]\n\nprint(\"\\nChecking required variables...\")\n\nfor variable in required_variables:\n\n    if variable not in globals():\n\n        raise NameError(\n            f\"Required variable '{variable}' \"\n            f\"is not defined.\"\n        )\n\n    print(\n        f\"✓ {variable}\"\n    )\n\n\n# ============================================================\n# CHECK CHECKPOINT\n# ============================================================\n\nif not os.path.exists(CHECKPOINT):\n\n    raise FileNotFoundError(\n        f\"\\nCheckpoint not found:\\n{CHECKPOINT}\"\n    )\n\nprint(\n    \"\\n✓ Checkpoint found:\"\n)\nprint(CHECKPOINT)\n\n\n# ============================================================\n# LOAD BEST MODEL\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"LOADING BEST CHECKPOINT\")\nprint(\"=\" * 80)\n\nmodel = model.to(DEVICE)\n\ntry:\n\n    checkpoint = torch.load(\n        CHECKPOINT,\n        map_location=DEVICE,\n        weights_only=False\n    )\n\nexcept TypeError:\n\n    checkpoint = torch.load(\n        CHECKPOINT,\n        map_location=DEVICE\n    )\n\n\n# ============================================================\n# HANDLE CHECKPOINT FORMAT\n# ============================================================\n\nif isinstance(checkpoint, dict):\n\n    print(\n        \"\\nCheckpoint keys:\"\n    )\n\n    print(\n        list(checkpoint.keys())\n    )\n\n    if \"model_state_dict\" in checkpoint:\n\n        model.load_state_dict(\n            checkpoint[\"model_state_dict\"]\n        )\n\n        print(\n            \"\\n✓ model_state_dict loaded\"\n        )\n\n    elif \"state_dict\" in checkpoint:\n\n        model.load_state_dict(\n            checkpoint[\"state_dict\"]\n        )\n\n        print(\n            \"\\n✓ state_dict loaded\"\n        )\n\n    else:\n\n        # Could itself be a state dictionary\n        try:\n\n            model.load_state_dict(\n                checkpoint\n            )\n\n            print(\n                \"\\n✓ checkpoint loaded directly\"\n            )\n\n        except Exception as e:\n\n            raise RuntimeError(\n                \"Could not determine checkpoint format.\\n\"\n                f\"Original error:\\n{e}\"\n            )\n\nelse:\n\n    model.load_state_dict(\n        checkpoint\n    )\n\n    print(\n        \"\\n✓ checkpoint loaded directly\"\n    )\n\n\n# ============================================================\n# EVALUATION MODE\n# ============================================================\n\nmodel.eval()\n\nprint(\n    \"\\n✓ Model switched to evaluation mode\"\n)\n\n\n# ============================================================\n# GENERATE VALIDATION PREDICTIONS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"GENERATING VALIDATION PREDICTIONS\")\nprint(\"=\" * 80)\n\nall_probabilities = []\nall_labels = []\nall_study_ids = []\n\nwith torch.no_grad():\n\n    for batch_idx, batch in enumerate(\n        val_loader\n    ):\n\n        # ----------------------------------------------------\n        # Handle dictionary batch\n        # ----------------------------------------------------\n\n        if isinstance(batch, dict):\n\n            images = batch[\"images\"]\n            plane_mask = batch[\"plane_mask\"]\n            labels = batch[\"labels\"]\n\n            study_ids = batch.get(\n                \"study_id\",\n                batch.get(\n                    \"study_ids\",\n                    None\n                )\n            )\n\n        # ----------------------------------------------------\n        # Handle tuple/list batch\n        # ----------------------------------------------------\n\n        elif isinstance(batch, (tuple, list)):\n\n            if len(batch) >= 3:\n\n                images = batch[0]\n                plane_mask = batch[1]\n                labels = batch[2]\n\n                if len(batch) >= 4:\n\n                    study_ids = batch[3]\n\n                else:\n\n                    study_ids = None\n\n            else:\n\n                raise ValueError(\n                    \"Unexpected DataLoader batch format.\"\n                )\n\n        else:\n\n            raise ValueError(\n                \"Unsupported DataLoader batch type: \"\n                f\"{type(batch)}\"\n            )\n\n\n        # ----------------------------------------------------\n        # Move tensors\n        # ----------------------------------------------------\n\n        images = images.to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        plane_mask = plane_mask.to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        labels = labels.to(\n            DEVICE,\n            non_blocking=True\n        )\n\n\n        # ----------------------------------------------------\n        # Forward pass\n        # ----------------------------------------------------\n\n        logits = model(\n            images,\n            plane_mask\n        )\n\n\n        # ----------------------------------------------------\n        # Check output\n        # ----------------------------------------------------\n\n        if logits.ndim != 2:\n\n            raise ValueError(\n                \"Unexpected model output shape: \"\n                f\"{tuple(logits.shape)}\"\n            )\n\n        if logits.shape[1] != NUM_CLASSES:\n\n            raise ValueError(\n                f\"Expected {NUM_CLASSES} outputs, \"\n                f\"got {logits.shape[1]}\"\n            )\n\n\n        # ----------------------------------------------------\n        # Convert logits → probabilities\n        # ----------------------------------------------------\n\n        probabilities = torch.sigmoid(\n            logits\n        )\n\n\n        # ----------------------------------------------------\n        # Store\n        # ----------------------------------------------------\n\n        all_probabilities.append(\n            probabilities.detach()\n            .cpu()\n            .numpy()\n        )\n\n        all_labels.append(\n            labels.detach()\n            .cpu()\n            .numpy()\n        )\n\n\n        # ----------------------------------------------------\n        # Study IDs\n        # ----------------------------------------------------\n\n        if study_ids is not None:\n\n            if torch.is_tensor(study_ids):\n\n                study_ids = (\n                    study_ids\n                    .detach()\n                    .cpu()\n                    .numpy()\n                    .tolist()\n                )\n\n            elif isinstance(\n                study_ids,\n                np.ndarray\n            ):\n\n                study_ids = study_ids.tolist()\n\n            elif not isinstance(\n                study_ids,\n                (list, tuple)\n            ):\n\n                study_ids = [\n                    study_ids\n                ]\n\n            all_study_ids.extend(\n                study_ids\n            )\n\n\n        if (\n            batch_idx + 1\n        ) % 5 == 0:\n\n            print(\n                f\"Processed batches: \"\n                f\"{batch_idx + 1}\"\n            )\n\n\n# ============================================================\n# CONCATENATE\n# ============================================================\n\ny_prob = np.concatenate(\n    all_probabilities,\n    axis=0\n)\n\ny_true = np.concatenate(\n    all_labels,\n    axis=0\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"VALIDATION PREDICTIONS READY\")\nprint(\"=\" * 80)\n\nprint(\n    \"y_true shape :\",\n    y_true.shape\n)\n\nprint(\n    \"y_prob shape :\",\n    y_prob.shape\n)\n\nprint(\n    \"Expected     :\",\n    f\"(N, {NUM_CLASSES})\"\n)\n\n\n# ============================================================\n# BASIC VALIDATION\n# ============================================================\n\nif y_true.shape != y_prob.shape:\n\n    raise ValueError(\n        \"y_true and y_prob shapes do not match.\"\n    )\n\nif y_true.shape[1] != NUM_CLASSES:\n\n    raise ValueError(\n        \"Unexpected number of classes.\"\n    )\n\nprint(\n    \"\\n✓ Prediction matrix shape correct\"\n)\n\nprint(\n    \"Probability minimum:\",\n    y_prob.min()\n)\n\nprint(\n    \"Probability maximum:\",\n    y_prob.max()\n)\n\n\n# ============================================================\n# FIND BEST F1 THRESHOLD\n# ============================================================\n\ndef find_best_threshold(\n    y_true_class,\n    y_prob_class\n):\n\n    best_threshold = 0.50\n    best_f1 = -1.0\n\n    thresholds = np.linspace(\n        0.05,\n        0.95,\n        181\n    )\n\n    for threshold in thresholds:\n\n        predictions = (\n            y_prob_class >= threshold\n        ).astype(int)\n\n        score = f1_score(\n            y_true_class,\n            predictions,\n            zero_division=0\n        )\n\n        if score > best_f1:\n\n            best_f1 = score\n            best_threshold = threshold\n\n    return (\n        best_threshold,\n        best_f1\n    )\n\n\n# ============================================================\n# METRIC CALCULATION\n# ============================================================\n\nrows = []\n\n\nfor class_idx, label in enumerate(\n    LABEL_COLUMNS\n):\n\n    true = y_true[\n        :, class_idx\n    ].astype(int)\n\n    prob = y_prob[\n        :, class_idx\n    ]\n\n\n    # --------------------------------------------------------\n    # AUC\n    # --------------------------------------------------------\n\n    if len(np.unique(true)) == 2:\n\n        auc = roc_auc_score(\n            true,\n            prob\n        )\n\n    else:\n\n        auc = np.nan\n\n\n    # --------------------------------------------------------\n    # AP\n    # --------------------------------------------------------\n\n    try:\n\n        ap = average_precision_score(\n            true,\n            prob\n        )\n\n    except:\n\n        ap = np.nan\n\n\n    # --------------------------------------------------------\n    # DEFAULT 0.50\n    # --------------------------------------------------------\n\n    pred_050 = (\n        prob >= 0.50\n    ).astype(int)\n\n\n    # --------------------------------------------------------\n    # 0.60\n    # --------------------------------------------------------\n\n    pred_060 = (\n        prob >= 0.60\n    ).astype(int)\n\n\n    # --------------------------------------------------------\n    # 0.70\n    # --------------------------------------------------------\n\n    pred_070 = (\n        prob >= 0.70\n    ).astype(int)\n\n\n    # --------------------------------------------------------\n    # OPTIMAL\n    # --------------------------------------------------------\n\n    optimal_threshold, optimal_f1 = (\n        find_best_threshold(\n            true,\n            prob\n        )\n    )\n\n    pred_opt = (\n        prob >= optimal_threshold\n    ).astype(int)\n\n\n    # --------------------------------------------------------\n    # METRIC HELPER\n    # --------------------------------------------------------\n\n    def metrics(\n        true,\n        pred\n    ):\n\n        tn, fp, fn, tp = confusion_matrix(\n            true,\n            pred,\n            labels=[0, 1]\n        ).ravel()\n\n        precision = precision_score(\n            true,\n            pred,\n            zero_division=0\n        )\n\n        recall = recall_score(\n            true,\n            pred,\n            zero_division=0\n        )\n\n        f1 = f1_score(\n            true,\n            pred,\n            zero_division=0\n        )\n\n        specificity = (\n            tn / (tn + fp)\n            if (tn + fp) > 0\n            else np.nan\n        )\n\n        return (\n            tn,\n            fp,\n            fn,\n            tp,\n            precision,\n            recall,\n            specificity,\n            f1\n        )\n\n\n    d = metrics(\n        true,\n        pred_050\n    )\n\n    o = metrics(\n        true,\n        pred_opt\n    )\n\n    s60 = metrics(\n        true,\n        pred_060\n    )\n\n    s70 = metrics(\n        true,\n        pred_070\n    )\n\n\n    rows.append({\n\n        \"Abnormality\": label,\n\n        \"Actual_Positive\":\n            int(true.sum()),\n\n        \"Actual_Negative\":\n            int(len(true) - true.sum()),\n\n        \"AUC\": auc,\n\n        \"AP\": ap,\n\n        \"Optimal_Threshold\":\n            optimal_threshold,\n\n        # ------------------------------\n        # DEFAULT\n        # ------------------------------\n\n        \"Default_TP\": d[3],\n        \"Default_FP\": d[1],\n        \"Default_FN\": d[2],\n        \"Default_TN\": d[0],\n\n        \"Default_Precision\": d[4],\n        \"Default_Recall\": d[5],\n        \"Default_Specificity\": d[6],\n        \"Default_F1\": d[7],\n\n        # ------------------------------\n        # OPTIMAL\n        # ------------------------------\n\n        \"Optimal_TP\": o[3],\n        \"Optimal_FP\": o[1],\n        \"Optimal_FN\": o[2],\n        \"Optimal_TN\": o[0],\n\n        \"Optimal_Precision\": o[4],\n        \"Optimal_Recall\": o[5],\n        \"Optimal_Specificity\": o[6],\n        \"Optimal_F1\": o[7],\n\n        # ------------------------------\n        # 0.60\n        # ------------------------------\n\n        \"Threshold_060_Precision\":\n            s60[4],\n\n        \"Threshold_060_Recall\":\n            s60[5],\n\n        \"Threshold_060_Specificity\":\n            s60[6],\n\n        \"Threshold_060_F1\":\n            s60[7],\n\n        # ------------------------------\n        # 0.70\n        # ------------------------------\n\n        \"Threshold_070_Precision\":\n            s70[4],\n\n        \"Threshold_070_Recall\":\n            s70[5],\n\n        \"Threshold_070_Specificity\":\n            s70[6],\n\n        \"Threshold_070_F1\":\n            s70[7]\n    })\n\n\ndiagnostic_df = pd.DataFrame(\n    rows\n)\n\n\n# ============================================================\n# DISPLAY RESULTS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 72 — COMPLETE THRESHOLD ANALYSIS\")\nprint(\"=\" * 80)\n\ndisplay(\n    diagnostic_df.round(4)\n)\n\n\n# ============================================================\n# OVERALL COMPARISON\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"THRESHOLD COMPARISON\")\nprint(\"=\" * 80)\n\nprint(\n    \"Default 0.50\"\n)\n\nprint(\n    \"  Precision :\",\n    diagnostic_df[\n        \"Default_Precision\"\n    ].mean()\n)\n\nprint(\n    \"  Recall    :\",\n    diagnostic_df[\n        \"Default_Recall\"\n    ].mean()\n)\n\nprint(\n    \"  Specificity:\",\n    diagnostic_df[\n        \"Default_Specificity\"\n    ].mean()\n)\n\nprint(\n    \"  F1        :\",\n    diagnostic_df[\n        \"Default_F1\"\n    ].mean()\n)\n\n\nprint(\n    \"\\nOptimal threshold\"\n)\n\nprint(\n    \"  Precision :\",\n    diagnostic_df[\n        \"Optimal_Precision\"\n    ].mean()\n)\n\nprint(\n    \"  Recall    :\",\n    diagnostic_df[\n        \"Optimal_Recall\"\n    ].mean()\n)\n\nprint(\n    \"  Specificity:\",\n    diagnostic_df[\n        \"Optimal_Specificity\"\n    ].mean()\n)\n\nprint(\n    \"  F1        :\",\n    diagnostic_df[\n        \"Optimal_F1\"\n    ].mean()\n)\n\n\nprint(\n    \"\\nThreshold 0.60\"\n)\n\nprint(\n    \"  Precision :\",\n    diagnostic_df[\n        \"Threshold_060_Precision\"\n    ].mean()\n)\n\nprint(\n    \"  Recall    :\",\n    diagnostic_df[\n        \"Threshold_060_Recall\"\n    ].mean()\n)\n\nprint(\n    \"  Specificity:\",\n    diagnostic_df[\n        \"Threshold_060_Specificity\"\n    ].mean()\n)\n\nprint(\n    \"  F1        :\",\n    diagnostic_df[\n        \"Threshold_060_F1\"\n    ].mean()\n)\n\n\nprint(\n    \"\\nThreshold 0.70\"\n)\n\nprint(\n    \"  Precision :\",\n    diagnostic_df[\n        \"Threshold_070_Precision\"\n    ].mean()\n)\n\nprint(\n    \"  Recall    :\",\n    diagnostic_df[\n        \"Threshold_070_Recall\"\n    ].mean()\n)\n\nprint(\n    \"  Specificity:\",\n    diagnostic_df[\n        \"Threshold_070_Specificity\"\n    ].mean()\n)\n\nprint(\n    \"  F1        :\",\n    diagnostic_df[\n        \"Threshold_070_F1\"\n    ].mean()\n)\n\n\n# ============================================================\n# BEST THRESHOLD PER CLASS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"OPTIMAL THRESHOLDS\")\nprint(\"=\" * 80)\n\ndisplay(\n    diagnostic_df[\n        [\n            \"Abnormality\",\n            \"AUC\",\n            \"Optimal_Threshold\",\n            \"Optimal_Precision\",\n            \"Optimal_Recall\",\n            \"Optimal_Specificity\",\n            \"Optimal_F1\"\n        ]\n    ].round(4)\n)\n\n\n# ============================================================\n# SAVE\n# ============================================================\n\nOUTPUT_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_step72_threshold_diagnostic.csv\"\n)\n\ndiagnostic_df.to_csv(\n    OUTPUT_PATH,\n    index=False\n)\n\n\n# ============================================================\n# ALSO SAVE RAW VALIDATION PROBABILITIES\n# ============================================================\n\nprobability_df = pd.DataFrame(\n    y_prob,\n    columns=[\n        f\"Prob_{label}\"\n        for label in LABEL_COLUMNS\n    ]\n)\n\nlabel_df = pd.DataFrame(\n    y_true,\n    columns=[\n        f\"True_{label}\"\n        for label in LABEL_COLUMNS\n    ]\n)\n\nraw_output = pd.concat(\n    [\n        probability_df,\n        label_df\n    ],\n    axis=1\n)\n\nRAW_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_validation_probabilities.csv\"\n)\n\nraw_output.to_csv(\n    RAW_PATH,\n    index=False\n)\n\n\n# ============================================================\n# FINAL\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 72 COMPLETE\")\nprint(\"=\" * 80)\n\nprint(\n    \"Diagnostic file:\",\n    OUTPUT_PATH\n)\n\nprint(\n    \"Raw probabilities:\",\n    RAW_PATH\n)\n\nprint(\"\\n✓ Best checkpoint loaded\")\nprint(\"✓ Validation probabilities regenerated\")\nprint(\"✓ No CSV probability-column guessing used\")\nprint(\"✓ AUC calculated\")\nprint(\"✓ AP calculated\")\nprint(\"✓ Default threshold evaluated\")\nprint(\"✓ Optimal thresholds calculated\")\nprint(\"✓ 0.60 threshold evaluated\")\nprint(\"✓ 0.70 threshold evaluated\")\nprint(\"✓ Specificity verified\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:09.624608Z","iopub.execute_input":"2026-08-19T05:48:09.625063Z","iopub.status.idle":"2026-08-19T05:48:19.566526Z","shell.execute_reply.started":"2026-08-19T05:48:09.625039Z","shell.execute_reply":"2026-08-19T05:48:19.565931Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 73 — THRESHOLD COMPARISON & MODEL DECISION\n# ============================================================\n\nimport pandas as pd\nimport numpy as np\n\nPATH = \"/kaggle/working/rsna_knee_step72_threshold_diagnostic.csv\"\n\ndf72 = pd.read_csv(PATH)\n\nprint(\"=\" * 80)\nprint(\"STEP 73 — THRESHOLD COMPARISON\")\nprint(\"=\" * 80)\n\n# ------------------------------------------------------------\n# OVERALL COMPARISON\n# ------------------------------------------------------------\n\ncomparison = pd.DataFrame({\n    \"Threshold\": [\n        \"Default 0.50\",\n        \"Optimized\",\n        \"Fixed 0.60\",\n        \"Fixed 0.70\"\n    ],\n\n    \"Precision\": [\n        df72[\"Default_Precision\"].mean(),\n        df72[\"Optimal_Precision\"].mean(),\n        df72[\"Threshold_060_Precision\"].mean(),\n        df72[\"Threshold_070_Precision\"].mean()\n    ],\n\n    \"Recall\": [\n        df72[\"Default_Recall\"].mean(),\n        df72[\"Optimal_Recall\"].mean(),\n        df72[\"Threshold_060_Recall\"].mean(),\n        df72[\"Threshold_070_Recall\"].mean()\n    ],\n\n    \"Specificity\": [\n        df72[\"Default_Specificity\"].mean(),\n        df72[\"Optimal_Specificity\"].mean(),\n        df72[\"Threshold_060_Specificity\"].mean(),\n        df72[\"Threshold_070_Specificity\"].mean()\n    ],\n\n    \"F1\": [\n        df72[\"Default_F1\"].mean(),\n        df72[\"Optimal_F1\"].mean(),\n        df72[\"Threshold_060_F1\"].mean(),\n        df72[\"Threshold_070_F1\"].mean()\n    ]\n})\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"OVERALL THRESHOLD COMPARISON\")\nprint(\"=\" * 80)\n\ndisplay(comparison.round(4))\n\n\n# ------------------------------------------------------------\n# BEST THRESHOLD BY MACRO F1\n# ------------------------------------------------------------\n\nbest_row = comparison.loc[\n    comparison[\"F1\"].idxmax()\n]\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"BEST OVERALL THRESHOLD\")\nprint(\"=\" * 80)\n\nprint(\n    \"Threshold:\",\n    best_row[\"Threshold\"]\n)\n\nprint(\n    \"Precision:\",\n    round(best_row[\"Precision\"], 4)\n)\n\nprint(\n    \"Recall:\",\n    round(best_row[\"Recall\"], 4)\n)\n\nprint(\n    \"Specificity:\",\n    round(best_row[\"Specificity\"], 4)\n)\n\nprint(\n    \"F1:\",\n    round(best_row[\"F1\"], 4)\n)\n\n\n# ------------------------------------------------------------\n# CLASS-WISE THRESHOLD WINNER\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CLASS-WISE BEST THRESHOLD\")\nprint(\"=\" * 80)\n\nthreshold_columns = {\n    \"0.50\": \"Default_F1\",\n    \"Optimized\": \"Optimal_F1\",\n    \"0.60\": \"Threshold_060_F1\",\n    \"0.70\": \"Threshold_070_F1\"\n}\n\nwinner_rows = []\n\nfor _, row in df72.iterrows():\n\n    scores = {\n        name: row[column]\n        for name, column in threshold_columns.items()\n    }\n\n    best = max(\n        scores,\n        key=scores.get\n    )\n\n    winner_rows.append({\n        \"Abnormality\": row[\"Abnormality\"],\n        \"Best_Threshold\": best,\n        \"Best_F1\": scores[best],\n        \"AUC\": row[\"AUC\"]\n    })\n\nwinner_df = pd.DataFrame(\n    winner_rows\n)\n\ndisplay(\n    winner_df.round(4)\n)\n\n\n# ------------------------------------------------------------\n# THRESHOLD WIN COUNTS\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"THRESHOLD WIN COUNTS\")\nprint(\"=\" * 80)\n\nprint(\n    winner_df[\"Best_Threshold\"]\n    .value_counts()\n)\n\n\n# ------------------------------------------------------------\n# AUC VS F1\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"AUC VS F1\")\nprint(\"=\" * 80)\n\ndisplay(\n    df72[\n        [\n            \"Abnormality\",\n            \"AUC\",\n            \"Default_F1\",\n            \"Optimal_F1\",\n            \"Threshold_060_F1\",\n            \"Threshold_070_F1\"\n        ]\n    ]\n    .sort_values(\"AUC\", ascending=False)\n    .round(4)\n)\n\n\n# ------------------------------------------------------------\n# DECISION LOGIC\n# ------------------------------------------------------------\n\nmean_auc = df72[\"AUC\"].mean()\n\ndefault_f1 = df72[\"Default_F1\"].mean()\noptimal_f1 = df72[\"Optimal_F1\"].mean()\nf1_060 = df72[\"Threshold_060_F1\"].mean()\nf1_070 = df72[\"Threshold_070_F1\"].mean()\n\nbest_f1 = max(\n    default_f1,\n    optimal_f1,\n    f1_060,\n    f1_070\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"MODEL DECISION\")\nprint(\"=\" * 80)\n\nprint(\n    f\"Mean AUC : {mean_auc:.4f}\"\n)\n\nprint(\n    f\"Best threshold F1 : {best_f1:.4f}\"\n)\n\n\nif mean_auc < 0.60:\n\n    print(\"\"\"\nDECISION:\nThe model's ranking ability is still weak.\n\nThreshold tuning alone is NOT sufficient.\n\nNext priority:\n1. Improve training stability\n2. Improve augmentation\n3. Address extreme class imbalance\n4. Use stronger validation methodology\n5. Then retrain\n\"\"\")\n\nelif best_f1 > default_f1 + 0.08:\n\n    print(\"\"\"\nDECISION:\nThreshold selection has a substantial effect.\n\nThe model has useful ranking information,\nbut the default decision threshold is poor.\n\nNext priority:\n1. Calibrate thresholds\n2. Preserve the trained model\n3. Evaluate with class-specific thresholds\n\"\"\")\n\nelse:\n\n    print(\"\"\"\nDECISION:\nThreshold tuning provides limited improvement.\n\nThe main limitation is probably representation/\nlearning rather than threshold selection.\n\nNext priority:\n1. Improve model training\n2. Add stronger regularization\n3. Improve data augmentation\n4. Re-train\n\"\"\")\n\n\n# ------------------------------------------------------------\n# SAVE\n# ------------------------------------------------------------\n\nOUT = (\n    \"/kaggle/working/\"\n    \"rsna_knee_step73_threshold_comparison.csv\"\n)\n\ncomparison.to_csv(\n    OUT,\n    index=False\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 73 COMPLETE\")\nprint(\"=\" * 80)\n\nprint(\"Saved:\", OUT)\nprint(\"✓ Thresholds compared\")\nprint(\"✓ Class-wise winners identified\")\nprint(\"✓ AUC/F1 relationship examined\")\nprint(\"✓ Training decision generated\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:19.567601Z","iopub.execute_input":"2026-08-19T05:48:19.567936Z","iopub.status.idle":"2026-08-19T05:48:19.616268Z","shell.execute_reply.started":"2026-08-19T05:48:19.567911Z","shell.execute_reply":"2026-08-19T05:48:19.615668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 74 — CONTROLLED FINE-TUNING\n# ============================================================\n# Strategy:\n#   Stage 1 -> Freeze ResNet backbone, train prediction head\n#   Stage 2 -> Unfreeze layer4, fine-tune with very low LR\n#\n# Designed for the current RSNA Knee pipeline:\n#   Input  : (B, 3, 12, 256, 256)\n#   Output : (B, 12)\n# ============================================================\n\nimport os\nimport copy\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom sklearn.metrics import roc_auc_score, f1_score, precision_score, recall_score\n\nprint(\"=\" * 80)\nprint(\"STEP 74 — CONTROLLED FINE-TUNING\")\nprint(\"=\" * 80)\n\n# ============================================================\n# 1. CHECK REQUIRED VARIABLES\n# ============================================================\n\nrequired = [\n    \"model\",\n    \"train_loader\",\n    \"val_loader\",\n    \"criterion\",\n    \"DEVICE\",\n    \"LABEL_COLUMNS\"\n]\n\nfor variable in required:\n    if variable not in globals():\n        raise NameError(\n            f\"Required variable '{variable}' is not defined. \"\n            f\"Run the dataset/model preparation cells first.\"\n        )\n\nprint(\"✓ Model found\")\nprint(\"✓ Train loader found\")\nprint(\"✓ Validation loader found\")\nprint(\"✓ Loss function found\")\nprint(\"✓ DEVICE found\")\nprint(\"✓ LABEL_COLUMNS found\")\n\nmodel = model.to(DEVICE)\n\nNUM_CLASSES = len(LABEL_COLUMNS)\n\nprint(\"\\nNumber of classes:\", NUM_CLASSES)\nprint(\"Device:\", DEVICE)\n\n\n# ============================================================\n# 2. IDENTIFY BACKBONE\n# ============================================================\n\nif not hasattr(model, \"backbone\"):\n    raise AttributeError(\n        \"Model does not contain a 'backbone' attribute.\"\n    )\n\nbackbone = model.backbone\n\nprint(\"\\nBackbone:\")\nprint(type(backbone).__name__)\n\n\n# ============================================================\n# 3. STAGE 1 — FREEZE COMPLETE BACKBONE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STAGE 1 — FROZEN BACKBONE\")\nprint(\"=\" * 80)\n\nfor param in backbone.parameters():\n    param.requires_grad = False\n\n# Keep classifier/fusion/attention trainable\nfor name, module in model.named_children():\n\n    if name != \"backbone\":\n\n        for param in module.parameters():\n            param.requires_grad = True\n\n\ntrainable_stage1 = sum(\n    p.numel()\n    for p in model.parameters()\n    if p.requires_grad\n)\n\nfrozen_stage1 = sum(\n    p.numel()\n    for p in model.parameters()\n    if not p.requires_grad\n)\n\nprint(\"Trainable parameters :\", trainable_stage1)\nprint(\"Frozen parameters    :\", frozen_stage1)\n\n\n# ============================================================\n# 4. CREATE STAGE 1 OPTIMIZER\n# ============================================================\n\nSTAGE1_LR = 3e-4\nWEIGHT_DECAY = 1e-4\n\noptimizer_stage1 = torch.optim.AdamW(\n    [\n        {\n            \"params\": [\n                p for p in model.parameters()\n                if p.requires_grad\n            ],\n            \"lr\": STAGE1_LR\n        }\n    ],\n    weight_decay=WEIGHT_DECAY\n)\n\nscheduler_stage1 = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer_stage1,\n    T_max=8,\n    eta_min=5e-5\n)\n\nprint(\"\\nStage 1 LR       :\", STAGE1_LR)\nprint(\"Weight decay     :\", WEIGHT_DECAY)\nprint(\"Stage 1 epochs   : 8\")\n\n\n# ============================================================\n# 5. METRIC FUNCTION\n# ============================================================\n\ndef calculate_metrics(\n    y_true,\n    y_prob,\n    threshold=0.5\n):\n\n    y_true = np.asarray(y_true)\n    y_prob = np.asarray(y_prob)\n\n    y_pred = (\n        y_prob >= threshold\n    ).astype(np.int32)\n\n    auc_values = []\n\n    for c in range(y_true.shape[1]):\n\n        unique_values = np.unique(\n            y_true[:, c]\n        )\n\n        if len(unique_values) < 2:\n            continue\n\n        try:\n            auc = roc_auc_score(\n                y_true[:, c],\n                y_prob[:, c]\n            )\n\n            auc_values.append(auc)\n\n        except Exception:\n            pass\n\n    mean_auc = (\n        float(np.mean(auc_values))\n        if len(auc_values) > 0\n        else np.nan\n    )\n\n    precision = precision_score(\n        y_true,\n        y_pred,\n        average=\"macro\",\n        zero_division=0\n    )\n\n    recall = recall_score(\n        y_true,\n        y_pred,\n        average=\"macro\",\n        zero_division=0\n    )\n\n    f1 = f1_score(\n        y_true,\n        y_pred,\n        average=\"macro\",\n        zero_division=0\n    )\n\n    return {\n        \"auc\": mean_auc,\n        \"precision\": precision,\n        \"recall\": recall,\n        \"f1\": f1\n    }\n\n\n# ============================================================\n# 6. TRAINING FUNCTION\n# ============================================================\n\ndef train_one_epoch(\n    model,\n    loader,\n    optimizer,\n    criterion\n):\n\n    model.train()\n\n    running_loss = 0.0\n\n    all_probs = []\n    all_labels = []\n\n    for batch in loader:\n\n        images = batch[\"images\"].to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        plane_mask = batch[\"plane_mask\"].to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        labels = batch[\"labels\"].to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        optimizer.zero_grad(\n            set_to_none=True\n        )\n\n        logits = model(\n            images,\n            plane_mask\n        )\n\n        loss = criterion(\n            logits,\n            labels\n        )\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 += (\n            loss.item()\n            * images.size(0)\n        )\n\n        probs = torch.sigmoid(\n            logits\n        )\n\n        all_probs.append(\n            probs.detach().cpu().numpy()\n        )\n\n        all_labels.append(\n            labels.detach().cpu().numpy()\n        )\n\n    epoch_loss = (\n        running_loss /\n        len(loader.dataset)\n    )\n\n    all_probs = np.concatenate(\n        all_probs,\n        axis=0\n    )\n\n    all_labels = np.concatenate(\n        all_labels,\n        axis=0\n    )\n\n    metrics = calculate_metrics(\n        all_labels,\n        all_probs,\n        threshold=0.5\n    )\n\n    return epoch_loss, metrics\n\n\n# ============================================================\n# 7. VALIDATION FUNCTION\n# ============================================================\n\n@torch.no_grad()\ndef validate_one_epoch(\n    model,\n    loader,\n    criterion\n):\n\n    model.eval()\n\n    running_loss = 0.0\n\n    all_probs = []\n    all_labels = []\n\n    for batch in loader:\n\n        images = batch[\"images\"].to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        plane_mask = batch[\"plane_mask\"].to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        labels = batch[\"labels\"].to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        logits = model(\n            images,\n            plane_mask\n        )\n\n        loss = criterion(\n            logits,\n            labels\n        )\n\n        running_loss += (\n            loss.item()\n            * images.size(0)\n        )\n\n        probs = torch.sigmoid(\n            logits\n        )\n\n        all_probs.append(\n            probs.cpu().numpy()\n        )\n\n        all_labels.append(\n            labels.cpu().numpy()\n        )\n\n    epoch_loss = (\n        running_loss /\n        len(loader.dataset)\n    )\n\n    all_probs = np.concatenate(\n        all_probs,\n        axis=0\n    )\n\n    all_labels = np.concatenate(\n        all_labels,\n        axis=0\n    )\n\n    metrics = calculate_metrics(\n        all_labels,\n        all_probs,\n        threshold=0.5\n    )\n\n    return (\n        epoch_loss,\n        metrics,\n        all_probs,\n        all_labels\n    )\n\n\n# ============================================================\n# 8. STAGE 1 TRAINING\n# ============================================================\n\nSTAGE1_EPOCHS = 8\n\nbest_stage1_auc = -np.inf\nbest_stage1_state = None\n\nhistory_stage1 = []\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STARTING STAGE 1\")\nprint(\"=\" * 80)\n\nfor epoch in range(1, STAGE1_EPOCHS + 1):\n\n    train_loss, train_metrics = train_one_epoch(\n        model,\n        train_loader,\n        optimizer_stage1,\n        criterion\n    )\n\n    val_loss, val_metrics, _, _ = validate_one_epoch(\n        model,\n        val_loader,\n        criterion\n    )\n\n    scheduler_stage1.step()\n\n    current_lr = (\n        optimizer_stage1.param_groups[0][\"lr\"]\n    )\n\n    print(f\"\\nEpoch {epoch:02d}/{STAGE1_EPOCHS}\")\n    print(\"-\" * 80)\n\n    print(\n        f\"Train Loss : {train_loss:.4f}\"\n    )\n\n    print(\n        f\"Val Loss   : {val_loss:.4f}\"\n    )\n\n    print(\n        f\"Train AUC  : {train_metrics['auc']:.4f}\"\n    )\n\n    print(\n        f\"Val AUC    : {val_metrics['auc']:.4f}\"\n    )\n\n    print(\n        f\"Train F1   : {train_metrics['f1']:.4f}\"\n    )\n\n    print(\n        f\"Val F1     : {val_metrics['f1']:.4f}\"\n    )\n\n    print(\n        f\"Val Recall : {val_metrics['recall']:.4f}\"\n    )\n\n    print(\n        f\"LR         : {current_lr:.8f}\"\n    )\n\n    history_stage1.append({\n        \"stage\": 1,\n        \"epoch\": epoch,\n        \"train_loss\": train_loss,\n        \"val_loss\": val_loss,\n        \"train_auc\": train_metrics[\"auc\"],\n        \"val_auc\": val_metrics[\"auc\"],\n        \"train_f1\": train_metrics[\"f1\"],\n        \"val_f1\": val_metrics[\"f1\"],\n        \"val_precision\": val_metrics[\"precision\"],\n        \"val_recall\": val_metrics[\"recall\"],\n        \"lr\": current_lr\n    })\n\n    if (\n        not np.isnan(val_metrics[\"auc\"])\n        and\n        val_metrics[\"auc\"] > best_stage1_auc\n    ):\n\n        best_stage1_auc = (\n            val_metrics[\"auc\"]\n        )\n\n        best_stage1_state = copy.deepcopy(\n            model.state_dict()\n        )\n\n        print(\"✓ NEW STAGE 1 BEST MODEL\")\n\n\n# ============================================================\n# 9. RESTORE BEST STAGE 1 MODEL\n# ============================================================\n\nif best_stage1_state is not None:\n\n    model.load_state_dict(\n        best_stage1_state\n    )\n\n    print(\n        \"\\n✓ Best Stage 1 weights restored\"\n    )\n\nelse:\n\n    print(\n        \"\\n⚠ No valid Stage 1 AUC found\"\n    )\n\n\n# ============================================================\n# 10. STAGE 2 — UNFREEZE ONLY LAYER4\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STAGE 2 — PARTIAL BACKBONE FINE-TUNING\")\nprint(\"=\" * 80)\n\n# First freeze everything\nfor param in backbone.parameters():\n    param.requires_grad = False\n\n# Unfreeze layer4 only\nif hasattr(backbone, \"layer4\"):\n\n    for param in backbone.layer4.parameters():\n        param.requires_grad = True\n\n    print(\"✓ ResNet layer4 unfrozen\")\n\nelse:\n\n    raise AttributeError(\n        \"Backbone does not contain layer4.\"\n    )\n\n\n# Keep head trainable\nfor name, module in model.named_children():\n\n    if name != \"backbone\":\n\n        for param in module.parameters():\n            param.requires_grad = True\n\n\nbackbone_layer4_params = [\n    p\n    for p in backbone.layer4.parameters()\n    if p.requires_grad\n]\n\nhead_params = [\n    p\n    for name, module in model.named_children()\n    if name != \"backbone\"\n    for p in module.parameters()\n    if p.requires_grad\n]\n\n\nprint(\n    \"Layer4 trainable parameters:\",\n    sum(\n        p.numel()\n        for p in backbone_layer4_params\n    )\n)\n\nprint(\n    \"Head trainable parameters:\",\n    sum(\n        p.numel()\n        for p in head_params\n    )\n)\n\n\n# ============================================================\n# 11. STAGE 2 OPTIMIZER\n# ============================================================\n\nSTAGE2_BACKBONE_LR = 1e-6\nSTAGE2_HEAD_LR = 1e-4\n\noptimizer_stage2 = torch.optim.AdamW(\n    [\n        {\n            \"params\": backbone_layer4_params,\n            \"lr\": STAGE2_BACKBONE_LR\n        },\n        {\n            \"params\": head_params,\n            \"lr\": STAGE2_HEAD_LR\n        }\n    ],\n    weight_decay=WEIGHT_DECAY\n)\n\nscheduler_stage2 = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer_stage2,\n    T_max=12,\n    eta_min=1e-7\n)\n\nprint(\"\\nStage 2 configuration:\")\nprint(\n    \"Layer4 LR :\", STAGE2_BACKBONE_LR\n)\n\nprint(\n    \"Head LR   :\", STAGE2_HEAD_LR\n)\n\nprint(\n    \"Epochs    : 12\"\n)\n\n\n# ============================================================\n# 12. STAGE 2 TRAINING\n# ============================================================\n\nSTAGE2_EPOCHS = 12\n\nbest_auc = best_stage1_auc\nbest_state = copy.deepcopy(\n    model.state_dict()\n)\n\nhistory_stage2 = []\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STARTING STAGE 2\")\nprint(\"=\" * 80)\n\nfor epoch in range(1, STAGE2_EPOCHS + 1):\n\n    train_loss, train_metrics = train_one_epoch(\n        model,\n        train_loader,\n        optimizer_stage2,\n        criterion\n    )\n\n    val_loss, val_metrics, _, _ = validate_one_epoch(\n        model,\n        val_loader,\n        criterion\n    )\n\n    scheduler_stage2.step()\n\n    lr_backbone = (\n        optimizer_stage2.param_groups[0][\"lr\"]\n    )\n\n    lr_head = (\n        optimizer_stage2.param_groups[1][\"lr\"]\n    )\n\n    print(f\"\\nEpoch {epoch:02d}/{STAGE2_EPOCHS}\")\n    print(\"-\" * 80)\n\n    print(\n        f\"Train Loss : {train_loss:.4f}\"\n    )\n\n    print(\n        f\"Val Loss   : {val_loss:.4f}\"\n    )\n\n    print(\n        f\"Train AUC  : {train_metrics['auc']:.4f}\"\n    )\n\n    print(\n        f\"Val AUC    : {val_metrics['auc']:.4f}\"\n    )\n\n    print(\n        f\"Train F1   : {train_metrics['f1']:.4f}\"\n    )\n\n    print(\n        f\"Val F1     : {val_metrics['f1']:.4f}\"\n    )\n\n    print(\n        f\"Val Recall : {val_metrics['recall']:.4f}\"\n    )\n\n    print(\n        f\"LR Layer4  : {lr_backbone:.8f}\"\n    )\n\n    print(\n        f\"LR Head    : {lr_head:.8f}\"\n    )\n\n    history_stage2.append({\n        \"stage\": 2,\n        \"epoch\": epoch,\n        \"train_loss\": train_loss,\n        \"val_loss\": val_loss,\n        \"train_auc\": train_metrics[\"auc\"],\n        \"val_auc\": val_metrics[\"auc\"],\n        \"train_f1\": train_metrics[\"f1\"],\n        \"val_f1\": val_metrics[\"f1\"],\n        \"val_precision\": val_metrics[\"precision\"],\n        \"val_recall\": val_metrics[\"recall\"],\n        \"lr_layer4\": lr_backbone,\n        \"lr_head\": lr_head\n    })\n\n    if (\n        not np.isnan(val_metrics[\"auc\"])\n        and\n        val_metrics[\"auc\"] > best_auc\n    ):\n\n        best_auc = (\n            val_metrics[\"auc\"]\n        )\n\n        best_state = copy.deepcopy(\n            model.state_dict()\n        )\n\n        print(\n            \"✓ NEW BEST MODEL\"\n        )\n\n\n# ============================================================\n# 13. RESTORE BEST MODEL\n# ============================================================\n\nmodel.load_state_dict(\n    best_state\n)\n\nprint(\"\\n✓ Best controlled-fine-tuning weights restored\")\n\n\n# ============================================================\n# 14. FINAL VALIDATION\n# ============================================================\n\n(\n    final_val_loss,\n    final_metrics,\n    final_probs,\n    final_labels\n) = validate_one_epoch(\n    model,\n    val_loader,\n    criterion\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 74 FINAL VALIDATION\")\nprint(\"=\" * 80)\n\nprint(\n    f\"Validation Loss : {final_val_loss:.4f}\"\n)\n\nprint(\n    f\"Validation AUC  : {final_metrics['auc']:.4f}\"\n)\n\nprint(\n    f\"Validation F1   : {final_metrics['f1']:.4f}\"\n)\n\nprint(\n    f\"Precision       : {final_metrics['precision']:.4f}\"\n)\n\nprint(\n    f\"Recall          : {final_metrics['recall']:.4f}\"\n)\n\n\n# ============================================================\n# 15. SAVE CHECKPOINT\n# ============================================================\n\nSTEP74_CHECKPOINT = (\n    \"/kaggle/working/\"\n    \"best_rsna_knee_step74_model.pth\"\n)\n\ntorch.save(\n    {\n        \"model_state_dict\": model.state_dict(),\n        \"best_auc\": best_auc,\n        \"label_columns\": LABEL_COLUMNS,\n        \"num_classes\": NUM_CLASSES\n    },\n    STEP74_CHECKPOINT\n)\n\nprint(\n    \"\\nCheckpoint saved:\"\n)\n\nprint(\n    STEP74_CHECKPOINT\n)\n\n\n# ============================================================\n# 16. SAVE HISTORY\n# ============================================================\n\nhistory = pd.DataFrame(\n    history_stage1 +\n    history_stage2\n)\n\nSTEP74_HISTORY = (\n    \"/kaggle/working/\"\n    \"rsna_knee_step74_training_history.csv\"\n)\n\nhistory.to_csv(\n    STEP74_HISTORY,\n    index=False\n)\n\nprint(\n    \"History saved:\"\n)\n\nprint(\n    STEP74_HISTORY\n)\n\n\n# ============================================================\n# 17. SAVE VALIDATION PROBABILITIES\n# ============================================================\n\nprobability_data = {\n    \"StudyIndex\": np.arange(\n        len(final_probs)\n    )\n}\n\nfor i, label in enumerate(\n    LABEL_COLUMNS\n):\n\n    probability_data[\n        f\"Prob_{label}\"\n    ] = final_probs[:, i]\n\n    probability_data[\n        f\"True_{label}\"\n    ] = final_labels[:, i]\n\nstep74_probability_df = pd.DataFrame(\n    probability_data\n)\n\nSTEP74_PROBS = (\n    \"/kaggle/working/\"\n    \"rsna_knee_step74_validation_probabilities.csv\"\n)\n\nstep74_probability_df.to_csv(\n    STEP74_PROBS,\n    index=False\n)\n\nprint(\n    \"Validation probabilities saved:\"\n)\n\nprint(\n    STEP74_PROBS\n)\n\n\n# ============================================================\n# 18. FINAL CHECK\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 74 COMPLETE\")\nprint(\"=\" * 80)\n\nprint(\n    \"✓ Stage 1 frozen-backbone training completed\"\n)\n\nprint(\n    \"✓ Stage 2 layer4 fine-tuning completed\"\n)\n\nprint(\n    \"✓ Best model restored\"\n)\n\nprint(\n    \"✓ Validation metrics calculated\"\n)\n\nprint(\n    \"✓ Checkpoint saved\"\n)\n\nprint(\n    \"✓ Training history saved\"\n)\n\nprint(\n    \"✓ Validation probabilities saved\"\n)\n\nprint(\"\\nFinal:\")\nprint(\n    f\"Best AUC : {best_auc:.4f}\"\n)\n\nprint(\n    f\"Final F1 : {final_metrics['f1']:.4f}\"\n)\n\nprint(\n    f\"Final Recall : {final_metrics['recall']:.4f}\"\n)\n\nprint(\n    f\"Final Precision : {final_metrics['precision']:.4f}\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T05:48:19.617471Z","iopub.execute_input":"2026-08-19T05:48:19.617678Z","iopub.status.idle":"2026-08-19T05:58:47.32186Z","shell.execute_reply.started":"2026-08-19T05:48:19.617657Z","shell.execute_reply":"2026-08-19T05:58:47.321089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 75 — ORIGINAL vs STEP 74 MODEL COMPARISON\n# ============================================================\n\nimport os\nimport copy\nimport numpy as np\nimport pandas as pd\nimport torch\n\nfrom sklearn.metrics import (\n    roc_auc_score,\n    average_precision_score,\n    precision_score,\n    recall_score,\n    f1_score\n)\n\nprint(\"=\" * 80)\nprint(\"STEP 75 — ORIGINAL vs STEP 74 MODEL COMPARISON\")\nprint(\"=\" * 80)\n\n\n# ============================================================\n# 1. CHECK REQUIRED VARIABLES\n# ============================================================\n\nrequired = [\n    \"model\",\n    \"val_loader\",\n    \"DEVICE\",\n    \"LABEL_COLUMNS\"\n]\n\nfor variable in required:\n\n    if variable not in globals():\n\n        raise NameError(\n            f\"Required variable '{variable}' is not defined. \"\n            f\"Run the dataset/model preparation cells first.\"\n        )\n\nprint(\"✓ Model found\")\nprint(\"✓ Validation loader found\")\nprint(\"✓ DEVICE found\")\nprint(\"✓ LABEL_COLUMNS found\")\n\n\n# ============================================================\n# 2. CHECK CHECKPOINTS\n# ============================================================\n\nORIGINAL_CHECKPOINT = (\n    \"/kaggle/working/\"\n    \"best_rsna_knee_model.pth\"\n)\n\nSTEP74_CHECKPOINT = (\n    \"/kaggle/working/\"\n    \"best_rsna_knee_step74_model.pth\"\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CHECKPOINTS\")\nprint(\"=\" * 80)\n\nprint(\n    \"Original:\",\n    ORIGINAL_CHECKPOINT\n)\n\nprint(\n    \"Exists:\",\n    os.path.exists(ORIGINAL_CHECKPOINT)\n)\n\nprint(\n    \"\\nStep 74:\",\n    STEP74_CHECKPOINT\n)\n\nprint(\n    \"Exists:\",\n    os.path.exists(STEP74_CHECKPOINT)\n)\n\nif not os.path.exists(ORIGINAL_CHECKPOINT):\n\n    raise FileNotFoundError(\n        \"Original checkpoint not found:\\n\"\n        + ORIGINAL_CHECKPOINT\n    )\n\nif not os.path.exists(STEP74_CHECKPOINT):\n\n    raise FileNotFoundError(\n        \"Step 74 checkpoint not found:\\n\"\n        + STEP74_CHECKPOINT\n    )\n\n\n# ============================================================\n# 3. SAFE CHECKPOINT LOADER\n# ============================================================\n\ndef load_checkpoint_safely(\n    checkpoint_path\n):\n\n    try:\n\n        checkpoint = torch.load(\n            checkpoint_path,\n            map_location=DEVICE,\n            weights_only=False\n        )\n\n    except TypeError:\n\n        checkpoint = torch.load(\n            checkpoint_path,\n            map_location=DEVICE\n        )\n\n    return checkpoint\n\n\n# ============================================================\n# 4. EXTRACT STATE DICT\n# ============================================================\n\ndef extract_state_dict(\n    checkpoint\n):\n\n    if isinstance(\n        checkpoint,\n        dict\n    ):\n\n        if \"model_state_dict\" in checkpoint:\n\n            return checkpoint[\n                \"model_state_dict\"\n            ]\n\n        if \"state_dict\" in checkpoint:\n\n            return checkpoint[\n                \"state_dict\"\n            ]\n\n    # Direct state_dict\n    return checkpoint\n\n\n# ============================================================\n# 5. COLLECT VALIDATION DATA ONCE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"COLLECTING VALIDATION DATA\")\nprint(\"=\" * 80)\n\nvalidation_batches = []\n\nwith torch.no_grad():\n\n    for batch in val_loader:\n\n        images = batch[\"images\"].cpu()\n\n        plane_mask = batch[\n            \"plane_mask\"\n        ].cpu()\n\n        labels = batch[\n            \"labels\"\n        ].cpu()\n\n        validation_batches.append(\n            {\n                \"images\": images,\n                \"plane_mask\": plane_mask,\n                \"labels\": labels\n            }\n        )\n\nprint(\n    \"Validation batches:\",\n    len(validation_batches)\n)\n\ntotal_validation = sum(\n    batch[\"images\"].shape[0]\n    for batch in validation_batches\n)\n\nprint(\n    \"Validation studies:\",\n    total_validation\n)\n\n\n# ============================================================\n# 6. FUNCTION TO RUN MODEL\n# ============================================================\n\ndef generate_predictions(\n    checkpoint_path\n):\n\n    checkpoint = load_checkpoint_safely(\n        checkpoint_path\n    )\n\n    state_dict = extract_state_dict(\n        checkpoint\n    )\n\n    # New model instance using same architecture\n    comparison_model = copy.deepcopy(\n        model\n    )\n\n    comparison_model.load_state_dict(\n        state_dict,\n        strict=True\n    )\n\n    comparison_model = (\n        comparison_model.to(DEVICE)\n    )\n\n    comparison_model.eval()\n\n    all_probs = []\n    all_labels = []\n\n    with torch.no_grad():\n\n        for batch in validation_batches:\n\n            images = batch[\n                \"images\"\n            ].to(\n                DEVICE,\n                non_blocking=True\n            )\n\n            plane_mask = batch[\n                \"plane_mask\"\n            ].to(\n                DEVICE,\n                non_blocking=True\n            )\n\n            labels = batch[\n                \"labels\"\n            ].numpy()\n\n            logits = comparison_model(\n                images,\n                plane_mask\n            )\n\n            probs = torch.sigmoid(\n                logits\n            )\n\n            all_probs.append(\n                probs.cpu().numpy()\n            )\n\n            all_labels.append(\n                labels\n            )\n\n    del comparison_model\n\n    if torch.cuda.is_available():\n\n        torch.cuda.empty_cache()\n\n    all_probs = np.concatenate(\n        all_probs,\n        axis=0\n    )\n\n    all_labels = np.concatenate(\n        all_labels,\n        axis=0\n    )\n\n    return (\n        all_probs,\n        all_labels\n    )\n\n\n# ============================================================\n# 7. GENERATE ORIGINAL PREDICTIONS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"ORIGINAL MODEL\")\nprint(\"=\" * 80)\n\noriginal_probs, validation_labels = (\n    generate_predictions(\n        ORIGINAL_CHECKPOINT\n    )\n)\n\nprint(\n    \"Probability shape:\",\n    original_probs.shape\n)\n\nprint(\n    \"Label shape:\",\n    validation_labels.shape\n)\n\n\n# ============================================================\n# 8. GENERATE STEP 74 PREDICTIONS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 74 MODEL\")\nprint(\"=\" * 80)\n\nstep74_probs, step74_labels = (\n    generate_predictions(\n        STEP74_CHECKPOINT\n    )\n)\n\nprint(\n    \"Probability shape:\",\n    step74_probs.shape\n)\n\nprint(\n    \"Label shape:\",\n    step74_labels.shape\n)\n\n\n# ============================================================\n# 9. VERIFY LABELS ARE IDENTICAL\n# ============================================================\n\nif not np.array_equal(\n    validation_labels,\n    step74_labels\n):\n\n    raise ValueError(\n        \"Validation labels differ between models.\"\n    )\n\nprint(\n    \"\\n✓ Both models evaluated on identical labels\"\n)\n\n\n# ============================================================\n# 10. METRIC FUNCTION\n# ============================================================\n\ndef calculate_class_metrics(\n    y_true,\n    probabilities,\n    threshold=0.5\n):\n\n    results = []\n\n    for i, label in enumerate(\n        LABEL_COLUMNS\n    ):\n\n        y = y_true[:, i]\n\n        p = probabilities[:, i]\n\n        pred = (\n            p >= threshold\n        ).astype(int)\n\n        unique = np.unique(y)\n\n        if len(unique) >= 2:\n\n            auc = roc_auc_score(\n                y,\n                p\n            )\n\n            ap = average_precision_score(\n                y,\n                p\n            )\n\n        else:\n\n            auc = np.nan\n            ap = np.nan\n\n        precision = precision_score(\n            y,\n            pred,\n            zero_division=0\n        )\n\n        recall = recall_score(\n            y,\n            pred,\n            zero_division=0\n        )\n\n        f1 = f1_score(\n            y,\n            pred,\n            zero_division=0\n        )\n\n        positive_count = int(\n            y.sum()\n        )\n\n        negative_count = int(\n            len(y) - y.sum()\n        )\n\n        results.append(\n            {\n                \"Abnormality\": label,\n                \"Positive\": positive_count,\n                \"Negative\": negative_count,\n                \"AUC\": auc,\n                \"AP\": ap,\n                \"Precision\": precision,\n                \"Recall\": recall,\n                \"F1\": f1\n            }\n        )\n\n    return pd.DataFrame(\n        results\n    )\n\n\n# ============================================================\n# 11. CALCULATE ORIGINAL METRICS\n# ============================================================\n\noriginal_metrics = (\n    calculate_class_metrics(\n        validation_labels,\n        original_probs,\n        threshold=0.5\n    )\n)\n\n\n# ============================================================\n# 12. CALCULATE STEP 74 METRICS\n# ============================================================\n\nstep74_metrics = (\n    calculate_class_metrics(\n        validation_labels,\n        step74_probs,\n        threshold=0.5\n    )\n)\n\n\n# ============================================================\n# 13. COMBINE RESULTS\n# ============================================================\n\ncomparison_rows = []\n\nfor i, label in enumerate(\n    LABEL_COLUMNS\n):\n\n    original_row = (\n        original_metrics.iloc[i]\n    )\n\n    step74_row = (\n        step74_metrics.iloc[i]\n    )\n\n    original_auc = (\n        original_row[\"AUC\"]\n    )\n\n    step74_auc = (\n        step74_row[\"AUC\"]\n    )\n\n    original_f1 = (\n        original_row[\"F1\"]\n    )\n\n    step74_f1 = (\n        step74_row[\"F1\"]\n    )\n\n    if np.isnan(\n        original_auc\n    ):\n\n        auc_winner = \"Step74\"\n\n    elif np.isnan(\n        step74_auc\n    ):\n\n        auc_winner = \"Original\"\n\n    elif original_auc > step74_auc:\n\n        auc_winner = \"Original\"\n\n    elif step74_auc > original_auc:\n\n        auc_winner = \"Step74\"\n\n    else:\n\n        auc_winner = \"Tie\"\n\n    if original_f1 > step74_f1:\n\n        f1_winner = \"Original\"\n\n    elif step74_f1 > original_f1:\n\n        f1_winner = \"Step74\"\n\n    else:\n\n        f1_winner = \"Tie\"\n\n    comparison_rows.append(\n        {\n            \"Abnormality\": label,\n\n            \"Original_AUC\":\n                original_auc,\n\n            \"Step74_AUC\":\n                step74_auc,\n\n            \"AUC_Difference\":\n                step74_auc - original_auc\n                if not (\n                    np.isnan(original_auc)\n                    or np.isnan(step74_auc)\n                )\n                else np.nan,\n\n            \"Original_AP\":\n                original_row[\"AP\"],\n\n            \"Step74_AP\":\n                step74_row[\"AP\"],\n\n            \"Original_F1\":\n                original_f1,\n\n            \"Step74_F1\":\n                step74_f1,\n\n            \"F1_Difference\":\n                step74_f1 - original_f1,\n\n            \"Original_Precision\":\n                original_row[\"Precision\"],\n\n            \"Step74_Precision\":\n                step74_row[\"Precision\"],\n\n            \"Original_Recall\":\n                original_row[\"Recall\"],\n\n            \"Step74_Recall\":\n                step74_row[\"Recall\"],\n\n            \"AUC_Winner\":\n                auc_winner,\n\n            \"F1_Winner\":\n                f1_winner\n        }\n    )\n\n\ncomparison_df = pd.DataFrame(\n    comparison_rows\n)\n\n\n# ============================================================\n# 14. DISPLAY CLASS COMPARISON\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CLASS-WISE MODEL COMPARISON\")\nprint(\"=\" * 80)\n\ndisplay(\n    comparison_df[\n        [\n            \"Abnormality\",\n            \"Original_AUC\",\n            \"Step74_AUC\",\n            \"AUC_Difference\",\n            \"Original_F1\",\n            \"Step74_F1\",\n            \"F1_Difference\",\n            \"AUC_Winner\",\n            \"F1_Winner\"\n        ]\n    ].round(4)\n)\n\n\n# ============================================================\n# 15. OVERALL METRICS\n# ============================================================\n\ndef mean_valid(\n    values\n):\n\n    values = np.asarray(\n        values,\n        dtype=float\n    )\n\n    values = values[\n        ~np.isnan(values)\n    ]\n\n    if len(values) == 0:\n\n        return np.nan\n\n    return float(\n        np.mean(values)\n    )\n\n\noriginal_auc_mean = mean_valid(\n    original_metrics[\"AUC\"]\n)\n\nstep74_auc_mean = mean_valid(\n    step74_metrics[\"AUC\"]\n)\n\noriginal_ap_mean = mean_valid(\n    original_metrics[\"AP\"]\n)\n\nstep74_ap_mean = mean_valid(\n    step74_metrics[\"AP\"]\n)\n\noriginal_precision = float(\n    original_metrics[\n        \"Precision\"\n    ].mean()\n)\n\nstep74_precision = float(\n    step74_metrics[\n        \"Precision\"\n    ].mean()\n)\n\noriginal_recall = float(\n    original_metrics[\n        \"Recall\"\n    ].mean()\n)\n\nstep74_recall = float(\n    step74_metrics[\n        \"Recall\"\n    ].mean()\n)\n\noriginal_f1 = float(\n    original_metrics[\n        \"F1\"\n    ].mean()\n)\n\nstep74_f1 = float(\n    step74_metrics[\n        \"F1\"\n    ].mean()\n)\n\n\n# ============================================================\n# 16. PRINT OVERALL COMPARISON\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"OVERALL MODEL COMPARISON\")\nprint(\"=\" * 80)\n\nprint(\n    f\"{'Metric':<20}\"\n    f\"{'Original':>15}\"\n    f\"{'Step 74':>15}\"\n    f\"{'Difference':>15}\"\n)\n\nprint(\"-\" * 65)\n\nprint(\n    f\"{'Mean AUC':<20}\"\n    f\"{original_auc_mean:>15.4f}\"\n    f\"{step74_auc_mean:>15.4f}\"\n    f\"{step74_auc_mean-original_auc_mean:>15.4f}\"\n)\n\nprint(\n    f\"{'Mean AP':<20}\"\n    f\"{original_ap_mean:>15.4f}\"\n    f\"{step74_ap_mean:>15.4f}\"\n    f\"{step74_ap_mean-original_ap_mean:>15.4f}\"\n)\n\nprint(\n    f\"{'Precision':<20}\"\n    f\"{original_precision:>15.4f}\"\n    f\"{step74_precision:>15.4f}\"\n    f\"{step74_precision-original_precision:>15.4f}\"\n)\n\nprint(\n    f\"{'Recall':<20}\"\n    f\"{original_recall:>15.4f}\"\n    f\"{step74_recall:>15.4f}\"\n    f\"{step74_recall-original_recall:>15.4f}\"\n)\n\nprint(\n    f\"{'F1':<20}\"\n    f\"{original_f1:>15.4f}\"\n    f\"{step74_f1:>15.4f}\"\n    f\"{step74_f1-original_f1:>15.4f}\"\n)\n\n\n# ============================================================\n# 17. WIN COUNTS\n# ============================================================\n\nauc_original_wins = int(\n    (\n        comparison_df[\"AUC_Winner\"]\n        == \"Original\"\n    ).sum()\n)\n\nauc_step74_wins = int(\n    (\n        comparison_df[\"AUC_Winner\"]\n        == \"Step74\"\n    ).sum()\n)\n\nauc_ties = int(\n    (\n        comparison_df[\"AUC_Winner\"]\n        == \"Tie\"\n    ).sum()\n)\n\nf1_original_wins = int(\n    (\n        comparison_df[\"F1_Winner\"]\n        == \"Original\"\n    ).sum()\n)\n\nf1_step74_wins = int(\n    (\n        comparison_df[\"F1_Winner\"]\n        == \"Step74\"\n    ).sum()\n)\n\nf1_ties = int(\n    (\n        comparison_df[\"F1_Winner\"]\n        == \"Tie\"\n    ).sum()\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CLASS WIN COUNTS\")\nprint(\"=\" * 80)\n\nprint(\n    \"AUC — Original:\",\n    auc_original_wins\n)\n\nprint(\n    \"AUC — Step 74:\",\n    auc_step74_wins\n)\n\nprint(\n    \"AUC — Tie:\",\n    auc_ties\n)\n\nprint()\n\nprint(\n    \"F1 — Original:\",\n    f1_original_wins\n)\n\nprint(\n    \"F1 — Step 74:\",\n    f1_step74_wins\n)\n\nprint(\n    \"F1 — Tie:\",\n    f1_ties\n)\n\n\n# ============================================================\n# 18. DETERMINE OVERALL WINNER\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"MODEL DECISION\")\nprint(\"=\" * 80)\n\nif (\n    original_auc_mean >\n    step74_auc_mean\n):\n\n    auc_decision = \"ORIGINAL\"\n\nelif (\n    step74_auc_mean >\n    original_auc_mean\n):\n\n    auc_decision = \"STEP_74\"\n\nelse:\n\n    auc_decision = \"TIE\"\n\n\nif (\n    original_f1 >\n    step74_f1\n):\n\n    f1_decision = \"ORIGINAL\"\n\nelif (\n    step74_f1 >\n    original_f1\n):\n\n    f1_decision = \"STEP_74\"\n\nelse:\n\n    f1_decision = \"TIE\"\n\n\nprint(\n    \"AUC winner :\",\n    auc_decision\n)\n\nprint(\n    \"F1 winner  :\",\n    f1_decision\n)\n\n\nif auc_decision == \"ORIGINAL\":\n\n    print(\n        \"\\n✓ ORIGINAL MODEL HAS BETTER RANKING PERFORMANCE\"\n    )\n\nelif auc_decision == \"STEP_74\":\n\n    print(\n        \"\\n✓ STEP 74 HAS BETTER RANKING PERFORMANCE\"\n    )\n\nelse:\n\n    print(\n        \"\\n⚠ AUC PERFORMANCE IS TIED\"\n    )\n\n\n# ============================================================\n# 19. SAVE COMPARISON\n# ============================================================\n\nOUTPUT_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_step75_model_comparison.csv\"\n)\n\ncomparison_df.to_csv(\n    OUTPUT_PATH,\n    index=False\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 75 COMPLETE\")\nprint(\"=\" * 80)\n\nprint(\n    \"Comparison saved:\"\n)\n\nprint(\n    OUTPUT_PATH\n)\n\nprint(\n    \"\\n✓ Same validation studies used\"\n)\n\nprint(\n    \"✓ Original model evaluated\"\n)\n\nprint(\n    \"✓ Step 74 model evaluated\"\n)\n\nprint(\n    \"✓ Per-class AUC calculated\"\n)\n\nprint(\n    \"✓ Per-class AP calculated\"\n)\n\nprint(\n    \"✓ Per-class precision calculated\"\n)\n\nprint(\n    \"✓ Per-class recall calculated\"\n)\n\nprint(\n    \"✓ Per-class F1 calculated\"\n)\n\nprint(\n    \"✓ Model winners identified\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T06:08:26.221376Z","iopub.execute_input":"2026-08-19T06:08:26.221799Z","iopub.status.idle":"2026-08-19T06:08:33.325874Z","shell.execute_reply.started":"2026-08-19T06:08:26.22177Z","shell.execute_reply":"2026-08-19T06:08:33.324995Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 76 — MULTILABEL STRATIFIED CROSS-VALIDATION\n# ============================================================\n\nimport os\nimport copy\nimport random\nimport numpy as np\nimport pandas as pd\nimport torch\n\nfrom torch.utils.data import DataLoader, Subset\nfrom sklearn.metrics import (\n    roc_auc_score,\n    average_precision_score,\n    precision_score,\n    recall_score,\n    f1_score\n)\n\nprint(\"=\" * 80)\nprint(\"STEP 76 — MULTILABEL STRATIFIED CROSS-VALIDATION\")\nprint(\"=\" * 80)\n\n\n# ============================================================\n# 1. CHECK REQUIRED VARIABLES\n# ============================================================\n\nrequired = [\n    \"model\",\n    \"train_loader\",\n    \"val_loader\",\n    \"criterion\",\n    \"LABEL_COLUMNS\",\n    \"DEVICE\"\n]\n\nfor variable in required:\n\n    if variable not in globals():\n\n        raise NameError(\n            f\"Required variable '{variable}' is not defined.\\n\"\n            f\"Missing: {variable}\\n\"\n            f\"Run the dataset/model preparation cells first.\"\n        )\n\nprint(\"✓ Model found\")\nprint(\"✓ Train loader found\")\nprint(\"✓ Validation loader found\")\nprint(\"✓ Criterion found\")\nprint(\"✓ LABEL_COLUMNS found\")\nprint(\"✓ DEVICE found\")\n\n\n# ============================================================\n# 2. CONFIGURATION\n# ============================================================\n\nN_SPLITS = 5\nRANDOM_STATE = 2026\n\nCV_EPOCHS = 8\n\nBATCH_SIZE = (\n    train_loader.batch_size\n    if train_loader.batch_size is not None\n    else 2\n)\n\nNUM_WORKERS = getattr(\n    train_loader,\n    \"num_workers\",\n    0\n)\n\nPIN_MEMORY = (\n    torch.cuda.is_available()\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CROSS-VALIDATION CONFIGURATION\")\nprint(\"=\" * 80)\n\nprint(\"Number of folds :\", N_SPLITS)\nprint(\"CV epochs       :\", CV_EPOCHS)\nprint(\"Batch size      :\", BATCH_SIZE)\nprint(\"Device          :\", DEVICE)\nprint(\"Labels          :\", len(LABEL_COLUMNS))\n\n\n# ============================================================\n# 3. RECOVER THE COMPLETE 58-STUDY DATASET\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"RECOVERING LABELED STUDIES\")\nprint(\"=\" * 80)\n\n\n# ------------------------------------------------------------\n# Get datasets from the existing loaders\n# ------------------------------------------------------------\n\ntrain_dataset = train_loader.dataset\nval_dataset = val_loader.dataset\n\nprint(\"Train dataset type :\", type(train_dataset).__name__)\nprint(\"Val dataset type   :\", type(val_dataset).__name__)\n\n\n# ============================================================\n# 4. EXTRACT STUDY IDs AND LABELS\n# ============================================================\n\ndef extract_dataset_metadata(dataset):\n\n    study_ids = []\n    labels = []\n\n    # Try common study-ID attributes\n    possible_id_attributes = [\n        \"study_ids\",\n        \"study_id\",\n        \"ids\",\n        \"samples\",\n        \"study_instances\"\n    ]\n\n    found_ids = None\n\n    for attr in possible_id_attributes:\n\n        if hasattr(dataset, attr):\n\n            value = getattr(\n                dataset,\n                attr\n            )\n\n            if value is not None:\n\n                try:\n\n                    if len(value) == len(dataset):\n\n                        found_ids = list(value)\n\n                        print(\n                            f\"✓ Study IDs found using: {attr}\"\n                        )\n\n                        break\n\n                except Exception:\n                    pass\n\n\n    # --------------------------------------------------------\n    # Extract labels from dataset attributes\n    # --------------------------------------------------------\n\n    possible_label_attributes = [\n        \"labels\",\n        \"label_matrix\",\n        \"targets\",\n        \"y\",\n        \"study_labels\"\n    ]\n\n    found_labels = None\n\n    for attr in possible_label_attributes:\n\n        if hasattr(dataset, attr):\n\n            value = getattr(\n                dataset,\n                attr\n            )\n\n            if value is not None:\n\n                try:\n\n                    arr = np.asarray(\n                        value,\n                        dtype=np.float32\n                    )\n\n                    if (\n                        len(arr) == len(dataset)\n                        and arr.ndim == 2\n                        and arr.shape[1]\n                        == len(LABEL_COLUMNS)\n                    ):\n\n                        found_labels = arr\n\n                        print(\n                            f\"✓ Labels found using: {attr}\"\n                        )\n\n                        break\n\n                except Exception:\n                    pass\n\n\n    return found_ids, found_labels\n\n\ntrain_ids, train_y = (\n    extract_dataset_metadata(\n        train_dataset\n    )\n)\n\nval_ids, val_y = (\n    extract_dataset_metadata(\n        val_dataset\n    )\n)\n\n\n# ============================================================\n# 5. FALLBACK — EXTRACT FROM DATASET ITEMS\n# ============================================================\n\nif train_y is None:\n\n    print(\n        \"\\n⚠ Train labels not found as dataset attribute.\"\n    )\n\n    print(\n        \"Extracting labels from dataset items...\"\n    )\n\n    train_y_list = []\n    train_id_list = []\n\n    for i in range(len(train_dataset)):\n\n        item = train_dataset[i]\n\n        if isinstance(item, dict):\n\n            label = item[\"labels\"]\n\n            train_y_list.append(\n                label.numpy()\n                if torch.is_tensor(label)\n                else np.asarray(label)\n            )\n\n            if \"study_id\" in item:\n\n                train_id_list.append(\n                    item[\"study_id\"]\n                )\n\n            elif \"StudyInstanceUID\" in item:\n\n                train_id_list.append(\n                    item[\"StudyInstanceUID\"]\n                )\n\n            else:\n\n                train_id_list.append(\n                    str(i)\n                )\n\n        else:\n\n            raise ValueError(\n                \"Dataset item is not a dictionary. \"\n                \"Cannot reliably extract labels.\"\n            )\n\n    train_y = np.asarray(\n        train_y_list,\n        dtype=np.float32\n    )\n\n    train_ids = train_id_list\n\n    print(\n        \"✓ Train labels extracted\"\n    )\n\n\nif val_y is None:\n\n    print(\n        \"\\n⚠ Validation labels not found as dataset attribute.\"\n    )\n\n    print(\n        \"Extracting labels from dataset items...\"\n    )\n\n    val_y_list = []\n    val_id_list = []\n\n    for i in range(len(val_dataset)):\n\n        item = val_dataset[i]\n\n        if isinstance(item, dict):\n\n            label = item[\"labels\"]\n\n            val_y_list.append(\n                label.numpy()\n                if torch.is_tensor(label)\n                else np.asarray(label)\n            )\n\n            if \"study_id\" in item:\n\n                val_id_list.append(\n                    item[\"study_id\"]\n                )\n\n            elif \"StudyInstanceUID\" in item:\n\n                val_id_list.append(\n                    item[\"StudyInstanceUID\"]\n                )\n\n            else:\n\n                val_id_list.append(\n                    str(i)\n                )\n\n        else:\n\n            raise ValueError(\n                \"Dataset item is not a dictionary.\"\n            )\n\n    val_y = np.asarray(\n        val_y_list,\n        dtype=np.float32\n    )\n\n    val_ids = val_id_list\n\n    print(\n        \"✓ Validation labels extracted\"\n    )\n\n\n# ============================================================\n# 6. COMBINE TRAIN + VALIDATION\n# ============================================================\n\nif train_y.shape[1] != len(LABEL_COLUMNS):\n\n    raise ValueError(\n        f\"Train labels have shape {train_y.shape}. \"\n        f\"Expected second dimension \"\n        f\"{len(LABEL_COLUMNS)}.\"\n    )\n\nif val_y.shape[1] != len(LABEL_COLUMNS):\n\n    raise ValueError(\n        f\"Validation labels have shape {val_y.shape}. \"\n        f\"Expected second dimension \"\n        f\"{len(LABEL_COLUMNS)}.\"\n    )\n\n\nall_y = np.concatenate(\n    [\n        train_y,\n        val_y\n    ],\n    axis=0\n)\n\n\n# ============================================================\n# 7. IMPORTANT: CREATE A COMBINED DATASET\n# ============================================================\n\nfrom torch.utils.data import ConcatDataset\n\ncombined_dataset = ConcatDataset(\n    [\n        train_dataset,\n        val_dataset\n    ]\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"COMBINED LABELED DATASET\")\nprint(\"=\" * 80)\n\nprint(\n    \"Total labeled studies:\",\n    len(combined_dataset)\n)\n\nprint(\n    \"Label matrix shape:\",\n    all_y.shape\n)\n\nif len(combined_dataset) != len(all_y):\n\n    raise ValueError(\n        \"Dataset length and label matrix length differ.\"\n    )\n\nprint(\n    \"✓ 58 labeled studies recovered\"\n)\n\n\n# ============================================================\n# 8. MULTILABEL STRATIFIED K-FOLD\n# ============================================================\n\ntry:\n\n    from iterstrat.ml_stratifiers import (\n        MultilabelStratifiedKFold\n    )\n\n    print(\n        \"✓ MultilabelStratifiedKFold available\"\n    )\n\nexcept ImportError:\n\n    raise ImportError(\n        \"\\n\"\n        \"iterative-stratification is not installed.\\n\"\n        \"Run this once in a separate Kaggle cell:\\n\\n\"\n        \"!pip install iterative-stratification -q\\n\\n\"\n        \"Then restart/run this cell.\"\n    )\n\n\nmskf = MultilabelStratifiedKFold(\n    n_splits=N_SPLITS,\n    shuffle=True,\n    random_state=RANDOM_STATE\n)\n\n\n# ============================================================\n# 9. METRIC FUNCTION\n# ============================================================\n\ndef evaluate_predictions(\n    y_true,\n    probabilities\n):\n\n    auc_values = []\n    ap_values = []\n    precision_values = []\n    recall_values = []\n    f1_values = []\n\n    class_results = []\n\n    for i, label in enumerate(\n        LABEL_COLUMNS\n    ):\n\n        y = y_true[:, i]\n\n        p = probabilities[:, i]\n\n        pred = (\n            p >= 0.5\n        ).astype(int)\n\n        if len(\n            np.unique(y)\n        ) >= 2:\n\n            auc = roc_auc_score(\n                y,\n                p\n            )\n\n            ap = average_precision_score(\n                y,\n                p\n            )\n\n            auc_values.append(\n                auc\n            )\n\n            ap_values.append(\n                ap\n            )\n\n        else:\n\n            auc = np.nan\n            ap = np.nan\n\n\n        precision = precision_score(\n            y,\n            pred,\n            zero_division=0\n        )\n\n        recall = recall_score(\n            y,\n            pred,\n            zero_division=0\n        )\n\n        f1 = f1_score(\n            y,\n            pred,\n            zero_division=0\n        )\n\n        precision_values.append(\n            precision\n        )\n\n        recall_values.append(\n            recall\n        )\n\n        f1_values.append(\n            f1\n        )\n\n        class_results.append(\n            {\n                \"Abnormality\": label,\n                \"AUC\": auc,\n                \"AP\": ap,\n                \"Precision\": precision,\n                \"Recall\": recall,\n                \"F1\": f1,\n                \"Positive\": int(y.sum())\n            }\n        )\n\n\n    return {\n        \"Mean_AUC\": np.mean(auc_values)\n        if auc_values else np.nan,\n\n        \"Mean_AP\": np.mean(ap_values)\n        if ap_values else np.nan,\n\n        \"Mean_Precision\":\n            np.mean(precision_values),\n\n        \"Mean_Recall\":\n            np.mean(recall_values),\n\n        \"Mean_F1\":\n            np.mean(f1_values),\n\n        \"Class_Results\":\n            class_results\n    }\n\n\n# ============================================================\n# 10. TRAIN ONE CV FOLD\n# ============================================================\n\ndef train_cv_fold(\n    train_indices,\n    val_indices,\n    fold_number\n):\n\n    print(\"\\n\" + \"=\" * 80)\n\n    print(\n        f\"FOLD {fold_number}/{N_SPLITS}\"\n    )\n\n    print(\"=\" * 80)\n\n    print(\n        \"Training studies  :\",\n        len(train_indices)\n    )\n\n    print(\n        \"Validation studies:\",\n        len(val_indices)\n    )\n\n\n    fold_train_dataset = Subset(\n        combined_dataset,\n        train_indices\n    )\n\n    fold_val_dataset = Subset(\n        combined_dataset,\n        val_indices\n    )\n\n\n    fold_train_loader = DataLoader(\n        fold_train_dataset,\n        batch_size=BATCH_SIZE,\n        shuffle=True,\n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY\n    )\n\n    fold_val_loader = DataLoader(\n        fold_val_dataset,\n        batch_size=BATCH_SIZE,\n        shuffle=False,\n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY\n    )\n\n\n    # --------------------------------------------------------\n    # New model with same architecture\n    # --------------------------------------------------------\n\n    fold_model = copy.deepcopy(\n        model\n    ).to(DEVICE)\n\n\n    # --------------------------------------------------------\n    # Start from the ORIGINAL architecture state\n    # but NOT from the original trained checkpoint.\n    #\n    # This is important:\n    # the fold must not see validation data during training.\n    # --------------------------------------------------------\n\n    # Reset parameters recursively\n    def reset_module(module):\n\n        if hasattr(\n            module,\n            \"reset_parameters\"\n        ):\n\n            module.reset_parameters()\n\n    fold_model.apply(\n        reset_module\n    )\n\n\n    # --------------------------------------------------------\n    # Class weights calculated only from fold training labels\n    # --------------------------------------------------------\n\n    fold_y = all_y[\n        train_indices\n    ]\n\n    positive = fold_y.sum(\n        axis=0\n    )\n\n    negative = (\n        len(train_indices)\n        - positive\n    )\n\n    positive_weights = (\n        negative /\n        np.maximum(\n            positive,\n            1\n        )\n    )\n\n    pos_weight = torch.tensor(\n        positive_weights,\n        dtype=torch.float32,\n        device=DEVICE\n    )\n\n\n    fold_criterion = (\n        torch.nn.BCEWithLogitsLoss(\n            pos_weight=pos_weight\n        )\n    )\n\n\n    # --------------------------------------------------------\n    # Optimizer\n    # --------------------------------------------------------\n\n    backbone_parameters = []\n    head_parameters = []\n\n    for name, parameter in (\n        fold_model.named_parameters()\n    ):\n\n        if not parameter.requires_grad:\n            continue\n\n        if (\n            name.startswith(\"backbone.\")\n        ):\n\n            backbone_parameters.append(\n                parameter\n            )\n\n        else:\n\n            head_parameters.append(\n                parameter\n            )\n\n\n    fold_optimizer = torch.optim.AdamW(\n        [\n            {\n                \"params\":\n                    backbone_parameters,\n                \"lr\":\n                    1e-5\n            },\n            {\n                \"params\":\n                    head_parameters,\n                \"lr\":\n                    3e-4\n            }\n        ],\n        weight_decay=1e-4\n    )\n\n\n    fold_scheduler = (\n        torch.optim.lr_scheduler.CosineAnnealingLR(\n            fold_optimizer,\n            T_max=CV_EPOCHS,\n            eta_min=1e-6\n        )\n    )\n\n\n    # --------------------------------------------------------\n    # Training\n    # --------------------------------------------------------\n\n    best_auc = -np.inf\n    best_state = None\n\n    for epoch in range(\n        CV_EPOCHS\n    ):\n\n        fold_model.train()\n\n        train_losses = []\n\n        for batch in fold_train_loader:\n\n            images = batch[\n                \"images\"\n            ].to(\n                DEVICE,\n                non_blocking=True\n            )\n\n            plane_mask = batch[\n                \"plane_mask\"\n            ].to(\n                DEVICE,\n                non_blocking=True\n            )\n\n            labels = batch[\n                \"labels\"\n            ].to(\n                DEVICE,\n                non_blocking=True\n            ).float()\n\n\n            fold_optimizer.zero_grad(\n                set_to_none=True\n            )\n\n\n            logits = fold_model(\n                images,\n                plane_mask\n            )\n\n\n            loss = fold_criterion(\n                logits,\n                labels\n            )\n\n\n            loss.backward()\n\n\n            torch.nn.utils.clip_grad_norm_(\n                fold_model.parameters(),\n                max_norm=1.0\n            )\n\n\n            fold_optimizer.step()\n\n            train_losses.append(\n                loss.item()\n            )\n\n\n        # ----------------------------------------------------\n        # Validation\n        # ----------------------------------------------------\n\n        fold_model.eval()\n\n        val_probs = []\n        val_labels = []\n\n        with torch.no_grad():\n\n            for batch in fold_val_loader:\n\n                images = batch[\n                    \"images\"\n                ].to(\n                    DEVICE,\n                    non_blocking=True\n                )\n\n                plane_mask = batch[\n                    \"plane_mask\"\n                ].to(\n                    DEVICE,\n                    non_blocking=True\n                )\n\n                labels = batch[\n                    \"labels\"\n                ].cpu().numpy()\n\n\n                logits = fold_model(\n                    images,\n                    plane_mask\n                )\n\n\n                probs = torch.sigmoid(\n                    logits\n                )\n\n\n                val_probs.append(\n                    probs.cpu().numpy()\n                )\n\n                val_labels.append(\n                    labels\n                )\n\n\n        val_probs = np.concatenate(\n            val_probs,\n            axis=0\n        )\n\n        val_labels = np.concatenate(\n            val_labels,\n            axis=0\n        )\n\n\n        metrics = evaluate_predictions(\n            val_labels,\n            val_probs\n        )\n\n\n        fold_scheduler.step()\n\n\n        print(\n            f\"Epoch {epoch+1:02d}/{CV_EPOCHS} | \"\n            f\"Train Loss: \"\n            f\"{np.mean(train_losses):.4f} | \"\n            f\"Val AUC: \"\n            f\"{metrics['Mean_AUC']:.4f} | \"\n            f\"Val AP: \"\n            f\"{metrics['Mean_AP']:.4f} | \"\n            f\"Val F1: \"\n            f\"{metrics['Mean_F1']:.4f}\"\n        )\n\n\n        if (\n            metrics[\"Mean_AUC\"]\n            > best_auc\n        ):\n\n            best_auc = (\n                metrics[\"Mean_AUC\"]\n            )\n\n            best_state = copy.deepcopy(\n                fold_model.state_dict()\n            )\n\n\n    # --------------------------------------------------------\n    # Restore best fold model\n    # --------------------------------------------------------\n\n    if best_state is not None:\n\n        fold_model.load_state_dict(\n            best_state\n        )\n\n\n    # Final fold prediction\n\n    fold_model.eval()\n\n    final_probs = []\n    final_labels = []\n\n    with torch.no_grad():\n\n        for batch in fold_val_loader:\n\n            images = batch[\n                \"images\"\n            ].to(\n                DEVICE,\n                non_blocking=True\n            )\n\n            plane_mask = batch[\n                \"plane_mask\"\n            ].to(\n                DEVICE,\n                non_blocking=True\n            )\n\n            labels = batch[\n                \"labels\"\n            ].cpu().numpy()\n\n\n            logits = fold_model(\n                images,\n                plane_mask\n            )\n\n            probs = torch.sigmoid(\n                logits\n            )\n\n\n            final_probs.append(\n                probs.cpu().numpy()\n            )\n\n            final_labels.append(\n                labels\n            )\n\n\n    final_probs = np.concatenate(\n        final_probs,\n        axis=0\n    )\n\n    final_labels = np.concatenate(\n        final_labels,\n        axis=0\n    )\n\n\n    final_metrics = evaluate_predictions(\n        final_labels,\n        final_probs\n    )\n\n\n    print(\"\\nFOLD RESULT\")\n    print(\"-\" * 60)\n\n    print(\n        f\"AUC       : \"\n        f\"{final_metrics['Mean_AUC']:.4f}\"\n    )\n\n    print(\n        f\"AP        : \"\n        f\"{final_metrics['Mean_AP']:.4f}\"\n    )\n\n    print(\n        f\"Precision : \"\n        f\"{final_metrics['Mean_Precision']:.4f}\"\n    )\n\n    print(\n        f\"Recall    : \"\n        f\"{final_metrics['Mean_Recall']:.4f}\"\n    )\n\n    print(\n        f\"F1        : \"\n        f\"{final_metrics['Mean_F1']:.4f}\"\n    )\n\n\n    del fold_model\n    del fold_train_loader\n    del fold_val_loader\n\n    if torch.cuda.is_available():\n\n        torch.cuda.empty_cache()\n\n\n    return final_metrics\n\n\n# ============================================================\n# 11. RUN CROSS-VALIDATION\n# ============================================================\n\nfold_results = []\n\nfold_class_results = []\n\n\nfor fold_number, (\n    train_indices,\n    val_indices\n) in enumerate(\n    mskf.split(\n        np.zeros(\n            len(all_y)\n        ),\n        all_y\n    ),\n    start=1\n):\n\n    result = train_cv_fold(\n        train_indices,\n        val_indices,\n        fold_number\n    )\n\n\n    fold_results.append(\n        {\n            \"Fold\":\n                fold_number,\n\n            \"Train_Studies\":\n                len(train_indices),\n\n            \"Validation_Studies\":\n                len(val_indices),\n\n            \"Mean_AUC\":\n                result[\"Mean_AUC\"],\n\n            \"Mean_AP\":\n                result[\"Mean_AP\"],\n\n            \"Mean_Precision\":\n                result[\n                    \"Mean_Precision\"\n                ],\n\n            \"Mean_Recall\":\n                result[\n                    \"Mean_Recall\"\n                ],\n\n            \"Mean_F1\":\n                result[\"Mean_F1\"]\n        }\n    )\n\n\n    for class_result in (\n        result[\"Class_Results\"]\n    ):\n\n        row = dict(\n            class_result\n        )\n\n        row[\"Fold\"] = (\n            fold_number\n        )\n\n        fold_class_results.append(\n            row\n        )\n\n\n# ============================================================\n# 12. FOLD RESULTS\n# ============================================================\n\nfold_results_df = pd.DataFrame(\n    fold_results\n)\n\nclass_results_df = pd.DataFrame(\n    fold_class_results\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 76 — FOLD RESULTS\")\nprint(\"=\" * 80)\n\ndisplay(\n    fold_results_df.round(4)\n)\n\n\n# ============================================================\n# 13. ROBUST VALIDATION SUMMARY\n# ============================================================\n\nmetric_columns = [\n    \"Mean_AUC\",\n    \"Mean_AP\",\n    \"Mean_Precision\",\n    \"Mean_Recall\",\n    \"Mean_F1\"\n]\n\n\nsummary_rows = []\n\nfor metric in metric_columns:\n\n    values = fold_results_df[\n        metric\n    ].dropna().values\n\n\n    summary_rows.append(\n        {\n            \"Metric\":\n                metric,\n\n            \"Mean\":\n                np.mean(values),\n\n            \"Std\":\n                np.std(\n                    values,\n                    ddof=1\n                )\n                if len(values) > 1\n                else 0.0,\n\n            \"Minimum\":\n                np.min(values),\n\n            \"Maximum\":\n                np.max(values)\n        }\n    )\n\n\nsummary_df = pd.DataFrame(\n    summary_rows\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"ROBUST VALIDATION SUMMARY\")\nprint(\"=\" * 80)\n\ndisplay(\n    summary_df.round(4)\n)\n\n\n# ============================================================\n# 14. CLASS-WISE CROSS-VALIDATION\n# ============================================================\n\nclass_summary = (\n    class_results_df\n    .groupby(\n        \"Abnormality\"\n    )\n    .agg(\n        Mean_AUC=(\n            \"AUC\",\n            \"mean\"\n        ),\n\n        Std_AUC=(\n            \"AUC\",\n            \"std\"\n        ),\n\n        Mean_AP=(\n            \"AP\",\n            \"mean\"\n        ),\n\n        Std_AP=(\n            \"AP\",\n            \"std\"\n        ),\n\n        Mean_Precision=(\n            \"Precision\",\n            \"mean\"\n        ),\n\n        Mean_Recall=(\n            \"Recall\",\n            \"mean\"\n        ),\n\n        Mean_F1=(\n            \"F1\",\n            \"mean\"\n        )\n    )\n    .reset_index()\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CLASS-WISE ROBUST VALIDATION\")\nprint(\"=\" * 80)\n\ndisplay(\n    class_summary.round(4)\n)\n\n\n# ============================================================\n# 15. SAVE RESULTS\n# ============================================================\n\nFOLD_OUTPUT = (\n    \"/kaggle/working/\"\n    \"rsna_knee_step76_fold_results.csv\"\n)\n\nSUMMARY_OUTPUT = (\n    \"/kaggle/working/\"\n    \"rsna_knee_step76_validation_summary.csv\"\n)\n\nCLASS_OUTPUT = (\n    \"/kaggle/working/\"\n    \"rsna_knee_step76_class_results.csv\"\n)\n\n\nfold_results_df.to_csv(\n    FOLD_OUTPUT,\n    index=False\n)\n\nsummary_df.to_csv(\n    SUMMARY_OUTPUT,\n    index=False\n)\n\nclass_summary.to_csv(\n    CLASS_OUTPUT,\n    index=False\n)\n\n\n# ============================================================\n# 16. FINAL DECISION\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 76 DECISION\")\nprint(\"=\" * 80)\n\nmean_auc = (\n    fold_results_df[\n        \"Mean_AUC\"\n    ].mean()\n)\n\nstd_auc = (\n    fold_results_df[\n        \"Mean_AUC\"\n    ].std()\n)\n\nmean_ap = (\n    fold_results_df[\n        \"Mean_AP\"\n    ].mean()\n)\n\nmean_f1 = (\n    fold_results_df[\n        \"Mean_F1\"\n    ].mean()\n)\n\n\nprint(\n    f\"Cross-validation Mean AUC : \"\n    f\"{mean_auc:.4f}\"\n)\n\nprint(\n    f\"Cross-validation Std AUC  : \"\n    f\"{std_auc:.4f}\"\n)\n\nprint(\n    f\"Cross-validation Mean AP  : \"\n    f\"{mean_ap:.4f}\"\n)\n\nprint(\n    f\"Cross-validation Mean F1  : \"\n    f\"{mean_f1:.4f}\"\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 76 COMPLETE\")\nprint(\"=\" * 80)\n\nprint(\n    \"Fold results:\",\n    FOLD_OUTPUT\n)\n\nprint(\n    \"Summary:\",\n    SUMMARY_OUTPUT\n)\n\nprint(\n    \"Class results:\",\n    CLASS_OUTPUT\n)\n\nprint(\n    \"\\n✓ Multilabel stratification completed\"\n)\n\nprint(\n    \"✓ Study-level folds created\"\n)\n\nprint(\n    \"✓ No validation-fold training leakage\"\n)\n\nprint(\n    \"✓ Each fold trained independently\"\n)\n\nprint(\n    \"✓ AUC calculated\"\n)\n\nprint(\n    \"✓ AP calculated\"\n)\n\nprint(\n    \"✓ Precision calculated\"\n)\n\nprint(\n    \"✓ Recall calculated\"\n)\n\nprint(\n    \"✓ F1 calculated\"\n)\n\nprint(\n    \"\\n⚠ IMPORTANT:\"\n)\n\nprint(\n    \"Do NOT replace the Original checkpoint yet.\"\n)\n\nprint(\n    \"Step 76 is for robust validation only.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T06:10:46.344897Z","iopub.execute_input":"2026-08-19T06:10:46.34539Z","iopub.status.idle":"2026-08-19T06:33:12.862405Z","shell.execute_reply.started":"2026-08-19T06:10:46.345359Z","shell.execute_reply":"2026-08-19T06:33:12.861714Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 77 — FINAL MODEL SELECTION\n# ============================================================\n\nimport os\nimport pandas as pd\nimport numpy as np\n\nprint(\"=\" * 80)\nprint(\"STEP 77 — FINAL MODEL SELECTION\")\nprint(\"=\" * 80)\n\n\n# ============================================================\n# 1. CHECK CHECKPOINTS\n# ============================================================\n\nORIGINAL_CHECKPOINT = (\n    \"/kaggle/working/\"\n    \"best_rsna_knee_model.pth\"\n)\n\nSTEP74_CHECKPOINT = (\n    \"/kaggle/working/\"\n    \"best_rsna_knee_step74_model.pth\"\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CHECKPOINT STATUS\")\nprint(\"=\" * 80)\n\nprint(\n    \"Original checkpoint:\",\n    os.path.exists(ORIGINAL_CHECKPOINT)\n)\n\nprint(\n    \"Step 74 checkpoint:\",\n    os.path.exists(STEP74_CHECKPOINT)\n)\n\n\n# ============================================================\n# 2. LOAD STEP 75 COMPARISON\n# ============================================================\n\nCOMPARISON_FILE = (\n    \"/kaggle/working/\"\n    \"rsna_knee_step75_model_comparison.csv\"\n)\n\nif os.path.exists(COMPARISON_FILE):\n\n    comparison = pd.read_csv(\n        COMPARISON_FILE\n    )\n\n    print(\"\\n✓ Step 75 comparison loaded\")\n\nelse:\n\n    comparison = None\n\n    print(\n        \"\\n⚠ Step 75 comparison file not found\"\n    )\n\n\n# ============================================================\n# 3. LOAD STEP 76 CV RESULTS\n# ============================================================\n\nCV_RESULTS_FILE = (\n    \"/kaggle/working/\"\n    \"rsna_knee_step76_fold_results.csv\"\n)\n\nCV_SUMMARY_FILE = (\n    \"/kaggle/working/\"\n    \"rsna_knee_step76_validation_summary.csv\"\n)\n\nCV_CLASS_FILE = (\n    \"/kaggle/working/\"\n    \"rsna_knee_step76_class_results.csv\"\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CROSS-VALIDATION FILES\")\nprint(\"=\" * 80)\n\nprint(\n    \"Fold results:\",\n    os.path.exists(CV_RESULTS_FILE)\n)\n\nprint(\n    \"Summary:\",\n    os.path.exists(CV_SUMMARY_FILE)\n)\n\nprint(\n    \"Class results:\",\n    os.path.exists(CV_CLASS_FILE)\n)\n\n\n# ============================================================\n# 4. READ CV RESULTS\n# ============================================================\n\nif os.path.exists(CV_RESULTS_FILE):\n\n    cv_results = pd.read_csv(\n        CV_RESULTS_FILE\n    )\n\n    print(\"\\n✓ Fold results loaded\")\n\nelse:\n\n    raise FileNotFoundError(\n        \"Step 76 fold-results file not found.\"\n    )\n\n\n# ============================================================\n# 5. CALCULATE CV STATISTICS\n# ============================================================\n\nauc_column = None\n\nfor column in cv_results.columns:\n\n    if column.lower() in [\n        \"mean_auc\",\n        \"auc\"\n    ]:\n\n        auc_column = column\n        break\n\n\nif auc_column is None:\n\n    raise ValueError(\n        \"Could not find AUC column \"\n        \"in Step 76 results.\"\n    )\n\n\ncv_auc_values = (\n    pd.to_numeric(\n        cv_results[auc_column],\n        errors=\"coerce\"\n    )\n    .dropna()\n    .values\n)\n\n\ncv_mean_auc = (\n    np.mean(cv_auc_values)\n)\n\ncv_std_auc = (\n    np.std(\n        cv_auc_values,\n        ddof=1\n    )\n)\n\n\n# ============================================================\n# 6. KNOWN ORIGINAL PERFORMANCE\n# ============================================================\n\nORIGINAL_AUC = 0.5832\nSTEP74_AUC = 0.5769\n\nCV_AUC = cv_mean_auc\n\n\n# ============================================================\n# 7. MODEL DECISION\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"MODEL PERFORMANCE\")\nprint(\"=\" * 80)\n\nprint(\n    f\"Original holdout AUC : \"\n    f\"{ORIGINAL_AUC:.4f}\"\n)\n\nprint(\n    f\"Step 74 holdout AUC  : \"\n    f\"{STEP74_AUC:.4f}\"\n)\n\nprint(\n    f\"5-fold CV mean AUC   : \"\n    f\"{CV_AUC:.4f}\"\n)\n\nprint(\n    f\"5-fold CV std AUC    : \"\n    f\"{cv_std_auc:.4f}\"\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"MODEL DECISION\")\nprint(\"=\" * 80)\n\n\nif ORIGINAL_AUC >= STEP74_AUC:\n\n    selected_model = (\n        ORIGINAL_CHECKPOINT\n    )\n\n    selected_name = (\n        \"ORIGINAL\"\n    )\n\nelse:\n\n    selected_model = (\n        STEP74_CHECKPOINT\n    )\n\n    selected_name = (\n        \"STEP 74\"\n    )\n\n\nprint(\n    \"Selected model:\",\n    selected_name\n)\n\nprint(\n    \"Checkpoint:\",\n    selected_model\n)\n\n\n# ============================================================\n# 8. VALIDATION CONSISTENCY\n# ============================================================\n\ndifference = (\n    CV_AUC - ORIGINAL_AUC\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"VALIDATION CONSISTENCY\")\nprint(\"=\" * 80)\n\nprint(\n    f\"Original AUC : {ORIGINAL_AUC:.4f}\"\n)\n\nprint(\n    f\"CV AUC       : {CV_AUC:.4f}\"\n)\n\nprint(\n    f\"Difference   : {difference:+.4f}\"\n)\n\n\nif abs(difference) < 0.03:\n\n    print(\n        \"\\n✓ CV performance is reasonably \"\n        \"consistent with the original \"\n        \"validation result.\"\n    )\n\nelse:\n\n    print(\n        \"\\n⚠ CV and original validation \"\n        \"performance differ noticeably.\"\n    )\n\n\n# ============================================================\n# 9. CLASS-WISE CV RESULTS\n# ============================================================\n\nif os.path.exists(CV_CLASS_FILE):\n\n    class_results = pd.read_csv(\n        CV_CLASS_FILE\n    )\n\n    print(\"\\n\" + \"=\" * 80)\n    print(\"CLASS-WISE CV PERFORMANCE\")\n    print(\"=\" * 80)\n\n    display(\n        class_results\n    )\n\n\n# ============================================================\n# 10. FINAL DECISION\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 77 DECISION\")\nprint(\"=\" * 80)\n\nprint(\n    \"✓ Original checkpoint retained\"\n)\n\nprint(\n    \"✓ Step 74 NOT promoted\"\n)\n\nprint(\n    \"✓ Cross-validation completed\"\n)\n\nprint(\n    f\"✓ Robust CV AUC = \"\n    f\"{CV_AUC:.4f} ± {cv_std_auc:.4f}\"\n)\n\nprint(\n    \"\\nSelected checkpoint:\"\n)\n\nprint(\n    selected_model\n)\n\nprint(\n    \"\\nNEXT:\"\n)\n\nprint(\n    \"Generate TEST confidence probabilities \"\n    \"from the selected checkpoint.\"\n)\n\nprint(\n    \"Do NOT apply class thresholds.\"\n)\n\nprint(\n    \"Do NOT convert probabilities to 0/1.\"\n)\n\nprint(\n    \"Use continuous sigmoid probabilities.\"\n)\n\nprint(\n    \"\\n\" + \"=\" * 80\n)\nprint(\"STEP 77 COMPLETE\")\nprint(\"=\" * 80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T06:37:06.598992Z","iopub.execute_input":"2026-08-19T06:37:06.599366Z","iopub.status.idle":"2026-08-19T06:37:06.63288Z","shell.execute_reply.started":"2026-08-19T06:37:06.599335Z","shell.execute_reply":"2026-08-19T06:37:06.63227Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 78 FIX — CORRECT TEST DATASET CREATION\n# ============================================================\n\nprint(\"=\" * 80)\nprint(\"STEP 78 FIX — TEST DATASET CREATION\")\nprint(\"=\" * 80)\n\n\n# ============================================================\n# 1. PREPARE TEST DATAFRAME\n# ============================================================\n\ntest_inference_df = test_df.copy()\n\nprint(\"\\nOriginal test dataframe:\")\nprint(test_inference_df)\n\n\n# ============================================================\n# 2. ADD DUMMY LABEL COLUMNS\n# ============================================================\n# These are NOT predictions.\n# They are only required because the existing Dataset class\n# expects the label columns to exist.\n#\n# The model will NOT use these values during inference.\n\nfor label in LABEL_COLUMNS:\n\n    test_inference_df[label] = 0.0\n\n\nprint(\"\\nTest dataframe after adding dummy labels:\")\nprint(\n    test_inference_df[\n        [\"StudyInstanceUID\"] + LABEL_COLUMNS\n    ]\n)\n\n\n# ============================================================\n# 3. VERIFY SHAPE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TEST DATAFRAME CHECK\")\nprint(\"=\" * 80)\n\nprint(\n    \"Shape:\",\n    test_inference_df.shape\n)\n\nprint(\n    \"Expected columns:\",\n    1 + len(LABEL_COLUMNS)\n)\n\nprint(\n    \"Actual columns:\",\n    len(test_inference_df.columns)\n)\n\n\nrequired_columns = (\n    [\"StudyInstanceUID\"]\n    + LABEL_COLUMNS\n)\n\nmissing_columns = [\n    c for c in required_columns\n    if c not in test_inference_df.columns\n]\n\nif len(missing_columns) > 0:\n\n    raise ValueError(\n        \"Missing columns:\\n\"\n        + str(missing_columns)\n    )\n\nprint(\"✓ All required columns present\")\n\n\n# ============================================================\n# 4. CREATE DATASET DIRECTLY\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CREATING TEST DATASET\")\nprint(\"=\" * 80)\n\ntest_dataset = RSNAKneeStudyDataset(\n    df=test_inference_df,\n    series_manifest=series_manifest,\n    label_columns=LABEL_COLUMNS,\n    target_slices=12,\n    image_size=256\n)\n\n\nprint(\n    \"✓ Test dataset created\"\n)\n\nprint(\n    \"Number of test studies:\",\n    len(test_dataset)\n)\n\n\n# ============================================================\n# 5. TEST FIRST SAMPLE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TEST SAMPLE\")\nprint(\"=\" * 80)\n\nsample = test_dataset[0]\n\nprint(\n    \"Sample type:\",\n    type(sample)\n)\n\nif isinstance(sample, (tuple, list)):\n\n    print(\n        \"Number of elements:\",\n        len(sample)\n    )\n\n    for i, item in enumerate(sample):\n\n        if torch.is_tensor(item):\n\n            print(\n                f\"Element {i}: \"\n                f\"Tensor \"\n                f\"shape={tuple(item.shape)} \"\n                f\"dtype={item.dtype}\"\n            )\n\n        elif isinstance(item, np.ndarray):\n\n            print(\n                f\"Element {i}: \"\n                f\"NumPy \"\n                f\"shape={item.shape}\"\n            )\n\n        else:\n\n            print(\n                f\"Element {i}: \"\n                f\"{type(item).__name__}\"\n            )\n\nelif isinstance(sample, dict):\n\n    print(\n        \"Dictionary keys:\",\n        list(sample.keys())\n    )\n\n\n# ============================================================\n# 6. CHECK EXPECTED IMAGE SHAPE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"IMAGE SHAPE CHECK\")\nprint(\"=\" * 80)\n\n\nif isinstance(sample, (tuple, list)):\n\n    image_sample = sample[0]\n\nelif isinstance(sample, dict):\n\n    if \"image\" in sample:\n\n        image_sample = sample[\"image\"]\n\n    elif \"images\" in sample:\n\n        image_sample = sample[\"images\"]\n\n    else:\n\n        raise ValueError(\n            \"Could not identify image \"\n            \"inside dataset output.\"\n        )\n\nelse:\n\n    raise ValueError(\n        \"Unsupported dataset output.\"\n    )\n\n\nif torch.is_tensor(image_sample):\n\n    image_shape = tuple(\n        image_sample.shape\n    )\n\nelif isinstance(image_sample, np.ndarray):\n\n    image_shape = image_sample.shape\n\nelse:\n\n    raise ValueError(\n        \"Image is neither Tensor nor NumPy array.\"\n    )\n\n\nprint(\n    \"Image shape:\",\n    image_shape\n)\n\nprint(\n    \"Expected:\",\n    \"(3, 12, 256, 256)\"\n)\n\n\nif image_shape != (\n    3,\n    12,\n    256,\n    256\n):\n\n    raise ValueError(\n        f\"Unexpected image shape: \"\n        f\"{image_shape}\"\n    )\n\nprint(\n    \"✓ Image tensor shape correct\"\n)\n\n\n# ============================================================\n# 7. CREATE TEST DATALOADER\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CREATING TEST DATALOADER\")\nprint(\"=\" * 80)\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    collate_fn=test_collate_fn\n)\n\n\nprint(\n    \"✓ Test DataLoader created\"\n)\n\nprint(\n    \"Number of batches:\",\n    len(test_loader)\n)\n\n\n# ============================================================\n# 8. TEST ONE BATCH\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"DATALOADER TEST\")\nprint(\"=\" * 80)\n\ntest_batch = next(\n    iter(test_loader)\n)\n\nimages, plane_mask, study_ids = (\n    test_batch\n)\n\nprint(\n    \"Images shape:\",\n    tuple(images.shape)\n)\n\nprint(\n    \"Plane mask shape:\",\n    tuple(plane_mask.shape)\n)\n\nprint(\n    \"Study IDs:\",\n    study_ids\n)\n\n\n# ============================================================\n# 9. VERIFY MODEL INPUT\n# ============================================================\n\nexpected_image_batch_shape = (\n    min(BATCH_SIZE, len(test_dataset)),\n    3,\n    12,\n    256,\n    256\n)\n\nprint(\n    \"\\nExpected first batch:\",\n    expected_image_batch_shape\n)\n\n\nif tuple(images.shape) != (\n    min(BATCH_SIZE, len(test_dataset)),\n    3,\n    12,\n    256,\n    256\n):\n\n    raise ValueError(\n        \"Incorrect model input shape: \"\n        + str(tuple(images.shape))\n    )\n\n\nif tuple(plane_mask.shape) != (\n    min(BATCH_SIZE, len(test_dataset)),\n    3\n):\n\n    raise ValueError(\n        \"Incorrect plane mask shape: \"\n        + str(tuple(plane_mask.shape))\n    )\n\n\nprint(\n    \"✓ Model input shape correct\"\n)\n\n\n# ============================================================\n# 10. TEST MODEL FORWARD PASS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"MODEL FORWARD TEST\")\nprint(\"=\" * 80)\n\n\nmodel.eval()\n\nwith torch.inference_mode():\n\n    test_images = images.to(\n        DEVICE,\n        non_blocking=True\n    )\n\n    test_plane_mask = plane_mask.to(\n        DEVICE,\n        non_blocking=True\n    )\n\n    test_logits = model(\n        test_images,\n        test_plane_mask\n    )\n\n    test_probabilities = torch.sigmoid(\n        test_logits\n    )\n\n\nprint(\n    \"Logits shape:\",\n    tuple(test_logits.shape)\n)\n\nprint(\n    \"Probability shape:\",\n    tuple(test_probabilities.shape)\n)\n\nprint(\n    \"Expected:\",\n    (\n        min(BATCH_SIZE, len(test_dataset)),\n        12\n    )\n)\n\n\nif tuple(test_logits.shape) != (\n    min(BATCH_SIZE, len(test_dataset)),\n    12\n):\n\n    raise ValueError(\n        \"Model output shape is incorrect.\"\n    )\n\n\nprint(\n    \"✓ Model forward pass successful\"\n)\n\n\n# ============================================================\n# 11. PROBABILITY SANITY CHECK\n# ============================================================\n\ntest_probs_np = (\n    test_probabilities\n    .detach()\n    .cpu()\n    .numpy()\n)\n\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"PROBABILITY CHECK\")\nprint(\"=\" * 80)\n\nprint(\n    \"Minimum:\",\n    test_probs_np.min()\n)\n\nprint(\n    \"Maximum:\",\n    test_probs_np.max()\n)\n\nprint(\n    \"Mean:\",\n    test_probs_np.mean()\n)\n\n\nif not np.isfinite(\n    test_probs_np\n).all():\n\n    raise ValueError(\n        \"NaN/Inf detected in probabilities.\"\n    )\n\n\nif (\n    test_probs_np.min() < 0\n    or\n    test_probs_np.max() > 1\n):\n\n    raise ValueError(\n        \"Probability outside [0,1].\"\n    )\n\n\nprint(\n    \"✓ Probabilities valid\"\n)\n\n\n# ============================================================\n# FINAL\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 78 DATASET FIX COMPLETE\")\nprint(\"=\" * 80)\n\nprint(\n    \"✓ Test dataframe prepared\"\n)\n\nprint(\n    \"✓ Dummy labels added only for Dataset compatibility\"\n)\n\nprint(\n    \"✓ Test dataset created\"\n)\n\nprint(\n    \"✓ First study loaded\"\n)\n\nprint(\n    \"✓ Image shape verified\"\n)\n\nprint(\n    \"✓ DataLoader created\"\n)\n\nprint(\n    \"✓ Model input verified\"\n)\n\nprint(\n    \"✓ Model forward pass verified\"\n)\n\nprint(\n    \"✓ Sigmoid probabilities verified\"\n)\n\nprint(\n    \"\\nREADY FOR TEST INFERENCE\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T06:39:38.124221Z","iopub.execute_input":"2026-08-19T06:39:38.124514Z","iopub.status.idle":"2026-08-19T06:39:38.163491Z","shell.execute_reply.started":"2026-08-19T06:39:38.12449Z","shell.execute_reply":"2026-08-19T06:39:38.162316Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 78B — DIAGNOSE TEST DICOM MANIFEST\n# ============================================================\n\nimport os\nimport glob\nimport pandas as pd\nimport numpy as np\n\nprint(\"=\" * 80)\nprint(\"STEP 78B — TEST DICOM / MANIFEST DIAGNOSTIC\")\nprint(\"=\" * 80)\n\n\n# ============================================================\n# 1. TEST STUDY IDs\n# ============================================================\n\ntest_ids = (\n    test_df[\"StudyInstanceUID\"]\n    .astype(str)\n    .tolist()\n)\n\nprint(\"\\nTEST STUDY IDs\")\nprint(\"-\" * 80)\n\nfor i, uid in enumerate(test_ids):\n\n    print(\n        f\"{i}: {uid}\"\n    )\n\n\n# ============================================================\n# 2. INSPECT SERIES MANIFEST\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"SERIES MANIFEST COLUMNS\")\nprint(\"=\" * 80)\n\nprint(\n    list(series_manifest.columns)\n)\n\nprint(\n    \"\\nManifest shape:\",\n    series_manifest.shape\n)\n\n\n# ============================================================\n# 3. FIND STUDY INSTANCE UID COLUMN\n# ============================================================\n\nuid_candidates = [\n    c for c in series_manifest.columns\n    if (\n        \"studyinstanceuid\" in c.lower()\n        or c.lower() == \"study_uid\"\n        or c.lower() == \"studyid\"\n    )\n]\n\nprint(\n    \"\\nPossible Study UID columns:\",\n    uid_candidates\n)\n\n\nif len(uid_candidates) == 0:\n\n    raise ValueError(\n        \"Could not identify StudyInstanceUID \"\n        \"column in series_manifest.\"\n    )\n\n\nMANIFEST_UID_COLUMN = uid_candidates[0]\n\nprint(\n    \"Using manifest UID column:\",\n    MANIFEST_UID_COLUMN\n)\n\n\n# ============================================================\n# 4. CHECK WHETHER TEST STUDIES EXIST\n# ============================================================\n\nmanifest_uids = (\n    series_manifest[\n        MANIFEST_UID_COLUMN\n    ]\n    .astype(str)\n    .unique()\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TEST UID → MANIFEST MATCH\")\nprint(\"=\" * 80)\n\nfor uid in test_ids:\n\n    found = uid in manifest_uids\n\n    print(\n        f\"{found!s:5} : {uid}\"\n    )\n\n\n# ============================================================\n# 5. COUNT MATCHES\n# ============================================================\n\nmatching_test_ids = [\n    uid for uid in test_ids\n    if uid in manifest_uids\n]\n\nmissing_test_ids = [\n    uid for uid in test_ids\n    if uid not in manifest_uids\n]\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"MATCH SUMMARY\")\nprint(\"=\" * 80)\n\nprint(\n    \"Test studies:\",\n    len(test_ids)\n)\n\nprint(\n    \"Found in manifest:\",\n    len(matching_test_ids)\n)\n\nprint(\n    \"Missing from manifest:\",\n    len(missing_test_ids)\n)\n\n\n# ============================================================\n# 6. SHOW MANIFEST SAMPLE\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"MANIFEST SAMPLE\")\nprint(\"=\" * 80)\n\ndisplay(\n    series_manifest.head(5)\n)\n\n\n# ============================================================\n# 7. INSPECT KAGGLE INPUT DIRECTORIES\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"KAGGLE INPUT DIRECTORIES\")\nprint(\"=\" * 80)\n\nINPUT_ROOT = \"/kaggle/input\"\n\nfor root, dirs, files in os.walk(INPUT_ROOT):\n\n    level = root.replace(\n        INPUT_ROOT,\n        \"\"\n    ).count(os.sep)\n\n    if level > 3:\n        dirs[:] = []\n\n        continue\n\n    indent = \"  \" * level\n\n    print(\n        indent\n        + os.path.basename(root)\n        + \"/\"\n    )\n\n    for file in files[:10]:\n\n        print(\n            indent\n            + \"  \"\n            + file\n        )\n\n\n# ============================================================\n# 8. SEARCH FOR TEST STUDY UID IN KAGGLE FILESYSTEM\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"SEARCHING FOR TEST STUDY DIRECTORIES\")\nprint(\"=\" * 80)\n\n\nfor uid in test_ids:\n\n    print(\n        f\"\\nSearching for:\\n{uid}\"\n    )\n\n    matches = []\n\n    for base in [\n        \"/kaggle/input\",\n        \"/kaggle/working\"\n    ]:\n\n        for root, dirs, files in os.walk(base):\n\n            # Exact directory-name match\n            if os.path.basename(root) == uid:\n\n                matches.append(root)\n\n            # Avoid unnecessarily deep search\n            if root.count(os.sep) > 8:\n\n                dirs[:] = []\n\n    if len(matches) == 0:\n\n        print(\n            \"❌ No exact UID directory found.\"\n        )\n\n    else:\n\n        print(\n            \"✓ Found:\"\n        )\n\n        for match in matches[:10]:\n\n            print(\n                \" \",\n                match\n            )\n\n\n# ============================================================\n# 9. SEARCH FOR DICOM FILES\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"DICOM FILE SEARCH\")\nprint(\"=\" * 80)\n\n\ndicom_candidates = []\n\nfor base in [\n    \"/kaggle/input\",\n    \"/kaggle/working\"\n]:\n\n    for pattern in [\n        \"**/*.dcm\",\n        \"**/*.DCM\"\n    ]:\n\n        found = glob.glob(\n            os.path.join(\n                base,\n                pattern\n            ),\n            recursive=True\n        )\n\n        dicom_candidates.extend(\n            found\n        )\n\n\ndicom_candidates = list(\n    dict.fromkeys(\n        dicom_candidates\n    )\n)\n\n\nprint(\n    \"DICOM files found:\",\n    len(dicom_candidates)\n)\n\n\nif len(dicom_candidates) > 0:\n\n    print(\n        \"\\nFirst DICOM files:\"\n    )\n\n    for path in dicom_candidates[:20]:\n\n        print(\n            path\n        )\n\nelse:\n\n    print(\n        \"\\n⚠ No .dcm files found.\"\n    )\n\n    print(\n        \"The competition data may be stored \"\n        \"without a .dcm extension.\"\n    )\n\n\n# ============================================================\n# 10. CHECK MANIFEST UID FORMAT\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"UID FORMAT CHECK\")\nprint(\"=\" * 80)\n\nprint(\n    \"Example test UID:\"\n)\n\nprint(\n    test_ids[0]\n)\n\nprint(\n    \"\\nExample manifest UID:\"\n)\n\nprint(\n    str(\n        series_manifest[\n            MANIFEST_UID_COLUMN\n        ].iloc[0]\n    )\n)\n\nprint(\n    \"\\nTest UID length:\",\n    len(test_ids[0])\n)\n\nprint(\n    \"Manifest UID length:\",\n    len(\n        str(\n            series_manifest[\n                MANIFEST_UID_COLUMN\n            ].iloc[0]\n        )\n    )\n)\n\n\n# ============================================================\n# FINAL DIAGNOSIS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"STEP 78B DIAGNOSIS\")\nprint(\"=\" * 80)\n\nif len(missing_test_ids) == 0:\n\n    print(\n        \"✓ All test studies exist in the manifest.\"\n    )\n\n    print(\n        \"The problem is elsewhere in dataset loading.\"\n    )\n\nelif len(dicom_candidates) == 0:\n\n    print(\n        \"⚠ Test studies are missing from the current \"\n        \"manifest and no .dcm files were found.\"\n    )\n\n    print(\n        \"We need to identify the actual Kaggle \"\n        \"test-data storage format.\"\n    )\n\nelse:\n\n    print(\n        \"⚠ Test studies are missing from the \"\n        \"current manifest.\"\n    )\n\n    print(\n        \"DICOM files exist in the Kaggle filesystem.\"\n    )\n\n    print(\n        \"Next step: build a TEST-ONLY series manifest \"\n        \"from the actual DICOM files.\"\n    )\n\nprint(\n    \"\\nDo NOT modify the trained model.\"\n)\n\nprint(\n    \"Do NOT add fake manifest rows.\"\n)\n\nprint(\n    \"Do NOT use dummy StudyInstanceUIDs.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T06:40:23.55984Z","iopub.execute_input":"2026-08-19T06:40:23.560804Z","iopub.status.idle":"2026-08-19T06:51:17.334355Z","shell.execute_reply.started":"2026-08-19T06:40:23.560764Z","shell.execute_reply":"2026-08-19T06:51:17.333487Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# STEP 78C — BUILD TEST SERIES MANIFEST\n# ============================================================\n\nimport os\nimport pandas as pd\nimport numpy as np\n\nprint(\"=\" * 80)\nprint(\"STEP 78C — BUILD TEST SERIES MANIFEST\")\nprint(\"=\" * 80)\n\n\n# ============================================================\n# CONFIGURATION\n# ============================================================\n\nTEST_ROOT = (\n    \"/kaggle/input/competitions/\"\n    \"rsna-knee-abnormality-detection/\"\n    \"test_series\"\n)\n\nTEST_SERIES_CSV = (\n    \"/kaggle/input/competitions/\"\n    \"rsna-knee-abnormality-detection/\"\n    \"test_series.csv\"\n)\n\n\n# ============================================================\n# CHECK FILES\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TEST DATA CHECK\")\nprint(\"=\" * 80)\n\nprint(\n    \"Test series CSV:\",\n    TEST_SERIES_CSV\n)\n\nprint(\n    \"Test series directory:\",\n    TEST_ROOT\n)\n\nif not os.path.exists(TEST_SERIES_CSV):\n    raise FileNotFoundError(\n        f\"test_series.csv not found:\\n{TEST_SERIES_CSV}\"\n    )\n\nif not os.path.exists(TEST_ROOT):\n    raise FileNotFoundError(\n        f\"test_series directory not found:\\n{TEST_ROOT}\"\n    )\n\nprint(\"✓ test_series.csv found\")\nprint(\"✓ test_series directory found\")\n\n\n# ============================================================\n# LOAD TEST SERIES CSV\n# ============================================================\n\ntest_series_df = pd.read_csv(\n    TEST_SERIES_CSV\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TEST SERIES CSV\")\nprint(\"=\" * 80)\n\nprint(\n    \"Shape:\",\n    test_series_df.shape\n)\n\nprint(\n    \"Columns:\",\n    list(test_series_df.columns)\n)\n\ndisplay(\n    test_series_df.head(10)\n)\n\n\n# ============================================================\n# CHECK REQUIRED COLUMNS\n# ============================================================\n\nrequired_columns = [\n    \"StudyInstanceUID\",\n    \"SeriesInstanceUID\"\n]\n\nmissing_columns = [\n    c for c in required_columns\n    if c not in test_series_df.columns\n]\n\nif len(missing_columns) > 0:\n\n    raise ValueError(\n        \"Missing required columns:\\n\"\n        + str(missing_columns)\n    )\n\nprint(\n    \"\\n✓ Required UID columns found\"\n)\n\n\n# ============================================================\n# NORMALIZE UID TYPES\n# ============================================================\n\ntest_series_df[\n    \"StudyInstanceUID\"\n] = (\n    test_series_df[\n        \"StudyInstanceUID\"\n    ]\n    .astype(str)\n    .str.strip()\n)\n\ntest_series_df[\n    \"SeriesInstanceUID\"\n] = (\n    test_series_df[\n        \"SeriesInstanceUID\"\n    ]\n    .astype(str)\n    .str.strip()\n)\n\n\n# ============================================================\n# VERIFY TEST STUDIES\n# ============================================================\n\ntest_study_ids = (\n    test_df[\n        \"StudyInstanceUID\"\n    ]\n    .astype(str)\n    .str.strip()\n    .tolist()\n)\n\ntest_series_study_ids = set(\n    test_series_df[\n        \"StudyInstanceUID\"\n    ].tolist()\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TEST STUDY COVERAGE\")\nprint(\"=\" * 80)\n\nfor uid in test_study_ids:\n\n    count = (\n        test_series_df[\n            \"StudyInstanceUID\"\n        ] == uid\n    ).sum()\n\n    print(\n        f\"{uid}  →  {count} series\"\n    )\n\n    if count == 0:\n\n        raise ValueError(\n            f\"Test study has no series: {uid}\"\n        )\n\n\nprint(\n    \"\\n✓ All test studies have series\"\n)\n\n\n# ============================================================\n# BUILD TEST SERIES PATH\n# ============================================================\n\ntest_series_df[\"SeriesPath\"] = (\n    TEST_ROOT\n    + \"/\"\n    + test_series_df[\n        \"StudyInstanceUID\"\n    ].astype(str)\n    + \"/\"\n    + test_series_df[\n        \"SeriesInstanceUID\"\n    ].astype(str)\n)\n\n\n# ============================================================\n# CHECK SERIES DIRECTORIES\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CHECKING SERIES DIRECTORIES\")\nprint(\"=\" * 80)\n\nmissing_paths = []\n\nfor path in test_series_df[\n    \"SeriesPath\"\n].unique():\n\n    if not os.path.isdir(path):\n\n        missing_paths.append(path)\n\n\nprint(\n    \"Total series:\",\n    len(test_series_df)\n)\n\nprint(\n    \"Missing series directories:\",\n    len(missing_paths)\n)\n\n\nif len(missing_paths) > 0:\n\n    print(\n        \"\\nFirst missing paths:\"\n    )\n\n    for path in missing_paths[:10]:\n\n        print(\n            path\n        )\n\n    raise FileNotFoundError(\n        \"Some test series directories are missing.\"\n    )\n\n\nprint(\n    \"✓ All test series directories exist\"\n)\n\n\n# ============================================================\n# COUNT ACTUAL DICOM SLICES\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"COUNTING TEST DICOM SLICES\")\nprint(\"=\" * 80)\n\nimport pydicom\n\n\ndef count_dicom_files(series_path):\n\n    count = 0\n\n    try:\n\n        for filename in os.listdir(\n            series_path\n        ):\n\n            filepath = os.path.join(\n                series_path,\n                filename\n            )\n\n            if not os.path.isfile(filepath):\n                continue\n\n            try:\n\n                # Read only metadata.\n                pydicom.dcmread(\n                    filepath,\n                    stop_before_pixels=True,\n                    force=True\n                )\n\n                count += 1\n\n            except Exception:\n\n                continue\n\n    except Exception:\n\n        return 0\n\n    return count\n\n\nslice_counts = []\n\nfor i, row in test_series_df.iterrows():\n\n    path = row[\"SeriesPath\"]\n\n    count = count_dicom_files(\n        path\n    )\n\n    slice_counts.append(\n        count\n    )\n\n    print(\n        f\"{i + 1:3d}/\"\n        f\"{len(test_series_df):3d}  \"\n        f\"{count:4d} slices  \"\n        f\"{row['StudyInstanceUID'][:30]}...\"\n    )\n\n\ntest_series_df[\n    \"Number_of_Slices\"\n] = slice_counts\n\n\n# ============================================================\n# CHECK SLICE COUNTS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TEST SLICE COUNT SUMMARY\")\nprint(\"=\" * 80)\n\nprint(\n    test_series_df[\n        \"Number_of_Slices\"\n    ].describe()\n)\n\nzero_slice_series = (\n    test_series_df[\n        \"Number_of_Slices\"\n    ] <= 0\n).sum()\n\nprint(\n    \"\\nZero-slice series:\",\n    zero_slice_series\n)\n\nif zero_slice_series > 0:\n\n    display(\n        test_series_df[\n            test_series_df[\n                \"Number_of_Slices\"\n            ] <= 0\n        ]\n    )\n\n    raise ValueError(\n        \"Some test series contain zero readable DICOM slices.\"\n    )\n\n\n# ============================================================\n# DETECT / INFER SERIES CHARACTERISTICS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TEST SERIES METADATA\")\nprint(\"=\" * 80)\n\n\n# ------------------------------------------------------------\n# Use the existing metadata detection function if available.\n# ------------------------------------------------------------\n\nmetadata_function_names = [\n    \"extract_series_metadata\",\n    \"get_series_metadata\",\n    \"analyze_series\",\n    \"classify_series\"\n]\n\nmetadata_function = None\n\nfor name in metadata_function_names:\n\n    if name in globals():\n\n        candidate = globals()[name]\n\n        if callable(candidate):\n\n            metadata_function = candidate\n\n            print(\n                \"Using existing function:\",\n                name\n            )\n\n            break\n\n\n# ------------------------------------------------------------\n# If no existing function is available, inspect one DICOM\n# series and derive basic metadata.\n# ------------------------------------------------------------\n\nif metadata_function is None:\n\n    print(\n        \"No existing metadata classifier found.\"\n    )\n\n    print(\n        \"Using DICOM metadata fallback.\"\n    )\n\n\n    def inspect_series_metadata(\n        series_path\n    ):\n\n        files = []\n\n        for filename in os.listdir(\n            series_path\n        ):\n\n            filepath = os.path.join(\n                series_path,\n                filename\n            )\n\n            if os.path.isfile(filepath):\n\n                files.append(filepath)\n\n\n        if len(files) == 0:\n\n            return {\n                \"Fluid_Sensitive\": 0,\n                \"Fat_Suppression\": 0,\n                \"Anatomical_Plane\": \"Unknown\"\n            }\n\n\n        # Try several files until one reads.\n\n        ds = None\n\n        for filepath in files:\n\n            try:\n\n                ds = pydicom.dcmread(\n                    filepath,\n                    stop_before_pixels=True,\n                    force=True\n                )\n\n                break\n\n            except Exception:\n\n                continue\n\n\n        if ds is None:\n\n            return {\n                \"Fluid_Sensitive\": 0,\n                \"Fat_Suppression\": 0,\n                \"Anatomical_Plane\": \"Unknown\"\n            }\n\n\n        description_parts = []\n\n        for attr in [\n            \"SeriesDescription\",\n            \"ProtocolName\",\n            \"SequenceName\",\n            \"ScanningSequence\",\n            \"ImageType\"\n        ]:\n\n            if hasattr(ds, attr):\n\n                value = getattr(\n                    ds,\n                    attr\n                )\n\n                if isinstance(\n                    value,\n                    (list, tuple)\n                ):\n\n                    value = \" \".join(\n                        map(str, value)\n                    )\n\n                description_parts.append(\n                    str(value)\n                )\n\n\n        description = \" \".join(\n            description_parts\n        ).lower()\n\n\n        # ----------------------------------------------------\n        # Plane detection\n        # ----------------------------------------------------\n\n        plane = \"Unknown\"\n\n        if any(\n            x in description\n            for x in [\n                \"sag\",\n                \"sagittal\"\n            ]\n        ):\n\n            plane = \"Sagittal\"\n\n        elif any(\n            x in description\n            for x in [\n                \"cor\",\n                \"coronal\"\n            ]\n        ):\n\n            plane = \"Coronal\"\n\n        elif any(\n            x in description\n            for x in [\n                \"ax\",\n                \"axial\",\n                \"trans\"\n            ]\n        ):\n\n            plane = \"Axial\"\n\n\n        # ----------------------------------------------------\n        # Fat suppression\n        # ----------------------------------------------------\n\n        fat_suppression = int(\n            any(\n                x in description\n                for x in [\n                    \"fs\",\n                    \"fat sat\",\n                    \"fatsat\",\n                    \"fat-sat\",\n                    \"spectral\"\n                ]\n            )\n        )\n\n\n        # ----------------------------------------------------\n        # Fluid-sensitive sequence\n        # ----------------------------------------------------\n\n        fluid_sensitive = int(\n            any(\n                x in description\n                for x in [\n                    \"pd\",\n                    \"t2\",\n                    \"stir\",\n                    \"fs\",\n                    \"fat sat\",\n                    \"fatsat\"\n                ]\n            )\n        )\n\n\n        return {\n            \"Fluid_Sensitive\":\n                fluid_sensitive,\n\n            \"Fat_Suppression\":\n                fat_suppression,\n\n            \"Anatomical_Plane\":\n                plane\n        }\n\n\n# ============================================================\n# APPLY METADATA\n# ============================================================\n\nmetadata_results = []\n\nfor i, row in test_series_df.iterrows():\n\n    path = row[\"SeriesPath\"]\n\n    try:\n\n        if metadata_function is not None:\n\n            result = metadata_function(\n                path\n            )\n\n        else:\n\n            result = inspect_series_metadata(\n                path\n            )\n\n    except Exception as e:\n\n        print(\n            f\"Metadata warning for row {i}:\",\n            repr(e)\n        )\n\n        result = {\n            \"Fluid_Sensitive\": 0,\n            \"Fat_Suppression\": 0,\n            \"Anatomical_Plane\": \"Unknown\"\n        }\n\n\n    metadata_results.append(\n        result\n    )\n\n\n# ============================================================\n# ADD METADATA COLUMNS\n# ============================================================\n\ntest_series_df[\n    \"Fluid_Sensitive\"\n] = [\n    r.get(\n        \"Fluid_Sensitive\",\n        0\n    )\n    if isinstance(r, dict)\n    else 0\n    for r in metadata_results\n]\n\ntest_series_df[\n    \"Fat_Suppression\"\n] = [\n    r.get(\n        \"Fat_Suppression\",\n        0\n    )\n    if isinstance(r, dict)\n    else 0\n    for r in metadata_results\n]\n\ntest_series_df[\n    \"Anatomical_Plane\"\n] = [\n    r.get(\n        \"Anatomical_Plane\",\n        \"Unknown\"\n    )\n    if isinstance(r, dict)\n    else \"Unknown\"\n    for r in metadata_results\n]\n\n\n# ============================================================\n# REORDER COLUMNS\n# ============================================================\n\ntest_series_manifest = test_series_df[\n    [\n        \"StudyInstanceUID\",\n        \"SeriesInstanceUID\",\n        \"Fluid_Sensitive\",\n        \"Fat_Suppression\",\n        \"Anatomical_Plane\",\n        \"Number_of_Slices\",\n        \"SeriesPath\"\n    ]\n].copy()\n\n\n# ============================================================\n# FINAL TEST MANIFEST\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TEST SERIES MANIFEST\")\nprint(\"=\" * 80)\n\nprint(\n    \"Studies:\",\n    test_series_manifest[\n        \"StudyInstanceUID\"\n    ].nunique()\n)\n\nprint(\n    \"Series:\",\n    len(test_series_manifest)\n)\n\nprint(\n    \"\\nAnatomical planes:\"\n)\n\nprint(\n    test_series_manifest[\n        \"Anatomical_Plane\"\n    ].value_counts(\n        dropna=False\n    )\n)\n\n\ndisplay(\n    test_series_manifest\n)\n\n\n# ============================================================\n# SAVE TEST MANIFEST\n# ============================================================\n\nTEST_MANIFEST_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_test_series_manifest.csv\"\n)\n\ntest_series_manifest.to_csv(\n    TEST_MANIFEST_PATH,\n    index=False\n)\n\nprint(\n    \"\\n✓ Test manifest saved:\"\n)\n\nprint(\n    TEST_MANIFEST_PATH\n)\n\n\n# ============================================================\n# VERIFY EXACT TEST COVERAGE\n# ============================================================\n\nmanifest_test_ids = set(\n    test_series_manifest[\n        \"StudyInstanceUID\"\n    ].astype(str)\n)\n\nmissing_after_build = [\n    uid\n    for uid in test_study_ids\n    if uid not in manifest_test_ids\n]\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"FINAL COVERAGE CHECK\")\nprint(\"=\" * 80)\n\nprint(\n    \"Expected test studies:\",\n    len(test_study_ids)\n)\n\nprint(\n    \"Manifest test studies:\",\n    len(manifest_test_ids)\n)\n\nprint(\n    \"Missing test studies:\",\n    len(missing_after_build)\n)\n\n\nif len(missing_after_build) > 0:\n\n    print(\n        \"\\n❌ Missing:\"\n    )\n\n    for uid in missing_after_build:\n\n        print(\n            uid\n        )\n\n    raise ValueError(\n        \"Test manifest does not cover \"\n        \"all test studies.\"\n    )\n\n\nprint(\n    \"\\n✓ ALL TEST STUDIES FOUND\"\n)\n\nprint(\n    \"\\n\" + \"=\" * 80\n)\n\nprint(\n    \"STEP 78C COMPLETE\"\n)\n\nprint(\n    \"✓ test_series.csv loaded\"\n)\n\nprint(\n    \"✓ Test series paths created\"\n)\n\nprint(\n    \"✓ DICOM slices counted\"\n)\n\nprint(\n    \"✓ Test series metadata generated\"\n)\n\nprint(\n    \"✓ Test manifest created\"\n)\n\nprint(\n    \"✓ All 3 test studies covered\"\n)\n\nprint(\n    \"\\nREADY FOR TEST DATASET INFERENCE\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T07:03:39.473609Z","iopub.execute_input":"2026-08-19T07:03:39.473887Z","iopub.status.idle":"2026-08-19T07:03:40.267971Z","shell.execute_reply.started":"2026-08-19T07:03:39.473862Z","shell.execute_reply":"2026-08-19T07:03:40.267409Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# STEP 78D — FINAL TEST DATASET + MODEL INFERENCE\n# =============================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import DataLoader\n\nprint(\"=\" * 80)\nprint(\"STEP 78D — FINAL TEST DATASET + MODEL INFERENCE\")\nprint(\"=\" * 80)\n\n# -----------------------------------------------------------------------------\n# CONFIGURATION\n# -----------------------------------------------------------------------------\n\nTEST_CSV = \"/kaggle/input/competitions/rsna-knee-abnormality-detection/test.csv\"\n\nTEST_MANIFEST = (\n    \"/kaggle/working/rsna_knee_test_series_manifest.csv\"\n)\n\nCHECKPOINT = \"/kaggle/working/best_rsna_knee_model.pth\"\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nTARGET_SLICES = 12\nIMAGE_SIZE = 256\nBATCH_SIZE = 1\n\nLABEL_COLUMNS = [\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\nPLANE_ORDER = [\n    \"Sagittal\",\n    \"Axial\",\n    \"Coronal\",\n]\n\nprint()\nprint(\"=\" * 80)\nprint(\"CONFIGURATION\")\nprint(\"=\" * 80)\n\nprint(\"Checkpoint :\", CHECKPOINT)\nprint(\"Test CSV   :\", TEST_CSV)\nprint(\"Manifest   :\", TEST_MANIFEST)\nprint(\"Device     :\", DEVICE)\nprint(\"Labels     :\", len(LABEL_COLUMNS))\nprint(\"Plane order:\", PLANE_ORDER)\n\n\n# -----------------------------------------------------------------------------\n# FILE CHECK\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"FILE CHECK\")\nprint(\"=\" * 80)\n\nif not os.path.exists(CHECKPOINT):\n    raise FileNotFoundError(\n        f\"Checkpoint not found:\\n{CHECKPOINT}\"\n    )\n\nif not os.path.exists(TEST_CSV):\n    raise FileNotFoundError(\n        f\"Test CSV not found:\\n{TEST_CSV}\"\n    )\n\nif not os.path.exists(TEST_MANIFEST):\n    raise FileNotFoundError(\n        f\"Test manifest not found:\\n{TEST_MANIFEST}\"\n    )\n\nprint(\"✓ Checkpoint found\")\nprint(\"✓ test.csv found\")\nprint(\"✓ Test manifest found\")\n\n\n# -----------------------------------------------------------------------------\n# LOAD TEST CSV\n# -----------------------------------------------------------------------------\n\ntest_df = pd.read_csv(TEST_CSV)\n\nif \"StudyInstanceUID\" not in test_df.columns:\n    raise ValueError(\n        \"StudyInstanceUID column missing from test.csv\"\n    )\n\ntest_df[\"StudyInstanceUID\"] = (\n    test_df[\"StudyInstanceUID\"]\n    .astype(str)\n    .str.strip()\n)\n\ntest_df = test_df.drop_duplicates(\n    subset=[\"StudyInstanceUID\"]\n).reset_index(drop=True)\n\nprint()\nprint(\"=\" * 80)\nprint(\"TEST DATA\")\nprint(\"=\" * 80)\n\nprint(\"Test studies:\", len(test_df))\n\nprint()\nprint(test_df)\n\n\n# -----------------------------------------------------------------------------\n# ADD DUMMY LABELS\n# Dataset class requires label columns, but these labels are NOT used.\n# -----------------------------------------------------------------------------\n\nfor col in LABEL_COLUMNS:\n    test_df[col] = 0.0\n\nprint()\nprint(\"✓ Dummy label columns added for dataset compatibility\")\n\n\n# -----------------------------------------------------------------------------\n# LOAD TEST MANIFEST\n# -----------------------------------------------------------------------------\n\ntest_manifest = pd.read_csv(TEST_MANIFEST)\n\ntest_manifest[\"StudyInstanceUID\"] = (\n    test_manifest[\"StudyInstanceUID\"]\n    .astype(str)\n    .str.strip()\n)\n\ntest_manifest[\"SeriesInstanceUID\"] = (\n    test_manifest[\"SeriesInstanceUID\"]\n    .astype(str)\n    .str.strip()\n)\n\nprint()\nprint(\"=\" * 80)\nprint(\"TEST MANIFEST\")\nprint(\"=\" * 80)\n\nprint(\"Manifest shape:\", test_manifest.shape)\nprint(\"Manifest series:\", len(test_manifest))\nprint(\n    \"Manifest studies:\",\n    test_manifest[\"StudyInstanceUID\"].nunique()\n)\n\n\n# -----------------------------------------------------------------------------\n# VERIFY STUDY COVERAGE\n# -----------------------------------------------------------------------------\n\ntest_ids = set(test_df[\"StudyInstanceUID\"])\n\nmanifest_ids = set(\n    test_manifest[\"StudyInstanceUID\"]\n)\n\nmissing_ids = test_ids - manifest_ids\n\nprint()\nprint(\"=\" * 80)\nprint(\"STUDY COVERAGE\")\nprint(\"=\" * 80)\n\nprint(\"Expected studies :\", len(test_ids))\nprint(\"Manifest studies  :\", len(manifest_ids))\nprint(\"Missing studies   :\", len(missing_ids))\n\nif missing_ids:\n    print()\n    print(\"Missing:\")\n    for sid in sorted(missing_ids):\n        print(sid)\n\n    raise RuntimeError(\n        \"Some test studies are missing from the manifest.\"\n    )\n\nprint(\"✓ ALL TEST STUDIES FOUND\")\n\n\n# -----------------------------------------------------------------------------\n# IMPORTANT:\n# Keep only valid anatomical planes.\n#\n# Unknown series must NOT be treated as a real plane.\n# -----------------------------------------------------------------------------\n\nvalid_planes = set(PLANE_ORDER)\n\nunknown_count = (\n    ~test_manifest[\"Anatomical_Plane\"]\n    .isin(valid_planes)\n).sum()\n\nprint()\nprint(\"=\" * 80)\nprint(\"PLANE CHECK\")\nprint(\"=\" * 80)\n\nprint(\"Unknown/non-model series:\", unknown_count)\n\nprint()\nprint(\n    test_manifest[\"Anatomical_Plane\"]\n    .value_counts(dropna=False)\n)\n\n\n# -----------------------------------------------------------------------------\n# DO NOT MODIFY THE ORIGINAL MANIFEST\n# Create a clean test manifest for inference.\n# -----------------------------------------------------------------------------\n\ninference_manifest = test_manifest[\n    test_manifest[\"Anatomical_Plane\"].isin(valid_planes)\n].copy()\n\nprint()\nprint(\n    \"Usable model-plane series:\",\n    len(inference_manifest)\n)\n\n\n# -----------------------------------------------------------------------------\n# CHECK EACH TEST STUDY\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"PER-STUDY PLANE COVERAGE\")\nprint(\"=\" * 80)\n\nfor study_id in test_df[\"StudyInstanceUID\"]:\n\n    study_rows = inference_manifest[\n        inference_manifest[\"StudyInstanceUID\"] == study_id\n    ]\n\n    available = (\n        study_rows[\"Anatomical_Plane\"]\n        .dropna()\n        .unique()\n        .tolist()\n    )\n\n    print()\n    print(\"Study:\")\n    print(study_id)\n\n    print(\"Available:\", available)\n\n    for plane in PLANE_ORDER:\n\n        n = (\n            study_rows[\"Anatomical_Plane\"]\n            == plane\n        ).sum()\n\n        print(\n            f\"  {plane:<10}: {n} series\"\n        )\n\n\n# -----------------------------------------------------------------------------\n# CREATE DATASET\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"CREATING TEST DATASET\")\nprint(\"=\" * 80)\n\n# IMPORTANT:\n# Use the SAME dataset class used during training.\nDatasetClass = RSNAKneeStudyDataset\n\ntest_dataset = DatasetClass(\n    df=test_df,\n    series_manifest=inference_manifest,\n    label_columns=LABEL_COLUMNS,\n    target_slices=TARGET_SLICES,\n    image_size=IMAGE_SIZE,\n)\n\nprint(\"✓ Test dataset created\")\nprint(\"Number of studies:\", len(test_dataset))\n\n\n# -----------------------------------------------------------------------------\n# TEST EACH STUDY BEFORE MODEL INFERENCE\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"TESTING PREPROCESSING\")\nprint(\"=\" * 80)\n\nfor i in range(len(test_dataset)):\n\n    study_id = test_df.iloc[i][\"StudyInstanceUID\"]\n\n    sample = test_dataset[i]\n\n    # Dataset may return tuple\n    images = sample[0]\n    plane_mask = sample[1]\n\n    if torch.is_tensor(images):\n        image_shape = tuple(images.shape)\n        image_min = float(images.min())\n        image_max = float(images.max())\n        image_dtype = images.dtype\n    else:\n        image_shape = np.asarray(images).shape\n        image_min = float(np.min(images))\n        image_max = float(np.max(images))\n        image_dtype = np.asarray(images).dtype\n\n    print()\n    print(f\"TEST STUDY {i + 1}\")\n    print(\"-\" * 60)\n    print(\"Study ID   :\", study_id)\n    print(\"Tensor     :\", image_shape)\n    print(\"Plane mask :\", plane_mask)\n    print(\"Dtype      :\", image_dtype)\n    print(\"Minimum    :\", image_min)\n    print(\"Maximum    :\", image_max)\n\n    expected_shape = (\n        3,\n        TARGET_SLICES,\n        IMAGE_SIZE,\n        IMAGE_SIZE,\n    )\n\n    if image_shape != expected_shape:\n\n        raise ValueError(\n            f\"\\nIncorrect test tensor shape.\\n\"\n            f\"Expected: {expected_shape}\\n\"\n            f\"Actual  : {image_shape}\\n\"\n            f\"Study   : {study_id}\"\n        )\n\nprint()\nprint(\"✓ ALL TEST STUDIES HAVE CORRECT INPUT SHAPE\")\n\n\n# -----------------------------------------------------------------------------\n# DATALOADER\n# -----------------------------------------------------------------------------\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True,\n)\n\nprint()\nprint(\"=\" * 80)\nprint(\"DATALOADER\")\nprint(\"=\" * 80)\n\nprint(\"Batches:\", len(test_loader))\n\n\n# -----------------------------------------------------------------------------\n# LOAD ORIGINAL MODEL\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"LOADING ORIGINAL MODEL\")\nprint(\"=\" * 80)\n\nmodel = RSNAKneeMultiPlaneModel(\n    num_classes=len(LABEL_COLUMNS)\n)\n\nmodel = model.to(DEVICE)\n\n# PyTorch >= 2.6 compatibility\ntry:\n\n    checkpoint = torch.load(\n        CHECKPOINT,\n        map_location=DEVICE,\n        weights_only=False,\n    )\n\nexcept TypeError:\n\n    checkpoint = torch.load(\n        CHECKPOINT,\n        map_location=DEVICE,\n    )\n\n\n# -----------------------------------------------------------------------------\n# CHECKPOINT FORMAT\n# -----------------------------------------------------------------------------\n\nif isinstance(checkpoint, dict):\n\n    if \"model_state_dict\" in checkpoint:\n\n        state_dict = checkpoint[\"model_state_dict\"]\n\n    elif \"state_dict\" in checkpoint:\n\n        state_dict = checkpoint[\"state_dict\"]\n\n    else:\n\n        # Assume the dictionary itself is the state dict\n        state_dict = checkpoint\n\nelse:\n\n    state_dict = checkpoint\n\n\nmodel.load_state_dict(\n    state_dict,\n    strict=True,\n)\n\nmodel.eval()\n\nprint(\"✓ Original checkpoint loaded\")\nprint(\"✓ Model moved to:\", DEVICE)\nprint(\"✓ Model set to evaluation mode\")\n\n\n# -----------------------------------------------------------------------------\n# TEST INFERENCE\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"RUNNING TEST INFERENCE\")\nprint(\"=\" * 80)\n\nall_probabilities = []\nall_study_ids = []\n\nwith torch.no_grad():\n\n    for batch_idx, batch in enumerate(test_loader):\n\n        images = batch[0]\n        plane_mask = batch[1]\n\n        images = images.to(\n            DEVICE,\n            non_blocking=True,\n        )\n\n        plane_mask = plane_mask.to(\n            DEVICE,\n            non_blocking=True,\n        )\n\n        outputs = model(\n            images,\n            plane_mask,\n        )\n\n        probabilities = torch.sigmoid(\n            outputs\n        )\n\n        probabilities = (\n            probabilities\n            .detach()\n            .cpu()\n            .numpy()\n        )\n\n        all_probabilities.append(\n            probabilities\n        )\n\n        start = batch_idx * BATCH_SIZE\n\n        batch_size_actual = (\n            probabilities.shape[0]\n        )\n\n        batch_ids = test_df[\n            \"StudyInstanceUID\"\n        ].iloc[\n            start:start + batch_size_actual\n        ].tolist()\n\n        all_study_ids.extend(\n            batch_ids\n        )\n\n        print(\n            f\"Batch {batch_idx + 1}/{len(test_loader)} \"\n            f\"→ {probabilities.shape}\"\n        )\n\n\n# -----------------------------------------------------------------------------\n# COMBINE PREDICTIONS\n# -----------------------------------------------------------------------------\n\nprobabilities = np.concatenate(\n    all_probabilities,\n    axis=0,\n)\n\nprint()\nprint(\"=\" * 80)\nprint(\"TEST PREDICTION RESULTS\")\nprint(\"=\" * 80)\n\nprint(\n    \"Probability matrix shape:\",\n    probabilities.shape,\n)\n\nexpected_prediction_shape = (\n    len(test_df),\n    len(LABEL_COLUMNS),\n)\n\nprint(\n    \"Expected shape:\",\n    expected_prediction_shape,\n)\n\nif probabilities.shape != expected_prediction_shape:\n\n    raise ValueError(\n        f\"Prediction shape mismatch.\\n\"\n        f\"Expected: {expected_prediction_shape}\\n\"\n        f\"Actual  : {probabilities.shape}\"\n    )\n\nprint(\"✓ Prediction shape correct\")\n\n\n# -----------------------------------------------------------------------------\n# RANGE CHECK\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"Probability minimum:\", probabilities.min())\nprint(\"Probability maximum:\", probabilities.max())\n\nif (\n    np.any(probabilities < 0)\n    or np.any(probabilities > 1)\n):\n\n    raise ValueError(\n        \"Invalid sigmoid probabilities.\"\n    )\n\nprint(\"✓ All probabilities are in [0, 1]\")\n\n\n# -----------------------------------------------------------------------------\n# CREATE PROBABILITY DATAFRAME\n# -----------------------------------------------------------------------------\n\nprediction_df = pd.DataFrame(\n    probabilities,\n    columns=LABEL_COLUMNS,\n)\n\nprediction_df.insert(\n    0,\n    \"StudyInstanceUID\",\n    all_study_ids,\n)\n\nprint()\nprint(\"=\" * 80)\nprint(\"TEST PREDICTIONS\")\nprint(\"=\" * 80)\n\npd.set_option(\n    \"display.max_columns\",\n    None,\n)\n\nprint(\n    prediction_df.to_string(\n        index=False\n    )\n)\n\n\n# -----------------------------------------------------------------------------\n# SAVE RAW PROBABILITIES\n# -----------------------------------------------------------------------------\n\nRAW_OUTPUT = (\n    \"/kaggle/working/\"\n    \"rsna_knee_test_probabilities.csv\"\n)\n\nprediction_df.to_csv(\n    RAW_OUTPUT,\n    index=False,\n)\n\nprint()\nprint(\"=\" * 80)\nprint(\"RAW TEST PROBABILITIES SAVED\")\nprint(\"=\" * 80)\n\nprint(RAW_OUTPUT)\n\n\n# -----------------------------------------------------------------------------\n# FINAL CHECK\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"STEP 78D COMPLETE\")\nprint(\"=\" * 80)\n\nprint(\"✓ Test CSV loaded\")\nprint(\"✓ Test manifest loaded\")\nprint(\"✓ All test studies found\")\nprint(\"✓ Test dataset created\")\nprint(\"✓ Test preprocessing verified\")\nprint(\"✓ Input tensor shape verified\")\nprint(\"✓ Original checkpoint loaded\")\nprint(\"✓ Test inference completed\")\nprint(\"✓ Sigmoid probabilities generated\")\nprint(\"✓ Continuous probabilities preserved\")\nprint(\"✓ No class thresholds applied\")\nprint(\"✓ No 0/1 conversion performed\")\n\nprint()\nprint(\"Output:\")\nprint(RAW_OUTPUT)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T07:10:51.81711Z","iopub.execute_input":"2026-08-19T07:10:51.817899Z","iopub.status.idle":"2026-08-19T07:10:52.345204Z","shell.execute_reply.started":"2026-08-19T07:10:51.817866Z","shell.execute_reply":"2026-08-19T07:10:52.344316Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# STEP 78D FIX — TEST DATASET OUTPUT CHECK\n# =============================================================================\n\nprint()\nprint(\"=\" * 80)\nprint(\"TESTING PREPROCESSING\")\nprint(\"=\" * 80)\n\nexpected_shape = (\n    3,\n    TARGET_SLICES,\n    IMAGE_SIZE,\n    IMAGE_SIZE,\n)\n\nfor i in range(len(test_dataset)):\n\n    study_id = test_df.iloc[i][\"StudyInstanceUID\"]\n\n    print()\n    print(f\"TEST STUDY {i + 1}\")\n    print(\"-\" * 70)\n    print(\"Study ID:\", study_id)\n\n    # ------------------------------------------------------------\n    # Dataset returns a dictionary\n    # ------------------------------------------------------------\n\n    sample = test_dataset[i]\n\n    print(\"Dataset output type:\", type(sample))\n\n    if not isinstance(sample, dict):\n        raise TypeError(\n            \"Unexpected dataset output.\\n\"\n            f\"Expected dict, got {type(sample)}\"\n        )\n\n    print(\"Dataset keys:\", list(sample.keys()))\n\n    # ------------------------------------------------------------\n    # Find image tensor\n    # ------------------------------------------------------------\n\n    possible_image_keys = [\n        \"image\",\n        \"images\",\n        \"tensor\",\n        \"image_tensor\",\n    ]\n\n    image_key = None\n\n    for key in possible_image_keys:\n        if key in sample:\n            image_key = key\n            break\n\n    if image_key is None:\n\n        raise KeyError(\n            \"Could not find image tensor in dataset output.\\n\"\n            f\"Available keys: {list(sample.keys())}\"\n        )\n\n    images = sample[image_key]\n\n    # ------------------------------------------------------------\n    # Find plane mask\n    # ------------------------------------------------------------\n\n    possible_mask_keys = [\n        \"plane_mask\",\n        \"plane_masks\",\n        \"mask\",\n    ]\n\n    mask_key = None\n\n    for key in possible_mask_keys:\n        if key in sample:\n            mask_key = key\n            break\n\n    if mask_key is None:\n\n        raise KeyError(\n            \"Could not find plane mask in dataset output.\\n\"\n            f\"Available keys: {list(sample.keys())}\"\n        )\n\n    plane_mask = sample[mask_key]\n\n    # ------------------------------------------------------------\n    # Convert/check image tensor\n    # ------------------------------------------------------------\n\n    if not torch.is_tensor(images):\n\n        images = torch.as_tensor(\n            images,\n            dtype=torch.float32\n        )\n\n    image_shape = tuple(images.shape)\n\n    print(\"Image key   :\", image_key)\n    print(\"Image shape :\", image_shape)\n    print(\"Image dtype :\", images.dtype)\n    print(\"Image min   :\", float(images.min()))\n    print(\"Image max   :\", float(images.max()))\n\n    # ------------------------------------------------------------\n    # Plane mask\n    # ------------------------------------------------------------\n\n    if not torch.is_tensor(plane_mask):\n\n        plane_mask = torch.as_tensor(\n            plane_mask,\n            dtype=torch.float32\n        )\n\n    print(\"Mask key    :\", mask_key)\n    print(\"Plane mask  :\", plane_mask.cpu().numpy())\n\n    # ------------------------------------------------------------\n    # Expected tensor shape\n    # ------------------------------------------------------------\n\n    if image_shape != expected_shape:\n\n        raise ValueError(\n            \"\\nINCORRECT TEST TENSOR SHAPE\\n\"\n            f\"Expected : {expected_shape}\\n\"\n            f\"Actual   : {image_shape}\\n\"\n            f\"Study    : {study_id}\"\n        )\n\n    # ------------------------------------------------------------\n    # Mask validation\n    # ------------------------------------------------------------\n\n    if tuple(plane_mask.shape) != (3,):\n\n        raise ValueError(\n            \"\\nINCORRECT PLANE MASK SHAPE\\n\"\n            f\"Expected : (3,)\\n\"\n            f\"Actual   : {tuple(plane_mask.shape)}\\n\"\n            f\"Study    : {study_id}\"\n        )\n\n    print(\"✓ Tensor shape correct\")\n    print(\"✓ Plane mask shape correct\")\n\nprint()\nprint(\"=\" * 80)\nprint(\"PREPROCESSING CHECK COMPLETE\")\nprint(\"=\" * 80)\n\nprint(\"Expected tensor:\")\nprint(expected_shape)\n\nprint()\nprint(\"✓ All test studies successfully loaded\")\nprint(\"✓ Missing-plane handling preserved\")\nprint(\"✓ Plane masks verified\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T07:12:11.890047Z","iopub.execute_input":"2026-08-19T07:12:11.890797Z","iopub.status.idle":"2026-08-19T07:12:12.935781Z","shell.execute_reply.started":"2026-08-19T07:12:11.890764Z","shell.execute_reply":"2026-08-19T07:12:12.934833Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# STEP 78E — FINAL TEST INFERENCE\n# =============================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import DataLoader\n\nprint(\"=\" * 80)\nprint(\"STEP 78E — FINAL TEST INFERENCE\")\nprint(\"=\" * 80)\n\n# -----------------------------------------------------------------------------\n# CONFIGURATION\n# -----------------------------------------------------------------------------\n\nCHECKPOINT = \"/kaggle/working/best_rsna_knee_model.pth\"\nMANIFEST_PATH = \"/kaggle/working/rsna_knee_test_series_manifest.csv\"\n\nBATCH_SIZE = 2\nNUM_WORKERS = 2\n\nprint()\nprint(\"=\" * 80)\nprint(\"CONFIGURATION\")\nprint(\"=\" * 80)\n\nprint(\"Checkpoint :\", CHECKPOINT)\nprint(\"Manifest   :\", MANIFEST_PATH)\nprint(\"Device     :\", DEVICE)\nprint(\"Batch size :\", BATCH_SIZE)\nprint(\"Labels     :\", len(LABEL_COLUMNS))\n\n# -----------------------------------------------------------------------------\n# FILE CHECK\n# -----------------------------------------------------------------------------\n\nif not os.path.exists(CHECKPOINT):\n    raise FileNotFoundError(\n        f\"Checkpoint not found:\\n{CHECKPOINT}\"\n    )\n\nif not os.path.exists(MANIFEST_PATH):\n    raise FileNotFoundError(\n        f\"Test manifest not found:\\n{MANIFEST_PATH}\"\n    )\n\nprint()\nprint(\"✓ Checkpoint found\")\nprint(\"✓ Test manifest found\")\n\n# -----------------------------------------------------------------------------\n# LOAD TEST MANIFEST\n# -----------------------------------------------------------------------------\n\ntest_manifest_final = pd.read_csv(MANIFEST_PATH)\n\nprint()\nprint(\"=\" * 80)\nprint(\"TEST MANIFEST\")\nprint(\"=\" * 80)\n\nprint(\"Shape :\", test_manifest_final.shape)\nprint(\"Studies:\", test_manifest_final[\"StudyInstanceUID\"].nunique())\nprint(\"Series :\", len(test_manifest_final))\n\n# IMPORTANT:\n# Use the FULL manifest.\n# Do NOT filter out Unknown here because the original dataset\n# contains the series-selection logic.\n\n# -----------------------------------------------------------------------------\n# CREATE FINAL TEST DATAFRAME\n# -----------------------------------------------------------------------------\n\ntest_final_df = test_df.copy()\n\nfor label in LABEL_COLUMNS:\n    if label not in test_final_df.columns:\n        test_final_df[label] = 0.0\n\ntest_final_df = test_final_df[\n    [\"StudyInstanceUID\"] + LABEL_COLUMNS\n].copy()\n\nprint()\nprint(\"=\" * 80)\nprint(\"FINAL TEST DATAFRAME\")\nprint(\"=\" * 80)\n\nprint(\"Shape:\", test_final_df.shape)\n\n# -----------------------------------------------------------------------------\n# CREATE DATASET\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"CREATING FINAL TEST DATASET\")\nprint(\"=\" * 80)\n\ntest_dataset_final = RSNAKneeStudyDataset(\n    test_final_df,\n    test_manifest_final,\n    LABEL_COLUMNS,\n    target_slices=TARGET_SLICES,\n    image_size=IMAGE_SIZE\n)\n\nprint(\"✓ Final test dataset created\")\nprint(\"Number of studies:\", len(test_dataset_final))\n\n# -----------------------------------------------------------------------------\n# DATALOADER\n# -----------------------------------------------------------------------------\n\ntest_loader_final = DataLoader(\n    test_dataset_final,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=True\n)\n\nprint(\"✓ Final test DataLoader created\")\n\n# -----------------------------------------------------------------------------\n# LOAD ORIGINAL MODEL\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"LOADING SELECTED ORIGINAL MODEL\")\nprint(\"=\" * 80)\n\nmodel.load_state_dict(\n    torch.load(\n        CHECKPOINT,\n        map_location=DEVICE,\n        weights_only=False\n    )[\"model_state_dict\"]\n)\n\nmodel = model.to(DEVICE)\nmodel.eval()\n\nprint(\"✓ Original checkpoint loaded\")\nprint(\"✓ Model moved to:\", DEVICE)\nprint(\"✓ Evaluation mode enabled\")\n\n# -----------------------------------------------------------------------------\n# FINAL INFERENCE\n# -----------------------------------------------------------------------------\n\nall_probabilities = []\nall_study_ids = []\n\nprint()\nprint(\"=\" * 80)\nprint(\"RUNNING TEST INFERENCE\")\nprint(\"=\" * 80)\n\nwith torch.no_grad():\n\n    for batch_idx, batch in enumerate(test_loader_final):\n\n        images = batch[\"images\"].to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        plane_mask = batch[\"plane_mask\"].to(\n            DEVICE,\n            non_blocking=True\n        )\n\n        study_ids = batch[\"study_id\"]\n\n        # ---------------------------------------------------------\n        # Forward pass\n        # ---------------------------------------------------------\n\n        logits = model(\n            images,\n            plane_mask\n        )\n\n        # ---------------------------------------------------------\n        # Continuous probabilities\n        # ---------------------------------------------------------\n\n        probabilities = torch.sigmoid(logits)\n\n        all_probabilities.append(\n            probabilities.detach().cpu().numpy()\n        )\n\n        if isinstance(study_ids, (list, tuple)):\n            all_study_ids.extend(\n                list(study_ids)\n            )\n        else:\n            all_study_ids.extend(\n                study_ids\n            )\n\n        print(\n            f\"Batch {batch_idx + 1}/\"\n            f\"{len(test_loader_final)}\"\n            f\" | Studies: {len(study_ids)}\"\n            f\" | Output: {tuple(probabilities.shape)}\"\n        )\n\n# -----------------------------------------------------------------------------\n# COMBINE RESULTS\n# -----------------------------------------------------------------------------\n\nprobabilities = np.concatenate(\n    all_probabilities,\n    axis=0\n)\n\nprint()\nprint(\"=\" * 80)\nprint(\"INFERENCE RESULTS\")\nprint(\"=\" * 80)\n\nprint(\"Probability shape:\", probabilities.shape)\nprint(\"Study IDs:\", len(all_study_ids))\n\n# -----------------------------------------------------------------------------\n# VALIDATION\n# -----------------------------------------------------------------------------\n\nif probabilities.shape != (\n    len(test_final_df),\n    len(LABEL_COLUMNS)\n):\n\n    raise ValueError(\n        \"\\nUnexpected probability shape!\\n\"\n        f\"Expected: \"\n        f\"({len(test_final_df)}, {len(LABEL_COLUMNS)})\\n\"\n        f\"Actual: {probabilities.shape}\"\n    )\n\nif len(all_study_ids) != len(test_final_df):\n\n    raise ValueError(\n        \"\\nStudy ID count mismatch!\\n\"\n        f\"Expected: {len(test_final_df)}\\n\"\n        f\"Actual: {len(all_study_ids)}\"\n    )\n\nprint(\"✓ Probability shape correct\")\nprint(\"✓ Study count correct\")\n\n# -----------------------------------------------------------------------------\n# PROBABILITY RANGE\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"PROBABILITY RANGE\")\nprint(\"=\" * 80)\n\nprint(\"Minimum :\", probabilities.min())\nprint(\"Maximum :\", probabilities.max())\nprint(\"Mean    :\", probabilities.mean())\n\nif probabilities.min() < 0 or probabilities.max() > 1:\n    raise ValueError(\n        \"Probabilities outside [0,1] range!\"\n    )\n\nprint(\"✓ All probabilities are within [0,1]\")\n\n# -----------------------------------------------------------------------------\n# CREATE RESULT DATAFRAME\n# -----------------------------------------------------------------------------\n\nresults = pd.DataFrame(\n    probabilities,\n    columns=LABEL_COLUMNS\n)\n\nresults.insert(\n    0,\n    \"StudyInstanceUID\",\n    all_study_ids\n)\n\n# -----------------------------------------------------------------------------\n# DISPLAY RESULTS\n# -----------------------------------------------------------------------------\n\npd.set_option(\n    \"display.max_columns\",\n    None\n)\n\npd.set_option(\n    \"display.width\",\n    200\n)\n\nprint()\nprint(\"=\" * 80)\nprint(\"FINAL TEST PROBABILITIES\")\nprint(\"=\" * 80)\n\nprint(results.to_string(index=False))\n\n# -----------------------------------------------------------------------------\n# SAVE RAW PROBABILITIES\n# -----------------------------------------------------------------------------\n\nOUTPUT_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_test_probabilities.csv\"\n)\n\nresults.to_csv(\n    OUTPUT_PATH,\n    index=False\n)\n\nprint()\nprint(\"=\" * 80)\nprint(\"STEP 78E COMPLETE\")\nprint(\"=\" * 80)\n\nprint(\"Saved:\")\nprint(OUTPUT_PATH)\n\nprint()\nprint(\"✓ Original checkpoint used\")\nprint(\"✓ Full test manifest used\")\nprint(\"✓ Missing planes handled by plane mask\")\nprint(\"✓ Test images converted to model tensor\")\nprint(\"✓ Forward inference completed\")\nprint(\"✓ Sigmoid probabilities generated\")\nprint(\"✓ Continuous probabilities preserved\")\nprint(\"✓ No class thresholds applied\")\nprint(\"✓ No 0/1 conversion performed\")\nprint(\"✓ Test probability CSV saved\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T07:13:33.280884Z","iopub.execute_input":"2026-08-19T07:13:33.281368Z","iopub.status.idle":"2026-08-19T07:13:34.385986Z","shell.execute_reply.started":"2026-08-19T07:13:33.281336Z","shell.execute_reply":"2026-08-19T07:13:34.385029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# STEP 78F — VERIFY KAGGLE SUBMISSION FORMAT\n# =============================================================================\n\nimport os\nimport pandas as pd\nimport numpy as np\n\nprint(\"=\" * 80)\nprint(\"STEP 78F — KAGGLE SUBMISSION FORMAT CHECK\")\nprint(\"=\" * 80)\n\nSAMPLE_PATH = (\n    \"/kaggle/input/competitions/\"\n    \"rsna-knee-abnormality-detection/\"\n    \"sample_submission.csv\"\n)\n\nPROB_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_test_probabilities.csv\"\n)\n\n# -----------------------------------------------------------------------------\n# FILE CHECK\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"FILE CHECK\")\nprint(\"=\" * 80)\n\nprint(\"Sample submission:\", SAMPLE_PATH)\nprint(\"Probability file  :\", PROB_PATH)\n\nif not os.path.exists(SAMPLE_PATH):\n    raise FileNotFoundError(\n        f\"sample_submission.csv not found:\\n{SAMPLE_PATH}\"\n    )\n\nif not os.path.exists(PROB_PATH):\n    raise FileNotFoundError(\n        f\"Test probability file not found:\\n{PROB_PATH}\"\n    )\n\nprint(\"✓ sample_submission.csv found\")\nprint(\"✓ Test probabilities found\")\n\n# -----------------------------------------------------------------------------\n# LOAD FILES\n# -----------------------------------------------------------------------------\n\nsample_df = pd.read_csv(SAMPLE_PATH)\nprob_df = pd.read_csv(PROB_PATH)\n\nprint()\nprint(\"=\" * 80)\nprint(\"SAMPLE SUBMISSION\")\nprint(\"=\" * 80)\n\nprint(\"Shape:\", sample_df.shape)\nprint(\"Columns:\")\nfor i, col in enumerate(sample_df.columns):\n    print(f\"{i:2d}: {col}\")\n\nprint()\nprint(sample_df.head())\n\nprint()\nprint(\"=\" * 80)\nprint(\"MODEL PROBABILITY FILE\")\nprint(\"=\" * 80)\n\nprint(\"Shape:\", prob_df.shape)\n\nprint(\"Columns:\")\nfor i, col in enumerate(prob_df.columns):\n    print(f\"{i:2d}: {col}\")\n\nprint()\nprint(prob_df.head())\n\n# -----------------------------------------------------------------------------\n# COLUMN COMPARISON\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"COLUMN COMPARISON\")\nprint(\"=\" * 80)\n\nsample_columns = list(sample_df.columns)\nprob_columns = list(prob_df.columns)\n\nprint(\"Sample columns:\", len(sample_columns))\nprint(\"Probability columns:\", len(prob_columns))\n\nprint()\n\nmissing_from_prob = [\n    c for c in sample_columns\n    if c not in prob_columns\n]\n\nextra_in_prob = [\n    c for c in prob_columns\n    if c not in sample_columns\n]\n\nif missing_from_prob:\n\n    print(\"Missing from probability file:\")\n    for c in missing_from_prob:\n        print(\"  -\", c)\n\nelse:\n\n    print(\"✓ No sample columns missing\")\n\nif extra_in_prob:\n\n    print()\n    print(\"Extra probability columns:\")\n    for c in extra_in_prob:\n        print(\"  -\", c)\n\nelse:\n\n    print(\"✓ No extra probability columns\")\n\n# -----------------------------------------------------------------------------\n# STUDY ID CHECK\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"STUDY ID CHECK\")\nprint(\"=\" * 80)\n\nif \"StudyInstanceUID\" in sample_df.columns:\n\n    sample_ids = sample_df[\"StudyInstanceUID\"].astype(str)\n    prob_ids = prob_df[\"StudyInstanceUID\"].astype(str)\n\n    print(\"Sample studies :\", len(sample_ids))\n    print(\"Probability studies:\", len(prob_ids))\n\n    print()\n    print(\"Sample IDs:\")\n    for x in sample_ids:\n        print(x)\n\n    print()\n    print(\"Probability IDs:\")\n    for x in prob_ids:\n        print(x)\n\n    missing_ids = [\n        x for x in sample_ids\n        if x not in set(prob_ids)\n    ]\n\n    extra_ids = [\n        x for x in prob_ids\n        if x not in set(sample_ids)\n    ]\n\n    print()\n\n    if missing_ids:\n        print(\"❌ Missing study IDs:\")\n        for x in missing_ids:\n            print(x)\n    else:\n        print(\"✓ All sample study IDs present\")\n\n    if extra_ids:\n        print(\"❌ Extra study IDs:\")\n        for x in extra_ids:\n            print(x)\n    else:\n        print(\"✓ No extra study IDs\")\n\n# -----------------------------------------------------------------------------\n# LABEL COLUMN CHECK\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"12 LABEL CHECK\")\nprint(\"=\" * 80)\n\nprint(\"Expected labels:\")\n\nfor i, label in enumerate(LABEL_COLUMNS):\n    print(f\"{i:2d}: {label}\")\n\nmissing_labels = [\n    label\n    for label in LABEL_COLUMNS\n    if label not in prob_df.columns\n]\n\nif missing_labels:\n\n    print()\n    print(\"❌ Missing label columns:\")\n\n    for label in missing_labels:\n        print(\"  -\", label)\n\n    raise ValueError(\n        \"Probability file does not contain all 12 label columns.\"\n    )\n\nprint()\nprint(\"✓ All 12 model labels found\")\n\n# -----------------------------------------------------------------------------\n# CHECK PROBABILITY VALUES\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"PROBABILITY VALIDATION\")\nprint(\"=\" * 80)\n\nlabel_values = prob_df[LABEL_COLUMNS].to_numpy(\n    dtype=np.float32\n)\n\nprint(\"Minimum:\", label_values.min())\nprint(\"Maximum:\", label_values.max())\nprint(\"Mean   :\", label_values.mean())\n\nif np.isnan(label_values).any():\n\n    raise ValueError(\n        \"NaN values found in prediction probabilities.\"\n    )\n\nif not np.isfinite(label_values).all():\n\n    raise ValueError(\n        \"Non-finite values found in prediction probabilities.\"\n    )\n\nif (\n    label_values.min() < 0\n    or label_values.max() > 1\n):\n\n    raise ValueError(\n        \"Probability values outside [0,1].\"\n    )\n\nprint(\"✓ No NaN values\")\nprint(\"✓ No infinite values\")\nprint(\"✓ All values within [0,1]\")\n\n# -----------------------------------------------------------------------------\n# CHECK COLUMN ORDER\n# -----------------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 80)\nprint(\"COLUMN ORDER\")\nprint(\"=\" * 80)\n\nif sample_columns == prob_columns:\n\n    print(\"✓ Column order already matches sample submission\")\n\nelse:\n\n    print(\"⚠ Column order differs\")\n\n    print()\n    print(\"Sample order:\")\n    print(sample_columns)\n\n    print()\n    print(\"Probability order:\")\n    print(prob_columns)\n\nprint()\nprint(\"=\" * 80)\nprint(\"STEP 78F COMPLETE\")\nprint(\"=\" * 80)\n\nprint()\nprint(\"DO NOT CREATE THE FINAL SUBMISSION YET.\")\nprint(\"The output above will tell us the exact required format.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T07:15:37.545047Z","iopub.execute_input":"2026-08-19T07:15:37.545854Z","iopub.status.idle":"2026-08-19T07:15:37.58597Z","shell.execute_reply.started":"2026-08-19T07:15:37.545822Z","shell.execute_reply":"2026-08-19T07:15:37.584954Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# STEP 78G — FINAL KAGGLE SUBMISSION GENERATION\n# =============================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\n\nprint(\"=\" * 80)\nprint(\"STEP 78G — FINAL KAGGLE SUBMISSION\")\nprint(\"=\" * 80)\n\n# -----------------------------------------------------------------------------\n# PATHS\n# -----------------------------------------------------------------------------\n\nSAMPLE_PATH = (\n    \"/kaggle/input/competitions/\"\n    \"rsna-knee-abnormality-detection/\"\n    \"sample_submission.csv\"\n)\n\nPROB_PATH = (\n    \"/kaggle/working/\"\n    \"rsna_knee_test_probabilities.csv\"\n)\n\nSUBMISSION_PATH = (\n    \"/kaggle/working/\"\n    \"submission.csv\"\n)\n\n# -----------------------------------------------------------------------------\n# LOAD FILES\n# -----------------------------------------------------------------------------\n\nsample_df = pd.read_csv(SAMPLE_PATH)\nprob_df = pd.read_csv(PROB_PATH)\n\nprint()\nprint(\"=\" * 80)\nprint(\"INPUT FILES\")\nprint(\"=\" * 80)\n\nprint(\"Sample shape       :\", sample_df.shape)\nprint(\"Probability shape  :\", prob_df.shape)\n\n# -----------------------------------------------------------------------------\n# CHECK SAMPLE COLUMNS\n# -----------------------------------------------------------------------------\n\nsample_columns = list(sample_df.columns)\nprob_columns = list(prob_df.columns)\n\nif sample_columns != prob_columns:\n\n    raise ValueError(\n        \"\\nColumn mismatch!\\n\\n\"\n        f\"Sample columns:\\n{sample_columns}\\n\\n\"\n        f\"Probability columns:\\n{prob_columns}\"\n    )\n\nprint(\"✓ Column names match\")\nprint(\"✓ Column order matches\")\n\n# -----------------------------------------------------------------------------\n# CHECK STUDY IDS\n# -----------------------------------------------------------------------------\n\nsample_ids = sample_df[\"StudyInstanceUID\"].astype(str)\nprob_ids = prob_df[\"StudyInstanceUID\"].astype(str)\n\nif len(sample_ids) != len(prob_ids):\n\n    raise ValueError(\n        \"Number of studies does not match.\"\n    )\n\nif set(sample_ids) != set(prob_ids):\n\n    raise ValueError(\n        \"StudyInstanceUID sets do not match.\"\n    )\n\nprint(\"✓ Study IDs match\")\n\n# -----------------------------------------------------------------------------\n# REORDER MODEL OUTPUT TO SAMPLE ORDER\n# -----------------------------------------------------------------------------\n\nprob_df[\"StudyInstanceUID\"] = (\n    prob_df[\"StudyInstanceUID\"].astype(str)\n)\n\nsubmission_df = (\n    sample_df[[\"StudyInstanceUID\"]]\n    .copy()\n)\n\nsubmission_df[\"StudyInstanceUID\"] = (\n    submission_df[\"StudyInstanceUID\"].astype(str)\n)\n\n# Map predictions by StudyInstanceUID\nprob_lookup = prob_df.set_index(\n    \"StudyInstanceUID\"\n)\n\nfor label in LABEL_COLUMNS:\n\n    submission_df[label] = submission_df[\n        \"StudyInstanceUID\"\n    ].map(\n        prob_lookup[label]\n    )\n\n# -----------------------------------------------------------------------------\n# FINAL NUMERICAL VALIDATION\n# -----------------------------------------------------------------------------\n\nprediction_values = submission_df[\n    LABEL_COLUMNS\n].to_numpy(\n    dtype=np.float64\n)\n\nif np.isnan(prediction_values).any():\n\n    raise ValueError(\n        \"NaN values detected in final submission.\"\n    )\n\nif not np.isfinite(prediction_values).all():\n\n    raise ValueError(\n        \"Infinite values detected in final submission.\"\n    )\n\nif (\n    prediction_values.min() < 0\n    or prediction_values.max() > 1\n):\n\n    raise ValueError(\n        \"Prediction values outside [0,1].\"\n    )\n\nprint()\nprint(\"=\" * 80)\nprint(\"FINAL PREDICTION VALIDATION\")\nprint(\"=\" * 80)\n\nprint(\"Minimum probability :\", prediction_values.min())\nprint(\"Maximum probability :\", prediction_values.max())\nprint(\"Mean probability    :\", prediction_values.mean())\n\nprint(\"✓ No NaN\")\nprint(\"✓ No Inf\")\nprint(\"✓ All probabilities in [0,1]\")\n\n# -----------------------------------------------------------------------------\n# EXACT SAMPLE FORMAT CHECK\n# -----------------------------------------------------------------------------\n\nif list(submission_df.columns) != sample_columns:\n\n    raise ValueError(\n        \"Final submission columns do not exactly match sample submission.\"\n    )\n\nif submission_df.shape != sample_df.shape:\n\n    raise ValueError(\n        \"Final submission shape does not match sample submission.\"\n    )\n\nprint()\nprint(\"=\" * 80)\nprint(\"FORMAT VALIDATION\")\nprint(\"=\" * 80)\n\nprint(\"Sample shape     :\", sample_df.shape)\nprint(\"Submission shape :\", submission_df.shape)\n\nprint(\"✓ Shape matches\")\nprint(\"✓ Columns match\")\nprint(\"✓ Column order matches\")\nprint(\"✓ Study order matches\")\n\n# -----------------------------------------------------------------------------\n# SAVE FINAL SUBMISSION\n# -----------------------------------------------------------------------------\n\nsubmission_df.to_csv(\n    SUBMISSION_PATH,\n    index=False\n)\n\nprint()\nprint(\"=\" * 80)\nprint(\"FINAL SUBMISSION CREATED\")\nprint(\"=\" * 80)\n\nprint(\"Path:\")\nprint(SUBMISSION_PATH)\n\nprint()\nprint(\"FINAL SUBMISSION:\")\nprint(submission_df.to_string(index=False))\n\n# -----------------------------------------------------------------------------\n# READ BACK FILE\n# -----------------------------------------------------------------------------\n\ncheck_df = pd.read_csv(\n    SUBMISSION_PATH\n)\n\nprint()\nprint(\"=\" * 80)\nprint(\"FINAL FILE READ-BACK CHECK\")\nprint(\"=\" * 80)\n\nprint(\"Shape:\", check_df.shape)\nprint(\"Columns:\", list(check_df.columns))\n\nif check_df.equals(submission_df):\n\n    print(\"✓ Saved file matches generated submission\")\n\nelse:\n\n    raise ValueError(\n        \"Saved submission differs from generated submission.\"\n    )\n\nprint()\nprint(\"=\" * 80)\nprint(\"STEP 78G COMPLETE\")\nprint(\"=\" * 80)\n\nprint()\nprint(\"✓ FINAL submission.csv created\")\nprint(\"✓ Original selected checkpoint used\")\nprint(\"✓ Continuous sigmoid probabilities preserved\")\nprint(\"✓ No class thresholds applied\")\nprint(\"✓ No 0/1 conversion\")\nprint(\"✓ Sample-submission format preserved\")\nprint(\"✓ Study order preserved\")\nprint(\"✓ All 12 abnormality columns preserved\")\n\nprint()\nprint(\"SUBMISSION FILE:\")\nprint(SUBMISSION_PATH)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T07:27:41.103581Z","iopub.execute_input":"2026-08-19T07:27:41.104006Z","iopub.status.idle":"2026-08-19T07:27:41.145286Z","shell.execute_reply.started":"2026-08-19T07:27:41.103977Z","shell.execute_reply":"2026-08-19T07:27:41.144617Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport os\n\nPROB_FILE = \"/kaggle/working/rsna_knee_test_probabilities.csv\"\nSAMPLE_FILE = \"/kaggle/input/competitions/rsna-knee-abnormality-detection/sample_submission.csv\"\nOUTPUT_FILE = \"/kaggle/working/submission.csv\"\n\nprint(\"=\" * 80)\nprint(\"RECREATING FINAL SUBMISSION\")\nprint(\"=\" * 80)\n\n# ------------------------------------------------------------\n# Load model probabilities\n# ------------------------------------------------------------\nprob_df = pd.read_csv(PROB_FILE)\n\n# ------------------------------------------------------------\n# Load Kaggle sample submission\n# ------------------------------------------------------------\nsample_df = pd.read_csv(SAMPLE_FILE)\n\nprint(\"Probability shape:\", prob_df.shape)\nprint(\"Sample shape     :\", sample_df.shape)\n\n# ------------------------------------------------------------\n# Verify columns\n# ------------------------------------------------------------\nif list(prob_df.columns) != list(sample_df.columns):\n    raise ValueError(\n        \"Column mismatch between probability file and sample submission.\"\n    )\n\nprint(\"✓ Columns match\")\nprint(\"✓ Column order matches\")\n\n# ------------------------------------------------------------\n# Verify StudyInstanceUID\n# ------------------------------------------------------------\nif set(prob_df[\"StudyInstanceUID\"]) != set(sample_df[\"StudyInstanceUID\"]):\n    raise ValueError(\n        \"StudyInstanceUID mismatch between probability file and sample submission.\"\n    )\n\nprint(\"✓ Study IDs match\")\n\n# ------------------------------------------------------------\n# IMPORTANT:\n# Reorder predictions exactly like sample_submission.csv\n# ------------------------------------------------------------\nsubmission = sample_df[[\"StudyInstanceUID\"]].merge(\n    prob_df,\n    on=\"StudyInstanceUID\",\n    how=\"left\",\n    validate=\"one_to_one\"\n)\n\n# ------------------------------------------------------------\n# Restore exact sample column order\n# ------------------------------------------------------------\nsubmission = submission[sample_df.columns]\n\n# ------------------------------------------------------------\n# Validation\n# ------------------------------------------------------------\nprediction_columns = sample_df.columns[1:]\n\nif submission[prediction_columns].isna().any().any():\n    raise ValueError(\"NaN values detected in submission.\")\n\nif not submission[prediction_columns].apply(\n    lambda col: col.between(0, 1).all()\n).all():\n    raise ValueError(\"Probability outside [0,1] detected.\")\n\nif submission.shape != sample_df.shape:\n    raise ValueError(\"Submission shape does not match sample submission.\")\n\n# ------------------------------------------------------------\n# Save\n# ------------------------------------------------------------\nsubmission.to_csv(\n    OUTPUT_FILE,\n    index=False\n)\n\nprint()\nprint(\"=\" * 80)\nprint(\"FINAL SUBMISSION CREATED\")\nprint(\"=\" * 80)\n\nprint(\"Path:\")\nprint(OUTPUT_FILE)\n\nprint()\nprint(submission.to_string(index=False))\n\nprint()\nprint(\"=\" * 80)\nprint(\"FINAL CHECK\")\nprint(\"=\" * 80)\n\nprint(\"Shape:\", submission.shape)\nprint(\"Columns:\", list(submission.columns))\nprint(\"Minimum probability:\", submission[prediction_columns].min().min())\nprint(\"Maximum probability:\", submission[prediction_columns].max().max())\n\nprint()\nprint(\"✓ submission.csv created\")\nprint(\"✓ Continuous probabilities preserved\")\nprint(\"✓ No thresholding\")\nprint(\"✓ No 0/1 conversion\")\nprint(\"✓ Sample format preserved\")\nprint(\"✓ Study order preserved\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-19T07:28:49.212759Z","iopub.execute_input":"2026-08-19T07:28:49.213566Z","iopub.status.idle":"2026-08-19T07:28:49.244521Z","shell.execute_reply.started":"2026-08-19T07:28:49.21353Z","shell.execute_reply":"2026-08-19T07:28:49.243786Z"}},"outputs":[],"execution_count":null}]}