{"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":[{"sourceType":"competition","sourceId":71549,"databundleVersionId":8561470}],"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":"2026-03-27T07:50:24.789168Z","iopub.execute_input":"2026-03-27T07:50:24.789333Z","iopub.status.idle":"2026-03-27T07:50:26.491080Z","shell.execute_reply.started":"2026-03-27T07:50:24.789317Z","shell.execute_reply":"2026-03-27T07:50:26.490151Z"}},"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":"2026-03-27T06:51:10.510747Z","iopub.execute_input":"2026-03-27T06:51:10.511006Z","iopub.status.idle":"2026-03-27T06:51:10.684808Z","shell.execute_reply.started":"2026-03-27T06:51:10.510981Z","shell.execute_reply":"2026-03-27T06:51:10.684029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_label_coordinates","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T06:51:10.685728Z","iopub.execute_input":"2026-03-27T06:51:10.686003Z","iopub.status.idle":"2026-03-27T06:51:10.699901Z","shell.execute_reply.started":"2026-03-27T06:51:10.685980Z","shell.execute_reply":"2026-03-27T06:51:10.699057Z"}},"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":"2026-03-27T06:51:10.700777Z","iopub.execute_input":"2026-03-27T06:51:10.701044Z","iopub.status.idle":"2026-03-27T06:51:10.739745Z","shell.execute_reply.started":"2026-03-27T06:51:10.701015Z","shell.execute_reply":"2026-03-27T06:51:10.738961Z"}},"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":"2026-03-27T06:51:10.741566Z","iopub.execute_input":"2026-03-27T06:51:10.741871Z","iopub.status.idle":"2026-03-27T06:51:10.754930Z","shell.execute_reply.started":"2026-03-27T06:51:10.741854Z","shell.execute_reply":"2026-03-27T06:51:10.754240Z"}},"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":"2026-03-27T06:51:10.755820Z","iopub.execute_input":"2026-03-27T06:51:10.756042Z","iopub.status.idle":"2026-03-27T06:51:10.773287Z","shell.execute_reply.started":"2026-03-27T06:51:10.756018Z","shell.execute_reply":"2026-03-27T06:51:10.772780Z"}},"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":"2026-03-27T06:51:10.773907Z","iopub.execute_input":"2026-03-27T06:51:10.774067Z","iopub.status.idle":"2026-03-27T06:51:11.619658Z","shell.execute_reply.started":"2026-03-27T06:51:10.774053Z","shell.execute_reply":"2026-03-27T06:51:11.618937Z"}},"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":"2026-03-27T06:51:11.620552Z","iopub.execute_input":"2026-03-27T06:51:11.620841Z","iopub.status.idle":"2026-03-27T06:51:12.127095Z","shell.execute_reply.started":"2026-03-27T06:51:11.620815Z","shell.execute_reply":"2026-03-27T06:51:12.126423Z"}},"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":"2026-03-27T06:51:12.127897Z","iopub.execute_input":"2026-03-27T06:51:12.128154Z","iopub.status.idle":"2026-03-27T06:51:12.202435Z","shell.execute_reply.started":"2026-03-27T06:51:12.128128Z","shell.execute_reply":"2026-03-27T06:51:12.201553Z"}},"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":"2026-03-27T06:51:12.203280Z","iopub.execute_input":"2026-03-27T06:51:12.203561Z","iopub.status.idle":"2026-03-27T06:51:12.462319Z","shell.execute_reply.started":"2026-03-27T06:51:12.203537Z","shell.execute_reply":"2026-03-27T06:51:12.461478Z"}},"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":"2026-03-27T06:51:12.463247Z","iopub.execute_input":"2026-03-27T06:51:12.463603Z","iopub.status.idle":"2026-03-27T06:51:12.765656Z","shell.execute_reply.started":"2026-03-27T06:51:12.463558Z","shell.execute_reply":"2026-03-27T06:51:12.764785Z"}},"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":"2026-03-27T06:51:12.766597Z","iopub.execute_input":"2026-03-27T06:51:12.766858Z","iopub.status.idle":"2026-03-27T06:51:12.774403Z","shell.execute_reply.started":"2026-03-27T06:51:12.766833Z","shell.execute_reply":"2026-03-27T06:51:12.773601Z"}},"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":"2026-03-27T06:51:12.775400Z","iopub.execute_input":"2026-03-27T06:51:12.775694Z","iopub.status.idle":"2026-03-27T07:04:46.932361Z","shell.execute_reply.started":"2026-03-27T06:51:12.775671Z","shell.execute_reply":"2026-03-27T07:04:46.931451Z"}},"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":"2026-03-27T07:06:07.990342Z","iopub.execute_input":"2026-03-27T07:06:07.991047Z","iopub.status.idle":"2026-03-27T07:06:10.410115Z","shell.execute_reply.started":"2026-03-27T07:06:07.991018Z","shell.execute_reply":"2026-03-27T07:06:10.409307Z"}},"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":"2026-03-27T07:06:19.395340Z","iopub.execute_input":"2026-03-27T07:06:19.395773Z","iopub.status.idle":"2026-03-27T07:06:19.420039Z","shell.execute_reply.started":"2026-03-27T07:06:19.395748Z","shell.execute_reply":"2026-03-27T07:06:19.419402Z"}},"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,"execution":{"iopub.status.busy":"2026-03-27T07:06:35.166871Z","iopub.execute_input":"2026-03-27T07:06:35.167731Z","iopub.status.idle":"2026-03-27T07:06:36.045101Z","shell.execute_reply.started":"2026-03-27T07:06:35.167704Z","shell.execute_reply":"2026-03-27T07:06:36.044439Z"}},"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,"execution":{"iopub.status.busy":"2026-03-27T07:06:36.046294Z","iopub.execute_input":"2026-03-27T07:06:36.046575Z","iopub.status.idle":"2026-03-27T07:06:40.224623Z","shell.execute_reply.started":"2026-03-27T07:06:36.046557Z","shell.execute_reply":"2026-03-27T07:06:40.223978Z"}},"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,"execution":{"iopub.status.busy":"2026-03-27T07:06:40.559356Z","iopub.execute_input":"2026-03-27T07:06:40.559965Z","iopub.status.idle":"2026-03-27T07:06:40.570723Z","shell.execute_reply.started":"2026-03-27T07:06:40.559945Z","shell.execute_reply":"2026-03-27T07:06:40.570108Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"# ---------- CLEANED SCRIPT (EfficientNet-B0 only) ----------\nimport os, random, time\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\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\nfrom PIL import Image\nimport timm\n\nfrom sklearn.preprocessing import LabelEncoder, label_binarize\nfrom sklearn.metrics import confusion_matrix, classification_report, roc_curve, auc\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# LOAD DATA\n# ----------------------------\ndf = train_ready.copy()\n\nif \"filename\" not in df.columns:\n    df[\"filename\"] = df[\"image_path\"].apply(lambda p: os.path.basename(p))\n\nle = LabelEncoder()\ndf[\"label\"] = le.fit_transform(df[\"severity\"].astype(str))\nclass_names = list(le.classes_)\nnum_classes = len(class_names)\n\n# class weights\nseverity_weight_map = {\"Normal/Mild\":1.0, \"Moderate\":2.0, \"Severe\":4.0}\nclass_weight_values = np.array([severity_weight_map[c] for c in class_names], dtype=np.float32)\n\n# ----------------------------\n# SPLIT\n# ----------------------------\nstudies = df[\"study_id\"].unique()\nnp.random.shuffle(studies)\n\nn_test = int(0.05 * len(studies))\ntest_studies = set(studies[:n_test])\nremaining = studies[n_test:]\n\nn_val = int(0.2 * len(remaining))\nval_studies = set(remaining[:n_val])\ntrain_studies = set(remaining[n_val:])\n\ntrain_df = df[df[\"study_id\"].isin(train_studies)].reset_index(drop=True)\nval_df   = df[df[\"study_id\"].isin(val_studies)].reset_index(drop=True)\ntest_df  = df[df[\"study_id\"].isin(test_studies)].reset_index(drop=True)\n\n# ----------------------------\n# DATASET\n# ----------------------------\ntrain_tf = transforms.Compose([\n    transforms.Resize((224,224)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(8),\n    transforms.ColorJitter(0.1,0.1),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])\n])\n\nval_tf = 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\n        self.patch_dir = patch_dir\n        self.transform = transform\n\n    def __len__(self): return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = Image.open(os.path.join(self.patch_dir, row[\"filename\"])).convert(\"RGB\")\n        if self.transform: img = self.transform(img)\n        return img, int(row[\"label\"])\n\ntrain_ds = SpineDataset(train_df, PATCH_DIR, train_tf)\nval_ds   = SpineDataset(val_df, PATCH_DIR, val_tf)\ntest_ds  = SpineDataset(test_df, PATCH_DIR, val_tf)\n\n# sampler\ncounts = train_df[\"label\"].value_counts().sort_index().values\nweights = 1.0 / (counts + 1e-8)\nsample_weights = [weights[l] for l in train_df[\"label\"]]\nsampler = WeightedRandomSampler(sample_weights, len(sample_weights))\n\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler, num_workers=NUM_WORKERS)\nval_loader   = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)\ntest_loader  = DataLoader(test_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)\n\n# ----------------------------\n# MODEL\n# ----------------------------\nmodel = timm.create_model(\"efficientnet_b0\", pretrained=True, num_classes=num_classes).to(DEVICE)\n\n# ----------------------------\n# TRAIN FUNCTION\n# ----------------------------\ndef train_model(model):\n    cw = torch.tensor(class_weight_values, device=DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=cw)\n    optimizer = optim.AdamW(model.parameters(), lr=LR)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=2)\n\n    best_loss = float(\"inf\")\n    history = {\"train_loss\":[], \"val_loss\":[], \"val_acc\":[]}\n\n    for epoch in range(NUM_EPOCHS):\n        model.train()\n        total_loss = 0\n\n        for x,y in train_loader:\n            x,y = x.to(DEVICE), y.to(DEVICE)\n            optimizer.zero_grad()\n            out = model(x)\n            loss = criterion(out,y)\n            loss.backward()\n            optimizer.step()\n            total_loss += loss.item()\n\n        # validation\n        model.eval()\n        val_loss, correct, total = 0,0,0\n        all_preds, all_probs, all_labels = [],[],[]\n\n        with torch.no_grad():\n            for x,y in val_loader:\n                x,y = x.to(DEVICE), y.to(DEVICE)\n                out = model(x)\n                loss = criterion(out,y)\n\n                probs = torch.softmax(out,1)\n                preds = probs.argmax(1)\n\n                val_loss += loss.item()\n                correct += (preds==y).sum().item()\n                total += y.size(0)\n\n                all_preds.extend(preds.cpu())\n                all_probs.extend(probs.cpu())\n                all_labels.extend(y.cpu())\n\n        val_acc = correct/total\n        history[\"train_loss\"].append(total_loss)\n        history[\"val_loss\"].append(val_loss)\n        history[\"val_acc\"].append(val_acc)\n\n        print(f\"Epoch {epoch+1}: val_acc={val_acc:.4f}\")\n\n        if val_loss < best_loss:\n            best_loss = val_loss\n            torch.save(model.state_dict(), \"/kaggle/working/efficientnet_b0_best.pth\")\n\n    return history, np.array(all_labels), np.array(all_preds), np.array(all_probs)\n\n# ----------------------------\n# TRAIN\n# ----------------------------\nhistory, y_true, y_pred, y_prob = train_model(model)\n\n# ----------------------------\n# TEST\n# ----------------------------\ndef test_model(model):\n    model.eval()\n    all_preds, all_probs, all_labels = [],[],[]\n\n    with torch.no_grad():\n        for x,y in test_loader:\n            x = x.to(DEVICE)\n            out = model(x)\n            probs = torch.softmax(out,1)\n            preds = probs.argmax(1)\n\n            all_preds.extend(preds.cpu())\n            all_probs.extend(probs.cpu())\n            all_labels.extend(y)\n\n    return np.array(all_labels), np.array(all_preds), np.array(all_probs)\n\ny_true_test, y_pred_test, y_prob_test = test_model(model)\n\n# ----------------------------\n# REPORT\n# ----------------------------\nprint(classification_report(y_true_test, y_pred_test, target_names=class_names))\n\n# ----------------------------\n# CONFUSION MATRIX\n# ----------------------------\ncm = confusion_matrix(y_true_test, y_pred_test)\nsns.heatmap(cm, annot=True, fmt='d')\nplt.title(\"Test Confusion Matrix\")\nplt.show()\n\n# ----------------------------\n# ROC\n# ----------------------------\ny_bin = label_binarize(y_true_test, classes=list(range(num_classes)))\nfor i in range(num_classes):\n    fpr,tpr,_ = roc_curve(y_bin[:,i], y_prob_test[:,i])\n    plt.plot(fpr,tpr,label=f\"{class_names[i]}\")\n\nplt.legend()\nplt.title(\"ROC Curve\")\nplt.show()\n\n# ----------------------------\n# LOSS / ACC PLOTS\n# ----------------------------\nplt.plot(history[\"val_loss\"])\nplt.title(\"Validation Loss\")\nplt.show()\n\nplt.plot(history[\"val_acc\"])\nplt.title(\"Validation Accuracy\")\nplt.show()\n\nprint(\"MODEL SAVED AT: /kaggle/working/efficientnet_b0_best.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T07:07:06.081537Z","iopub.execute_input":"2026-03-27T07:07:06.081930Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.utils.prune as prune\n\ndef apply_pruning(model):\n    for name, module in model.named_modules():\n        if isinstance(module, nn.Conv2d):\n            prune.l1_unstructured(module, name='weight', amount=0.2)  # prune 20%\n\n    print(\"Pruning applied (20% weights removed)\")\n    return model","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = apply_pruning(model)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def remove_pruning(model):\n    for module in model.modules():\n        if isinstance(module, nn.Conv2d):\n            try:\n                prune.remove(module, 'weight')\n            except:\n                pass\n    return model\n\nmodel = remove_pruning(model)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def apply_quantization(model):\n    model.cpu()  # required\n\n    quantized_model = torch.quantization.quantize_dynamic(\n        model,\n        {nn.Linear},\n        dtype=torch.qint8\n    )\n\n    print(\"Quantization applied (INT8)\")\n    return quantized_model","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"quantized_model = apply_quantization(model)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(quantized_model.state_dict(), \"efficientnet_b0_quantized.pth\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\n\ndef measure_inference(model, sample):\n    model.eval()\n    start = time.time()\n\n    with torch.no_grad():\n        _ = model(sample)\n\n    end = time.time()\n    return end - start\n\n\nsample = torch.randn(1,3,224,224)\n\nt1 = measure_inference(model, sample)\nt2 = measure_inference(quantized_model, sample)\n\nprint(\"Original model time:\", t1)\nprint(\"Quantized model time:\", t2)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}