{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport os\n\n# Path where Kaggle mounts the competition dataset\ndata_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification\"\n\n# List what’s inside\nprint(os.listdir(data_path))\n\n# Load train.csv\ntrain_df = pd.read_csv(os.path.join(data_path, \"train.csv\"))\n\ntrain_df.shape\ntrain_df.head()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:34:20.946925Z","iopub.execute_input":"2025-11-20T17:34:20.947161Z","iopub.status.idle":"2025-11-20T17:34:23.493396Z","shell.execute_reply.started":"2025-11-20T17:34:20.947137Z","shell.execute_reply":"2025-11-20T17:34:23.492641Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_label_coordinates = pd.read_csv(os.path.join(data_path, \"train_label_coordinates.csv\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:34:31.125927Z","iopub.execute_input":"2025-11-20T17:34:31.126201Z","iopub.status.idle":"2025-11-20T17:34:31.255908Z","shell.execute_reply.started":"2025-11-20T17:34:31.126174Z","shell.execute_reply":"2025-11-20T17:34:31.254958Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_label_coordinates","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:34:31.513494Z","iopub.execute_input":"2025-11-20T17:34:31.513781Z","iopub.status.idle":"2025-11-20T17:34:31.530879Z","shell.execute_reply.started":"2025-11-20T17:34:31.513760Z","shell.execute_reply":"2025-11-20T17:34:31.529969Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\ndata_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification\"\n\n# list a few study folders\nprint(os.listdir(os.path.join(data_path, \"train_images\"))[:5])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:34:31.823244Z","iopub.execute_input":"2025-11-20T17:34:31.823539Z","iopub.status.idle":"2025-11-20T17:34:31.870545Z","shell.execute_reply.started":"2025-11-20T17:34:31.823517Z","shell.execute_reply":"2025-11-20T17:34:31.869734Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"study_id = \"4003253\"\nstudy_path = os.path.join(data_path, \"train_images\", study_id)\n\nprint(\"Series inside study:\", os.listdir(study_path)[:5])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:34:32.142305Z","iopub.execute_input":"2025-11-20T17:34:32.142586Z","iopub.status.idle":"2025-11-20T17:34:32.153255Z","shell.execute_reply.started":"2025-11-20T17:34:32.142564Z","shell.execute_reply":"2025-11-20T17:34:32.152482Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"series_id = os.listdir(study_path)[0]  # take first series\nseries_path = os.path.join(study_path, series_id)\n\nprint(\"DICOM slices inside series:\", os.listdir(series_path)[:5])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:34:39.319108Z","iopub.execute_input":"2025-11-20T17:34:39.319887Z","iopub.status.idle":"2025-11-20T17:34:39.335706Z","shell.execute_reply.started":"2025-11-20T17:34:39.319860Z","shell.execute_reply":"2025-11-20T17:34:39.335088Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\nimport matplotlib.pyplot as plt\n\ndcm_file = os.path.join(series_path, os.listdir(series_path)[0])  # first slice\ndcm = pydicom.dcmread(dcm_file)\n\nprint(dcm)  # metadata\nplt.imshow(dcm.pixel_array, cmap='gray')\nplt.axis('off')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:34:39.650030Z","iopub.execute_input":"2025-11-20T17:34:39.650289Z","iopub.status.idle":"2025-11-20T17:34:40.539883Z","shell.execute_reply.started":"2025-11-20T17:34:39.650269Z","shell.execute_reply":"2025-11-20T17:34:40.539199Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\nimport matplotlib.pyplot as plt\nimport os\n\n# Pick one study and series (from your earlier code)\nstudy_id = \"4003253\"\nseries_id = os.listdir(os.path.join(data_path, \"train_images\", study_id))[0]\nseries_path = os.path.join(data_path, \"train_images\", study_id, series_id)\n\n# Get all slices in this series and sort by filename (important!)\ndcm_files = sorted(os.listdir(series_path))\n\n# Pick 5 slices spread across the series\nsample_slices = [dcm_files[i] for i in [0, len(dcm_files)//4, len(dcm_files)//2, 3*len(dcm_files)//4, -1]]\n\n# Plot them\nfig, axes = plt.subplots(1, 5, figsize=(20, 5))\n\nfor ax, fname in zip(axes, sample_slices):\n    dcm_path = os.path.join(series_path, fname)\n    dcm = pydicom.dcmread(dcm_path)\n    ax.imshow(dcm.pixel_array, cmap='gray')\n    ax.set_title(fname)\n    ax.axis('off')\n\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:34:40.540888Z","iopub.execute_input":"2025-11-20T17:34:40.541095Z","iopub.status.idle":"2025-11-20T17:34:41.083323Z","shell.execute_reply.started":"2025-11-20T17:34:40.541078Z","shell.execute_reply":"2025-11-20T17:34:41.082403Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport os\n\ncoords_path = os.path.join(data_path, \"train_label_coordinates.csv\")\ncoords_df = pd.read_csv(coords_path)\n\nprint(coords_df.shape)\nprint(coords_df.head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:34:41.084173Z","iopub.execute_input":"2025-11-20T17:34:41.084420Z","iopub.status.idle":"2025-11-20T17:34:41.153758Z","shell.execute_reply.started":"2025-11-20T17:34:41.084400Z","shell.execute_reply":"2025-11-20T17:34:41.152948Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\nimport matplotlib.pyplot as plt\nimport os\n\n# Pick one row from coords_df\nrow = coords_df.iloc[0]\n\nstudy_id = str(row.study_id)\nseries_id = str(row.series_id)\ninstance_number = int(row.instance_number)\n\n# Build path to the series folder\nseries_path = os.path.join(data_path, \"train_images\", study_id, series_id)\n\n# Get all slices (files) in this series\ndcm_files = sorted(os.listdir(series_path))\n\n# instance_number in metadata starts from 1, so adjust index\ndcm_file = os.path.join(series_path, dcm_files[instance_number - 1])\n\n# Load the DICOM slice\ndcm = pydicom.dcmread(dcm_file)\nimg = dcm.pixel_array\n\n# Plot image with coordinate marked\nplt.figure(figsize=(6,6))\nplt.imshow(img, cmap='gray')\nplt.scatter(row.x, row.y, c='red', s=40, label=f\"{row.condition} {row.level}\")\nplt.legend()\nplt.axis('off')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:34:41.155456Z","iopub.execute_input":"2025-11-20T17:34:41.156043Z","iopub.status.idle":"2025-11-20T17:34:41.404407Z","shell.execute_reply.started":"2025-11-20T17:34:41.156022Z","shell.execute_reply":"2025-11-20T17:34:41.403621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\nimport matplotlib.pyplot as plt\n\n# Pick a study to visualize\nstudy_id = \"4003253\"\n\n# Filter coords for this study\nstudy_coords = coords_df[coords_df.study_id == int(study_id)]\n\n# Pick one series from this study\nseries_id = str(study_coords.series_id.iloc[0])\nseries_path = os.path.join(data_path, \"train_images\", study_id, series_id)\n\n# Get all slices in this series\ndcm_files = sorted(os.listdir(series_path))\n\n# Choose a slice number that has multiple labels (for demonstration)\nslice_num = study_coords.instance_number.iloc[0]\ndcm_file = os.path.join(series_path, dcm_files[slice_num - 1])\n\n# Load the slice\ndcm = pydicom.dcmread(dcm_file)\nimg = dcm.pixel_array\n\n# Plot slice with all coordinates from this slice\nplt.figure(figsize=(7,7))\nplt.imshow(img, cmap='gray')\n\nfor _, row in study_coords[study_coords.instance_number == slice_num].iterrows():\n    plt.scatter(row.x, row.y, c='red', s=40)\n    plt.text(row.x+5, row.y+5, f\"{row.level}\", color='yellow', fontsize=9)\n\nplt.title(f\"Study {study_id} - Slice {slice_num}\")\nplt.axis('off')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:34:41.405118Z","iopub.execute_input":"2025-11-20T17:34:41.405307Z","iopub.status.idle":"2025-11-20T17:34:41.694289Z","shell.execute_reply.started":"2025-11-20T17:34:41.405293Z","shell.execute_reply":"2025-11-20T17:34:41.693460Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_study_ids = set(train_df.study_id.astype(str))\nimage_study_ids = set(os.listdir(os.path.join(data_path, \"train_images\")))\n\nmissing_studies = train_study_ids - image_study_ids\nextra_studies = image_study_ids - train_study_ids\n\nprint(\"Missing studies in images:\", len(missing_studies))\nprint(\"Extra studies not in train.csv:\", len(extra_studies))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:34:41.695456Z","iopub.execute_input":"2025-11-20T17:34:41.695701Z","iopub.status.idle":"2025-11-20T17:34:41.703672Z","shell.execute_reply.started":"2025-11-20T17:34:41.695683Z","shell.execute_reply":"2025-11-20T17:34:41.702825Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Normalise and cropping","metadata":{}},{"cell_type":"code","source":"import os\nimport pydicom\nimport numpy as np\nimport cv2\nfrom tqdm import tqdm\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\n# Base path (same as before)\ndata_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification\"\n\n# Parameters\nPATCH_SIZE = 128\nOUTPUT_DIR = \"/kaggle/working/patches_windowed\"\n\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n\n# Load label coordinates\ncoords_df = pd.read_csv(os.path.join(data_path, \"train_label_coordinates.csv\"))\n\n# ============== Helper Functions ==============\n\ndef window_image(img, center, width):\n    \"\"\"Apply DICOM windowing to convert 16-bit -> 8-bit visible range.\"\"\"\n    img_min = center - width / 2\n    img_max = center + width / 2\n    img = np.clip(img, img_min, img_max)\n    img = (img - img_min) / (img_max - img_min)\n    img = (img * 255).astype(np.uint8)\n    return img\n\ndef safe_crop(img, center_x, center_y, size=PATCH_SIZE):\n    \"\"\"Crop square patch centered at (x,y) safely even near borders.\"\"\"\n    h, w = img.shape\n    half = size // 2\n    x1, x2 = center_x - half, center_x + half\n    y1, y2 = center_y - half, center_y + half\n    patch = np.zeros((size, size), dtype=img.dtype)\n    x1_img, x2_img = max(0, x1), min(w, x2)\n    y1_img, y2_img = max(0, y1), min(h, y2)\n    x1_patch, y1_patch = x1_img - x1, y1_img - y1\n    x2_patch, y2_patch = x1_patch + (x2_img - x1_img), y1_patch + (y2_img - y1_img)\n    patch[y1_patch:y2_patch, x1_patch:x2_patch] = img[y1_img:y2_img, x1_img:x2_img]\n    return patch\n\n# ============== Cropping Loop ==============\n\nsaved, skipped = 0, 0\nsample_patches = []  # to visualize\n\nfor idx, row in tqdm(coords_df.iterrows(), total=len(coords_df)):\n    study_id = str(row[\"study_id\"])\n    series_id = str(row[\"series_id\"])\n    inst_no = int(row[\"instance_number\"])\n    x, y = int(row[\"x\"]), int(row[\"y\"])\n\n    condition = str(row['condition']).replace(\" \", \"_\")\n    level = str(row['level']).replace(\"/\", \"-\")\n\n    series_path = os.path.join(data_path, \"train_images\", study_id, series_id)\n\n    try:\n        dcm_file = os.path.join(series_path, f\"{inst_no}.dcm\")\n        dcm = pydicom.dcmread(dcm_file)\n        img = dcm.pixel_array.astype(np.float32)\n\n        # Apply DICOM windowing\n        wc = float(getattr(dcm, \"WindowCenter\", 300))\n        ww = float(getattr(dcm, \"WindowWidth\", 600))\n        img = window_image(img, wc, ww)\n\n        patch = safe_crop(img, x, y, PATCH_SIZE)\n\n        # Skip if patch is nearly blank (to avoid pure black)\n        if np.mean(patch) < 5:\n            skipped += 1\n            continue\n\n        out_path = os.path.join(OUTPUT_DIR, f\"{study_id}_{series_id}_{inst_no}_{condition}_{level}.png\")\n        ok = cv2.imwrite(out_path, patch)\n\n        if ok:\n            saved += 1\n            # collect a few random samples for preview\n            if len(sample_patches) < 5 and np.random.rand() < 0.001:\n                sample_patches.append((patch, f\"{condition}_{level}\"))\n        else:\n            skipped += 1\n    except Exception as e:\n        skipped += 1\n        continue\n\nprint(f\"✅ Done. Saved: {saved} patches | Skipped: {skipped}\")\nprint(f\"✅ Output dir: {OUTPUT_DIR}\")\n\n# ============== Quick Visualization ==============\n\nif sample_patches:\n    plt.figure(figsize=(15,3))\n    for i, (patch, title) in enumerate(sample_patches):\n        plt.subplot(1, len(sample_patches), i+1)\n        plt.imshow(patch, cmap='gray')\n        plt.title(title)\n        plt.axis('off')\n    plt.tight_layout()\n    plt.show()\nelse:\n    print(\"⚠️ No sample patches collected for preview. (Try increasing probability)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:34:47.971922Z","iopub.execute_input":"2025-11-20T17:34:47.972185Z","iopub.status.idle":"2025-11-20T17:47:58.374614Z","shell.execute_reply.started":"2025-11-20T17:34:47.972166Z","shell.execute_reply":"2025-11-20T17:47:58.373816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport random\nimport cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nPATCH_DIR = \"/kaggle/working/patches_windowed\"\n\n# List all patches\nall_patches = os.listdir(PATCH_DIR)\nprint(f\"Total patches found: {len(all_patches)}\")\n\n# Randomly sample up to 12 patches for visualization\nsample_files = random.sample(all_patches, min(12, len(all_patches)))\n\nplt.figure(figsize=(12, 8))\nfor i, fname in enumerate(sample_files):\n    img_path = os.path.join(PATCH_DIR, fname)\n    img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n    plt.subplot(3, 4, i + 1)\n    plt.imshow(img, cmap='gray')\n    plt.title(fname.split(\"_\")[3])  # condition\n    plt.axis('off')\nplt.tight_layout()\nplt.show()\n\n# ---------- Black Pixel Check ----------\ndef is_black_patch(img, threshold=10, black_ratio=0.9):\n    \"\"\"\n    Returns True if more than 90% of pixels are below intensity 10.\n    \"\"\"\n    return np.mean(img < threshold) > black_ratio\n\nblack_count = 0\ncheck_limit = min(1000, len(all_patches))  # check up to 1000 randomly\n\nfor fname in random.sample(all_patches, check_limit):\n    img_path = os.path.join(PATCH_DIR, fname)\n    img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n    if img is not None and is_black_patch(img):\n        black_count += 1\n\nblack_ratio = black_count / check_limit * 100\nprint(f\"\\n🧩 Quality check summary:\")\nprint(f\"Checked {check_limit} random patches\")\nprint(f\"Black/empty patches: {black_count} ({black_ratio:.2f}%)\")\n\nif black_ratio > 20:\n    print(\"⚠️ Warning: Too many black patches! You may need to adjust crop or window settings.\")\nelse:\n    print(\"✅ Looks good! Most patches contain visible anatomical content.\")\nimport os\nimport random\nimport cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nPATCH_DIR = \"/kaggle/working/patches_windowed\"\n\n# List all patches\nall_patches = os.listdir(PATCH_DIR)\nprint(f\"Total patches found: {len(all_patches)}\")\n\n# Randomly sample up to 12 patches for visualization\nsample_files = random.sample(all_patches, min(12, len(all_patches)))\n\nplt.figure(figsize=(12, 8))\nfor i, fname in enumerate(sample_files):\n    img_path = os.path.join(PATCH_DIR, fname)\n    img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n    plt.subplot(3, 4, i + 1)\n    plt.imshow(img, cmap='gray')\n    plt.title(fname.split(\"_\")[3])  # condition\n    plt.axis('off')\nplt.tight_layout()\nplt.show()\n\n# ---------- Black Pixel Check ----------\ndef is_black_patch(img, threshold=10, black_ratio=0.9):\n    \"\"\"\n    Returns True if more than 90% of pixels are below intensity 10.\n    \"\"\"\n    return np.mean(img < threshold) > black_ratio\n\nblack_count = 0\ncheck_limit = min(1000, len(all_patches))  # check up to 1000 randomly\n\nfor fname in random.sample(all_patches, check_limit):\n    img_path = os.path.join(PATCH_DIR, fname)\n    img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n    if img is not None and is_black_patch(img):\n        black_count += 1\n\nblack_ratio = black_count / check_limit * 100\nprint(f\"\\n🧩 Quality check summary:\")\nprint(f\"Checked {check_limit} random patches\")\nprint(f\"Black/empty patches: {black_count} ({black_ratio:.2f}%)\")\n\nif black_ratio > 20:\n    print(\"⚠️ Warning: Too many black patches! You may need to adjust crop or window settings.\")\nelse:\n    print(\"✅ Looks good! Most patches contain visible anatomical content.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:50:04.799637Z","iopub.execute_input":"2025-11-20T17:50:04.800321Z","iopub.status.idle":"2025-11-20T17:50:07.444580Z","shell.execute_reply.started":"2025-11-20T17:50:04.800296Z","shell.execute_reply":"2025-11-20T17:50:07.443928Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train data prep","metadata":{}},{"cell_type":"code","source":"train_csv = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv\")\nprint(train_csv.head())\nprint(train_csv.columns.tolist())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T17:52:39.530006Z","iopub.execute_input":"2025-11-20T17:52:39.530840Z","iopub.status.idle":"2025-11-20T17:52:39.551913Z","shell.execute_reply.started":"2025-11-20T17:52:39.530807Z","shell.execute_reply":"2025-11-20T17:52:39.551258Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport os\n\n# Paths\ndata_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification\"\n\n# Load CSVs\ncoords_df = pd.read_csv(os.path.join(data_path, \"train_label_coordinates.csv\"))\ntrain_df = pd.read_csv(os.path.join(data_path, \"train.csv\"))\n\n# Define mapping from \"condition\" + \"level\" → column name in train.csv\ndef get_column_name(row):\n    cond = row[\"condition\"]\n    lvl = row[\"level\"].replace(\"/\", \"_\")\n    \n    if \"Spinal Canal\" in cond:\n        return f\"spinal_canal_stenosis_{lvl.lower()}\"\n    elif \"Left Neural\" in cond:\n        return f\"left_neural_foraminal_narrowing_{lvl.lower()}\"\n    elif \"Right Neural\" in cond:\n        return f\"right_neural_foraminal_narrowing_{lvl.lower()}\"\n    elif \"Left Subarticular\" in cond:\n        return f\"left_subarticular_stenosis_{lvl.lower()}\"\n    elif \"Right Subarticular\" in cond:\n        return f\"right_subarticular_stenosis_{lvl.lower()}\"\n    else:\n        return None\n\ncoords_df[\"target_col\"] = coords_df.apply(get_column_name, axis=1)\n\n# Merge severity from train.csv\nmerged_df = coords_df.merge(\n    train_df,\n    on=\"study_id\",\n    how=\"left\"\n)\n\n# Add actual severity labels from the right column\nmerged_df[\"severity\"] = merged_df.apply(lambda r: r[r[\"target_col\"]] if pd.notnull(r[\"target_col\"]) else None, axis=1)\n\nmerged_df = merged_df[[\"study_id\", \"series_id\", \"instance_number\", \"condition\", \"level\", \"x\", \"y\", \"severity\"]]\nmerged_df.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport os\nimport pandas as pd\n\nPATCH_DIR = \"/kaggle/working/patches_windowed\"  # ✅ use your real patch folder\n\n# 1. Get filenames\npatch_files = os.listdir(PATCH_DIR)\npatch_df = pd.DataFrame({\"filename\": patch_files})\n\n# 2. Parse study_id, series_id, instance_number, condition, level\ndef parse_filename(fname):\n    parts = fname.replace(\".png\", \"\").split(\"_\")\n    study_id, series_id, instance_number = parts[:3]\n    condition = \"_\".join(parts[3:-1])  # includes 'Stenosis' or 'Narrowing'\n    level = parts[-1]                  # e.g., 'L4-L5'\n    return pd.Series([study_id, series_id, int(instance_number), condition, level])\n\npatch_df[[\"study_id\", \"series_id\", \"instance_number\", \"condition\", \"level\"]] = \\\n    patch_df[\"filename\"].apply(parse_filename)\n\n\ndef normalize_text(s):\n    return (\n        s.strip()\n         .replace(\" \", \"_\")\n         .replace(\"/\", \"_\")\n         .replace(\"-\", \"_\")\n         .replace(\"__\", \"_\")\n         .lower()\n    )\n\n# Normalize text fields\nmerged_df[\"condition_norm\"] = merged_df[\"condition\"].apply(normalize_text)\nmerged_df[\"level_norm\"] = merged_df[\"level\"].apply(normalize_text)\npatch_df[\"condition_norm\"] = patch_df[\"condition\"].apply(normalize_text)\npatch_df[\"level_norm\"] = patch_df[\"level\"].apply(normalize_text)\n\n# Ensure matching datatypes\nfor col in [\"study_id\", \"series_id\", \"instance_number\"]:\n    merged_df[col] = merged_df[col].astype(str)\n    patch_df[col] = patch_df[col].astype(str)\n\n# Merge labels\ntrain_ready = patch_df.merge(\n    merged_df[\n        [\"study_id\", \"series_id\", \"instance_number\", \"condition_norm\", \"level_norm\", \"severity\"]\n    ],\n    on=[\"study_id\", \"series_id\", \"instance_number\", \"condition_norm\", \"level_norm\"],\n    how=\"left\"\n)\n\n# Drop unmatched\ntrain_ready = train_ready.dropna(subset=[\"severity\"]).reset_index(drop=True)\n\nprint(\"✅ Patches matched correctly:\", len(train_ready))\nprint(\"🔸 Unique severity labels:\", train_ready[\"severity\"].unique())\ntrain_ready.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Encode string labels → numeric\nseverity_map = {\n    \"Normal/Mild\": 0,\n    \"Moderate\": 1,\n    \"Severe\": 2\n}\ntrain_ready[\"severity_encoded\"] = train_ready[\"severity\"].map(severity_map)\n\nprint(train_ready[\"severity_encoded\"].value_counts())\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"# ---------- Fixed script (minimal changes) ----------\nimport os, random, time\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nfrom collections import defaultdict\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torchvision import transforms, models\nfrom PIL import Image\nimport timm\n\nfrom sklearn.preprocessing import LabelEncoder, label_binarize\nfrom sklearn.metrics import (\n    confusion_matrix, classification_report,\n    roc_curve, auc\n)\n\n# ----------------------------\n# CONFIG\n# ----------------------------\nPATCH_DIR = \"/kaggle/working/patches_windowed\"\nBATCH_SIZE = 32\nNUM_EPOCHS = 15\nLR = 1e-4\nPATIENCE = 3\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nSEED = 42\nNUM_WORKERS = 4\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif DEVICE.type == \"cuda\":\n    torch.cuda.manual_seed_all(SEED)\n\n# ----------------------------\n# 1. LOAD METADATA (train_ready)\n# ----------------------------\ndf = train_ready.copy()   # <-- your variable from earlier cell\n\n# make sure we have filename\nif \"filename\" not in df.columns:\n    if \"image_path\" in df.columns:\n        df[\"filename\"] = df[\"image_path\"].apply(lambda p: os.path.basename(p))\n    else:\n        raise RuntimeError(\"train_ready must contain 'filename' or 'image_path'.\")\n\n# severity encoding\nif \"severity\" in df.columns:\n    le = LabelEncoder()\n    df[\"severity_str\"] = df[\"severity\"].astype(str)\n    df[\"label\"] = le.fit_transform(df[\"severity_str\"])\n    class_names = list(le.classes_)\nelif \"severity_encoded\" in df.columns:\n    df[\"label\"] = df[\"severity_encoded\"].astype(int)\n    class_names = [\"Normal/Mild\", \"Moderate\", \"Severe\"]\nelse:\n    raise RuntimeError(\"train_ready must contain 'severity' or 'severity_encoded'.\")\n\nnum_classes = len(np.unique(df[\"label\"]))\nprint(\"Classes:\", class_names, \"num_classes:\", num_classes)\n\n# ---- FIXED CLASS WEIGHTS 1, 2, 4 ----\nif len(class_names) != 3:\n    raise RuntimeError(\"Expected exactly 3 severity classes for weights [1,2,4].\")\n\nseverity_weight_map = {\n    \"Normal/Mild\": 1.0,\n    \"Moderate\": 2.0,\n    \"Severe\": 4.0,\n}\nclass_weight_values = np.array(\n    [severity_weight_map.get(c, 1.0) for c in class_names],\n    dtype=np.float32\n)\nprint(\"Using class weights:\", dict(zip(class_names, class_weight_values)))\n\n\n# ----------------------------\n# 2. TRAIN / VAL / TEST SPLIT (by study_id if exists)\n# ----------------------------\nif \"study_id\" in df.columns:\n    studies = df[\"study_id\"].unique()\n    np.random.shuffle(studies)\n    n = len(studies)\n\n    n_test = max(1, int(0.05 * n))  # 5% test\n    test_studies = set(studies[:n_test])\n\n    remaining = studies[n_test:]\n    n_rem = len(remaining)\n    n_val = int(0.2 * n_rem)        # 20% of remaining as val\n\n    val_studies = set(remaining[:n_val])\n    train_studies = set(remaining[n_val:])\n\n    train_df = df[df[\"study_id\"].isin(train_studies)].reset_index(drop=True)\n    val_df   = df[df[\"study_id\"].isin(val_studies)].reset_index(drop=True)\n    test_df  = df[df[\"study_id\"].isin(test_studies)].reset_index(drop=True)\n\nelse:\n    # fallback if no study_id\n    m = len(df)\n    idx = np.arange(m)\n    np.random.shuffle(idx)\n\n    n_test = max(1, int(0.05 * m))\n    test_idx = idx[:n_test]\n    remaining_idx = idx[n_test:]\n\n    n_val = int(0.2 * len(remaining_idx))\n    val_idx = remaining_idx[:n_val]\n    train_idx = remaining_idx[n_val:]\n\n    train_df = df.iloc[train_idx].reset_index(drop=True)\n    val_df   = df.iloc[val_idx].reset_index(drop=True)\n    test_df  = df.iloc[test_idx].reset_index(drop=True)\n\nprint(\"Train rows:\", len(train_df), \"Val rows:\", len(val_df), \"Test rows:\", len(test_df))\nprint(\"Train class counts:\\n\", train_df[\"label\"].value_counts().sort_index())\nprint(\"Test class counts:\\n\",  test_df[\"label\"].value_counts().sort_index())\n\n\n# ----------------------------\n# 3. DATASET + DATALOADERS  (FIXED: create train_loader, val_loader, test_loader)\n# ----------------------------\ntrain_transform = transforms.Compose([\n    transforms.Resize((224,224)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(8),\n    transforms.ColorJitter(brightness=0.12, contrast=0.12),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406], [0.229,0.224,0.225])\n])\n\nval_transform = transforms.Compose([\n    transforms.Resize((224,224)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406], [0.229,0.224,0.225])\n])\n\nclass SpineDataset(Dataset):\n    def __init__(self, df, patch_dir, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.patch_dir = patch_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = os.path.join(self.patch_dir, row[\"filename\"])\n        img = Image.open(img_path).convert(\"RGB\")\n        if self.transform:\n            img = self.transform(img)\n        label = int(row[\"label\"])\n        return img, label\n\ntrain_ds = SpineDataset(train_df, PATCH_DIR, transform=train_transform)\nval_ds   = SpineDataset(val_df,   PATCH_DIR, transform=val_transform)\ntest_ds  = SpineDataset(test_df,  PATCH_DIR, transform=val_transform)\n\n# Weighted sampler for train (you earlier built sample_weights; use it here)\nclass_counts = train_df[\"label\"].value_counts().sort_index().values\nclass_weights_inv = 1.0 / (class_counts + 1e-8)\nsample_weights = np.array([class_weights_inv[l] for l in train_df[\"label\"]], dtype=np.float32)\nsampler = WeightedRandomSampler(sample_weights, num_samples=len(sample_weights), replacement=True)\n\n# Use sampler for train_loader (do NOT set shuffle=True when using sampler)\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n                          num_workers=NUM_WORKERS, pin_memory=True)\nval_loader   = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=NUM_WORKERS, pin_memory=True)\ntest_loader  = DataLoader(test_ds, batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=NUM_WORKERS, pin_memory=True)\n\n\n# ----------------------------\n# 4. MODEL BUILDERS\n# ----------------------------\ndef build_resnet50(num_classes):\n    m = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)\n    m.fc = nn.Linear(m.fc.in_features, num_classes)\n    return m\n\ndef build_mobilenetv2(num_classes):\n    m = models.mobilenet_v2(weights=models.MobileNet_V2_Weights.IMAGENET1K_V1)\n    num_feats = m.classifier[1].in_features\n    m.classifier = nn.Sequential(\n        nn.Dropout(0.3),\n        nn.Linear(num_feats, 256),\n        nn.ReLU(),\n        nn.Dropout(0.3),\n        nn.Linear(256, num_classes)\n    )\n    return m\n\ndef build_densenet169(num_classes):\n    m = models.densenet169(weights=models.DenseNet169_Weights.IMAGENET1K_V1)\n    num_feats = m.classifier.in_features\n    m.classifier = nn.Sequential(\n        nn.Dropout(0.3),\n        nn.Linear(num_feats, 256),\n        nn.ReLU(),\n        nn.Dropout(0.3),\n        nn.Linear(256, num_classes)\n    )\n    return m\n\ndef build_effnet_b0(num_classes):\n    m = timm.create_model(\"efficientnet_b0\", pretrained=True, num_classes=num_classes)\n    return m\n\nclass ResNetEffB0Ensemble(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n        self.resnet = build_resnet50(num_classes)\n        self.effb0  = build_effnet_b0(num_classes)\n\n    def forward(self, x):\n        logits1 = self.resnet(x)\n        logits2 = self.effb0(x)\n        return (logits1 + logits2) / 2.0\n\n\n# ----------------------------\n# 5. TRAIN + EVAL FUNCTION\n# ----------------------------\ndef train_and_evaluate(model, name, num_epochs=NUM_EPOCHS, lr=LR, patience=PATIENCE):\n    model = model.to(DEVICE)\n\n    # ---- FIXED CLASS WEIGHTS 1,2,4 ----\n    cw_tensor = torch.tensor(class_weight_values, dtype=torch.float32, device=DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=cw_tensor)\n\n    optimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-5)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode=\"min\", patience=2, factor=0.5, verbose=False\n    )\n\n    best_val_loss = float(\"inf\")\n    patience_counter = 0\n\n    history = {\"train_loss\": [], \"val_loss\": [], \"train_acc\": [], \"val_acc\": []}\n    best_y_true = best_y_pred = best_y_prob = None\n\n    for epoch in range(num_epochs):\n        # ---- TRAIN ----\n        model.train()\n        running_loss = 0.0\n        running_correct = 0\n        running_total = 0\n\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n\n            running_loss += loss.item() * imgs.size(0)\n            preds = outputs.argmax(1)\n            running_correct += (preds == labels).sum().item()\n            running_total += labels.size(0)\n\n        train_loss = running_loss / running_total\n        train_acc  = running_correct / running_total\n\n        # ---- VALIDATION ----\n        model.eval()\n        val_loss = 0.0\n        val_correct = 0\n        val_total = 0\n        all_preds, all_probs, all_labels = [], [], []\n\n        with torch.no_grad():\n            for imgs, labels in val_loader:\n                imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n                outputs = model(imgs)\n                loss = criterion(outputs, labels)\n\n                probs = torch.softmax(outputs, dim=1)\n                preds = probs.argmax(1)\n\n                val_loss += loss.item() * imgs.size(0)\n                val_correct += (preds == labels).sum().item()\n                val_total += labels.size(0)\n\n                all_preds.extend(preds.cpu().numpy())\n                all_probs.extend(probs.cpu().numpy())\n                all_labels.extend(labels.cpu().numpy())\n\n        val_loss_epoch = val_loss / val_total\n        val_acc_epoch  = val_correct / val_total\n\n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss_epoch)\n        history[\"train_acc\"].append(train_acc)\n        history[\"val_acc\"].append(val_acc_epoch)\n\n        print(f\"[{name}] Epoch {epoch+1}/{num_epochs} \"\n              f\"train_loss={train_loss:.4f} train_acc={train_acc:.4f}  \"\n              f\"val_loss={val_loss_epoch:.4f} val_acc={val_acc_epoch:.4f}\")\n\n        scheduler.step(val_loss_epoch)\n\n        if val_loss_epoch < best_val_loss - 1e-6:\n            best_val_loss = val_loss_epoch\n            patience_counter = 0\n            best_y_true = np.array(all_labels)\n            best_y_pred = np.array(all_preds)\n            best_y_prob = np.array(all_probs)\n            torch.save(model.state_dict(), f\"best_{name}_val.pth\")\n        else:\n            patience_counter += 1\n            if patience_counter >= patience:\n                print(f\"[{name}] Early stopping at epoch {epoch+1}\")\n                break\n\n    # reload best validation weights before returning\n    model.load_state_dict(torch.load(f\"best_{name}_val.pth\", map_location=DEVICE))\n\n    return history, best_y_true, best_y_pred, best_y_prob\n\n\ndef evaluate_on_test(model, name):\n    model = model.to(DEVICE)\n    model.eval()\n\n    cw_tensor = torch.tensor(class_weight_values, dtype=torch.float32, device=DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=cw_tensor)\n\n    test_loss = 0.0\n    test_correct = 0\n    test_total = 0\n    all_preds, all_probs, all_labels = [], [], []\n\n    with torch.no_grad():\n        for imgs, labels in test_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n\n            probs = torch.softmax(outputs, dim=1)\n            preds = probs.argmax(1)\n\n            test_loss += loss.item() * imgs.size(0)\n            test_correct += (preds == labels).sum().item()\n            test_total += labels.size(0)\n\n            all_preds.extend(preds.cpu().numpy())\n            all_probs.extend(probs.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n\n    test_loss = test_loss / test_total if test_total else float('nan')\n    test_acc  = test_correct / test_total if test_total else float('nan')\n    print(f\"[{name}] TEST: loss={test_loss:.4f} acc={test_acc:.4f}\")\n\n    return test_loss, test_acc, np.array(all_labels), np.array(all_preds), np.array(all_probs)\n\n\n# ----------------------------\n# 6. RUN MODELS\n# ----------------------------\nmodels_dict = {\n    \"resnet50\":            build_resnet50(num_classes),\n    \"mobilenet_v2\":        build_mobilenetv2(num_classes),\n    \"densenet169\":         build_densenet169(num_classes),\n    \"efficientnet_b0\":     build_effnet_b0(num_classes),\n}\n\nresults = {}\n\nfor name, model in models_dict.items():\n    print(\"\\n\" + \"=\"*60)\n    print(\"Training (train/val):\", name)\n    start = time.time()\n\n    history, y_true_val, y_pred_val, y_prob_val = train_and_evaluate(model, name)\n    elapsed = (time.time() - start) / 60\n    print(f\"Done {name} train/val in {elapsed:.2f} min\")\n\n    # VAL REPORT\n    print(f\"\\n{name} — Validation classification report:\")\n    print(classification_report(y_true_val, y_pred_val, target_names=class_names, digits=4))\n\n    # Evaluate (no retrain) on test using the best validation weights already loaded\n    test_loss, test_acc, y_true_test, y_pred_test, y_prob_test = evaluate_on_test(model, name)\n\n    # store\n    results[name] = {\n        \"history\": history,\n        \"val_y_true\": y_true_val,\n        \"val_y_pred\": y_pred_val,\n        \"val_y_prob\": y_prob_val,\n        \"test_y_true\": y_true_test,\n        \"test_y_pred\": y_pred_test,\n        \"test_y_prob\": y_prob_test,\n        \"test_acc\": test_acc,\n        \"test_loss\": test_loss,\n    }\n\n\n    cm_val = confusion_matrix(y_true_val, y_pred_val)\n    plt.figure(figsize=(5,4))\n    sns.heatmap(cm_val, annot=True, fmt='d', cmap='Blues',\n                xticklabels=class_names, yticklabels=class_names)\n    plt.title(f\"Validation Confusion Matrix — {name}\")\n    plt.xlabel(\"Predicted\"); plt.ylabel(\"Actual\")\n    plt.tight_layout()\n    plt.show()\n\n    # VAL ROC (per model)\n    y_true_val_bin = label_binarize(y_true_val, classes=list(range(num_classes)))\n    plt.figure(figsize=(6,5))\n    for i in range(num_classes):\n        fpr, tpr, _ = roc_curve(y_true_val_bin[:, i], y_prob_val[:, i])\n        model_auc = auc(fpr, tpr)\n        plt.plot(fpr, tpr, label=f\"{class_names[i]} (AUC={model_auc:.2f})\")\n    plt.plot([0,1],[0,1],'k--',lw=1)\n    plt.title(f\"Validation ROC — {name}\")\n    plt.xlabel(\"False Positive Rate\")\n    plt.ylabel(\"True Positive Rate\")\n    plt.legend()\n    plt.grid(True)\n    plt.tight_layout()\n    plt.show()\n\n    # ---- DIRECT TEST EVAL (NO RETRAIN) ----\n    print(f\"\\nEvaluating on TEST set (no retraining) for {name}...\")\n    test_loss, test_acc, y_true_test, y_pred_test, y_prob_test = \\\n        evaluate_on_test(model, name)\n\n\n    print(f\"\\n{name} — TEST classification report:\")\n    print(classification_report(y_true_test, y_pred_test,\n                                target_names=class_names, digits=4))\n\n    cm_test = confusion_matrix(y_true_test, y_pred_test)\n    plt.figure(figsize=(5,4))\n    sns.heatmap(cm_test, annot=True, fmt='d', cmap='Blues',\n                xticklabels=class_names, yticklabels=class_names)\n    plt.title(f\"Test Confusion Matrix — {name}\")\n    plt.xlabel(\"Predicted\"); plt.ylabel(\"Actual\")\n    plt.tight_layout()\n    plt.show()\n\n    # TEST ROC (per model)\n    y_true_test_bin = label_binarize(y_true_test, classes=list(range(num_classes)))\n    plt.figure(figsize=(6,5))\n    for i in range(num_classes):\n        fpr, tpr, _ = roc_curve(y_true_test_bin[:, i], y_prob_test[:, i])\n        model_auc = auc(fpr, tpr)\n        plt.plot(fpr, tpr, label=f\"{class_names[i]} (AUC={model_auc:.2f})\")\n    plt.plot([0,1],[0,1],'k--',lw=1)\n    plt.title(f\"Test ROC — {name}\")\n    plt.xlabel(\"False Positive Rate\")\n    plt.ylabel(\"True Positive Rate\")\n    plt.legend()\n    plt.grid(True)\n    plt.tight_layout()\n    plt.show()\n\n    # store everything\n    results[name] = {\n        \"history\": history,\n        \"val_y_true\": y_true_val,\n        \"val_y_pred\": y_pred_val,\n        \"val_y_prob\": y_prob_val,\n        \"test_y_true\": y_true_test,\n        \"test_y_pred\": y_pred_test,\n        \"test_y_prob\": y_prob_test,\n        \"test_acc\": test_acc,\n        \"test_loss\": test_loss,\n    }\n\n# Validation loss comparison\nplt.figure(figsize=(12,5))\nmax_epochs = max(len(r[\"history\"][\"val_loss\"]) for r in results.values())\n\ndef pad(arr, n):\n    return arr + [np.nan]*(n - len(arr))\n\nfor name, res in results.items():\n    epochs = np.arange(1, max_epochs+1)\n    val_loss = pad(res[\"history\"][\"val_loss\"], max_epochs)\n    plt.plot(epochs, val_loss, marker='o', label=name)\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Validation Loss\")\nplt.title(\"Validation Loss — All Models\")\nplt.legend()\nplt.grid(True)\nplt.tight_layout()\nplt.show()\n\n\n# Validation accuracy comparison\nplt.figure(figsize=(12,5))\nfor name, res in results.items():\n    epochs = np.arange(1, max_epochs+1)\n    val_acc = pad(res[\"history\"][\"val_acc\"], max_epochs)\n    plt.plot(epochs, val_acc, marker='o', label=name)\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Validation Accuracy\")\nplt.title(\"Validation Accuracy — All Models\")\nplt.legend()\nplt.grid(True)\nplt.tight_layout()\nplt.show()\n\n# Validation ROC comparison per class\nfor class_idx, cname in enumerate(class_names):\n    plt.figure(figsize=(7,6))\n    for name, res in results.items():\n        y_true = res[\"val_y_true\"]\n        y_prob = res[\"val_y_prob\"]\n        y_true_bin = label_binarize(y_true, classes=list(range(num_classes)))\n        fpr, tpr, _ = roc_curve(y_true_bin[:, class_idx], y_prob[:, class_idx])\n        model_auc = auc(fpr, tpr)\n        plt.plot(fpr, tpr, lw=2, label=f\"{name} (AUC={model_auc:.2f})\")\n    plt.plot([0,1],[0,1],'k--',lw=1)\n    plt.title(f\"Validation ROC Comparison — {cname}\")\n    plt.xlabel(\"False Positive Rate\")\n    plt.ylabel(\"True Positive Rate\")\n    plt.legend()\n    plt.grid(True)\n    plt.tight_layout()\n    plt.show()\n\n# ----------------------------\n# TEST COMPARISON PLOTS\n# ----------------------------\n\n# Test accuracy comparison (bar plot)\nplt.figure(figsize=(8,5))\nmodel_names = []\ntest_accs = []\nfor name, res in results.items():\n    model_names.append(name)\n    test_accs.append(res[\"test_acc\"])\n\nplt.bar(model_names, test_accs)\nplt.ylabel(\"Test Accuracy\")\nplt.title(\"Test Accuracy — All Models\")\nplt.xticks(rotation=45, ha=\"right\")\nplt.grid(axis='y')\nplt.tight_layout()\nplt.show()\n\n# Test ROC comparison per class (all models)\nfor class_idx, cname in enumerate(class_names):\n    plt.figure(figsize=(7,6))\n    for name, res in results.items():\n        y_true = res[\"test_y_true\"]\n        y_prob = res[\"test_y_prob\"]\n        y_true_bin = label_binarize(y_true, classes=list(range(num_classes)))\n        fpr, tpr, _ = roc_curve(y_true_bin[:, class_idx], y_prob[:, class_idx])\n        model_auc = auc(fpr, tpr)\n        plt.plot(fpr, tpr, lw=2, label=f\"{name} (AUC={model_auc:.2f})\")\n    plt.plot([0,1],[0,1],'k--',lw=1)\n    plt.title(f\"Test ROC Comparison — {cname}\")\n    plt.xlabel(\"False Positive Rate\")\n    plt.ylabel(\"True Positive Rate\")\n    plt.legend()\n    plt.grid(True)\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T18:29:55.062100Z","iopub.execute_input":"2025-11-20T18:29:55.062375Z","iopub.status.idle":"2025-11-20T18:29:55.135494Z","shell.execute_reply.started":"2025-11-20T18:29:55.062351Z","shell.execute_reply":"2025-11-20T18:29:55.134525Z"}},"outputs":[],"execution_count":null}]}