{"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":"gpu","dataSources":[{"sourceType":"competition","sourceId":71549,"databundleVersionId":8561470},{"sourceType":"datasetVersion","sourceId":14353046,"datasetId":9009659,"databundleVersionId":15162491}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q git+https://github.com/qubvel/segmentation_models.pytorch","metadata":{"_uuid":"2ea6e542-b242-4947-b836-cefca876bc1f","_cell_guid":"8f5725ae-569a-4946-8309-d778c7911aef","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:08:45.073776Z","iopub.execute_input":"2025-12-31T12:08:45.074386Z","iopub.status.idle":"2025-12-31T12:09:58.976195Z","shell.execute_reply.started":"2025-12-31T12:08:45.074363Z","shell.execute_reply":"2025-12-31T12:09:58.975275Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## IMPORTS ##\nimport os\nimport glob\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport matplotlib.pyplot as plt\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torchvision\nimport segmentation_models_pytorch as smp\nfrom tqdm import tqdm\nimport gc\nfrom torch.utils.data import Dataset, DataLoader, ConcatDataset","metadata":{"_uuid":"6f685617-bba3-415e-b137-d26ff06a80d6","_cell_guid":"f67c3508-65d4-4a37-aa7e-7c36fcfe6e40","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:26:56.400391Z","iopub.execute_input":"2025-12-31T12:26:56.400700Z","iopub.status.idle":"2025-12-31T12:26:56.405179Z","shell.execute_reply.started":"2025-12-31T12:26:56.400676Z","shell.execute_reply":"2025-12-31T12:26:56.404351Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nENCODER_NAME = 'resnet18'\n\n## MRI SIZE ##\nPATCH_H = 512\nPATCH_W = 512\n\nBASE_PATH_IMG = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images\"\nSAGITTAL_MODEL_PATH = '/kaggle/input/lumbar-spine-keypoint-detection-models/Sagittal_T2_sagittal_level_segmentation_2_v2'\nAXIAL_MODEL_PATH = '/kaggle/input/lumbar-spine-keypoint-detection-models/Axial_T2_axial_side_segmentation_1'","metadata":{"_uuid":"ad5f8f50-f2c4-40ed-bb3c-1fe45ba2c748","_cell_guid":"e5b3ef5b-e875-45e3-a857-eb4002402581","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:25:20.249167Z","iopub.execute_input":"2025-12-31T12:25:20.249481Z","iopub.status.idle":"2025-12-31T12:25:20.314533Z","shell.execute_reply.started":"2025-12-31T12:25:20.249456Z","shell.execute_reply":"2025-12-31T12:25:20.313829Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myUNet(nn.Module):\n    def __init__(self, classes):\n        super(myUNet, self).__init__()\n        self.classes = classes\n        self.UNet = smp.Unet(\n            encoder_name=ENCODER_NAME,\n            classes=classes,\n            in_channels=1\n        ).to(DEVICE)\n\n    def forward(self, X):\n        H, W = X.shape[-2:]\n        x = self.UNet(X.view(-1, 1, H, W)).view(-1, H*W)\n        # MinMaxScaling along the class plane to generate a heatmap\n        min_values = x.min(-1)[0].view(-1, 1)\n        max_values = x.max(-1)[0].view(-1, 1)\n        d = (max_values - min_values)\n        d[d == 0] = 1\n        x = (x - min_values) / d\n        return x.view(-1, self.classes, H, W)","metadata":{"_uuid":"c63ded64-8f13-4487-b96b-85609edeadad","_cell_guid":"1fe03673-2dc9-498d-aa6a-925387de47b5","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:25:20.599763Z","iopub.execute_input":"2025-12-31T12:25:20.600507Z","iopub.status.idle":"2025-12-31T12:25:20.605996Z","shell.execute_reply.started":"2025-12-31T12:25:20.600479Z","shell.execute_reply":"2025-12-31T12:25:20.605215Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def pixel_to_patient_3d(dcm, x, y):\n    \"\"\"\n    Converts 2D pixel coordinates (x, y) to 3D Patient Coordinates (x, y, z).\n    x: Column index\n    y: Row index\n    \"\"\"\n    ipp = np.asarray(dcm.ImagePositionPatient, dtype=np.float64)\n    iop = np.asarray(dcm.ImageOrientationPatient, dtype=np.float64)\n    row_cosines = iop[:3]\n    col_cosines = iop[3:]\n    row_spacing, col_spacing = map(float, dcm.PixelSpacing)\n\n    return (\n        ipp\n        + x * col_spacing * col_cosines\n        + y * row_spacing * row_cosines\n    )","metadata":{"_uuid":"c3d13144-7660-4378-8d09-d2d7f3caf28e","_cell_guid":"10e13b90-2b94-4e53-8431-f3c56798f189","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:25:20.844510Z","iopub.execute_input":"2025-12-31T12:25:20.845378Z","iopub.status.idle":"2025-12-31T12:25:20.849976Z","shell.execute_reply.started":"2025-12-31T12:25:20.845348Z","shell.execute_reply":"2025-12-31T12:25:20.849178Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def collect_dicom_files(series_folder):\n    patterns = [\"*.dcm\", \"*.DCM\", \"*.dicom\"]\n    files = []\n    for p in patterns:\n        files.extend(glob.glob(os.path.join(series_folder, p)))\n    # Sort by instance number to ensure correct volume order\n    files.sort(key=lambda f: pydicom.dcmread(f, stop_before_pixels=True).InstanceNumber)\n    return files","metadata":{"_uuid":"241fce34-e125-40b7-b846-fb463017a2b0","_cell_guid":"fcaa8c2f-2d77-4c3f-bb84-ebd6879cdd03","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:25:22.560159Z","iopub.execute_input":"2025-12-31T12:25:22.560429Z","iopub.status.idle":"2025-12-31T12:25:22.565055Z","shell.execute_reply.started":"2025-12-31T12:25:22.560409Z","shell.execute_reply":"2025-12-31T12:25:22.564404Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_models():\n    \"\"\"Loads the models. Requires the class definition above to be in scope.\"\"\"\n    print(\"Loading Sagittal Model...\")\n    sag_model = torch.load(SAGITTAL_MODEL_PATH, map_location=DEVICE, weights_only=False)\n    sag_model.to(DEVICE)\n    sag_model.eval()\n\n    print(\"Loading Axial Model...\")\n    ax_model = torch.load(AXIAL_MODEL_PATH, map_location=DEVICE, weights_only=False)\n    ax_model.to(DEVICE)\n    ax_model.eval()\n    \n    return sag_model, ax_model","metadata":{"_uuid":"d574213b-261b-4665-97cf-211ae6fe6c51","_cell_guid":"52096c9b-23e1-4ecf-bd89-4c4da78a660c","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:25:22.720062Z","iopub.execute_input":"2025-12-31T12:25:22.720807Z","iopub.status.idle":"2025-12-31T12:25:22.725082Z","shell.execute_reply.started":"2025-12-31T12:25:22.720780Z","shell.execute_reply":"2025-12-31T12:25:22.724444Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch_resize = torchvision.transforms.Resize((PATCH_H, PATCH_W), antialias=True)\n\ndef preprocess_image(pixel_array):\n    image = pixel_array.astype(np.float32)\n    H, W = image.shape\n    \n    h_start, w_start = 0, 0\n    crop_size = 0\n    \n    if H > W:\n        crop_size = W\n        h_start = (H - crop_size) // 2\n        image = image[h_start : h_start + crop_size, :]\n    elif H < W:\n        crop_size = H\n        w_start = (W - crop_size) // 2\n        image = image[:, w_start : w_start + crop_size]\n    else:\n        crop_size = H\n        \n    img_max = np.max(image)\n    if img_max > 0:\n        image = image / img_max\n        \n    img_tensor = torch.tensor(image).unsqueeze(0) \n\n    ### IMPORTANT ### \n    # use the torchvision Resize for resizing, torch.interpolate gave bad results\n    # Resize to PATCH_H, PATCH_W using torchvision (Training standard)\n    img_tensor = torch_resize(img_tensor)    \n    img_tensor = img_tensor.unsqueeze(0).float().to(DEVICE)\n    \n    return img_tensor, (h_start, w_start, crop_size), (H, W)","metadata":{"_uuid":"fc375210-ecc0-4182-9903-2c8ea328597c","_cell_guid":"afcd63a1-c27c-44a5-8408-c4005a548c45","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:25:22.921498Z","iopub.execute_input":"2025-12-31T12:25:22.922545Z","iopub.status.idle":"2025-12-31T12:25:22.928527Z","shell.execute_reply.started":"2025-12-31T12:25:22.922511Z","shell.execute_reply":"2025-12-31T12:25:22.927808Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def extract_coordinates(heatmaps, crop_info, original_shape):\n    \"\"\"\n    Converts model output heatmaps back to original image coordinates.\n    \"\"\"\n    h_start, w_start, crop_size = crop_info\n    # orig_H, orig_W = original_shape # Not strictly needed for coord calc, but good for validation\n    \n    bs, n_classes, h_map, w_map = heatmaps.shape\n    coords = []\n    \n    for c in range(n_classes):\n        hm = heatmaps[0, c, :, :].detach().cpu().numpy()\n        \n        # Argmax gives (row, col) i.e. (y, x)\n        y_idx, x_idx = np.unravel_index(np.argmax(hm), hm.shape)\n        \n        # Map resize -> Crop\n        # Note: h_map and w_map should be 512\n        x_crop = x_idx * (crop_size / w_map)\n        y_crop = y_idx * (crop_size / h_map)\n        \n        # Map Crop -> Original\n        x_orig = x_crop + w_start\n        y_orig = y_crop + h_start\n        \n        coords.append((x_orig, y_orig))\n        \n    return coords","metadata":{"_uuid":"7387a896-d298-4166-b9db-23777bb68c4e","_cell_guid":"a6ca521b-cf9a-45e9-a1fe-10f7773cf5bb","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:25:23.081165Z","iopub.execute_input":"2025-12-31T12:25:23.081442Z","iopub.status.idle":"2025-12-31T12:25:23.086841Z","shell.execute_reply.started":"2025-12-31T12:25:23.081420Z","shell.execute_reply":"2025-12-31T12:25:23.085993Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## LOAD THE MODELS\nsagittal_model, axial_model = load_models()","metadata":{"_uuid":"6f7eadab-8e71-40a1-b117-28034b99d016","_cell_guid":"08ce90aa-5bb3-416f-a605-14dd261d9a7c","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:25:26.363732Z","iopub.execute_input":"2025-12-31T12:25:26.364018Z","iopub.status.idle":"2025-12-31T12:25:27.382715Z","shell.execute_reply.started":"2025-12-31T12:25:26.363997Z","shell.execute_reply":"2025-12-31T12:25:27.382089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Logic to perform inference on data. \nimport pandas as pd\ntrain_coords = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv\")\ntrain_desc = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv\")\nsample_data = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv\")\ntrain_coords.drop(columns=['instance_number'], inplace = True)\ntrain_desc = train_desc[train_desc['series_description'] != 'Sagittal T1']","metadata":{"_uuid":"8b1e11de-0437-40d9-94a9-8d6e7bb0649b","_cell_guid":"aa9c6e68-dde6-4a83-9cd9-c288eecccbc8","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:25:27.383692Z","iopub.execute_input":"2025-12-31T12:25:27.383906Z","iopub.status.idle":"2025-12-31T12:25:27.500328Z","shell.execute_reply.started":"2025-12-31T12:25:27.383889Z","shell.execute_reply":"2025-12-31T12:25:27.499756Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = sample_data.merge(train_desc, on=['study_id'], how='inner')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T12:46:20.915850Z","iopub.execute_input":"2025-12-31T12:46:20.916488Z","iopub.status.idle":"2025-12-31T12:46:20.928043Z","shell.execute_reply.started":"2025-12-31T12:46:20.916465Z","shell.execute_reply":"2025-12-31T12:46:20.927456Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T12:46:23.513260Z","iopub.execute_input":"2025-12-31T12:46:23.513946Z","iopub.status.idle":"2025-12-31T12:46:23.535297Z","shell.execute_reply.started":"2025-12-31T12:46:23.513922Z","shell.execute_reply":"2025-12-31T12:46:23.534603Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"levels = [\n    \"spinal_canal_stenosis_l1_l2\",\n    \"spinal_canal_stenosis_l2_l3\",\n    \"spinal_canal_stenosis_l3_l4\",\n    \"spinal_canal_stenosis_l4_l5\",\n    \"spinal_canal_stenosis_l5_s1\",\n]\n\ndef spinal_canal_class_frequency(df):\n    freq = {}\n    for lvl in levels:\n        freq[lvl] = df[lvl].value_counts()\n    class_freq = pd.DataFrame(freq).fillna(0).astype(int)\n    class_freq.columns = class_freq.columns.str.replace(\n        \"spinal_canal_stenosis_\", \"\", regex=False\n    )\n    return class_freq\n\n# Usage\nclass_freq = spinal_canal_class_frequency(sample_data)\nprint(class_freq)","metadata":{"_uuid":"55980bd7-261d-4ad3-ac3e-2dd87913f922","_cell_guid":"2ebffcc1-3188-4737-8ea0-142f49a79004","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:25:28.200619Z","iopub.execute_input":"2025-12-31T12:25:28.200900Z","iopub.status.idle":"2025-12-31T12:25:28.219312Z","shell.execute_reply.started":"2025-12-31T12:25:28.200880Z","shell.execute_reply":"2025-12-31T12:25:28.218614Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Dataset Class**","metadata":{"_uuid":"7e1b9d37-00c5-4cef-a9be-75e017fde9e2","_cell_guid":"fea10113-a094-4227-a001-f5734abf4133","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"markdown","source":"**Removed OTSU Thresholding - Only Sagittal & Axial slices**","metadata":{"_uuid":"92fc52d5-6ef0-4cf2-b31d-a457fb1bc0bc","_cell_guid":"179e8bcf-0127-4bfa-a6ca-0ba9e03046c9","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def load_data_into_memory(\n    df,\n    base_img_path,\n    sag_model,\n    device=\"cuda\",\n    final_size=384,        # Final output size (Reference uses 384)\n    sag_offsets=(-1, 0, 1),\n    num_axial_slices=3\n):\n    processed = []\n    levels = [\"L1/L2\", \"L2/L3\", \"L3/L4\", \"L4/L5\", \"L5/S1\"]\n    label_map = {\"Normal/Mild\": 0, \"Moderate\": 1, \"Severe\": 2}\n    \n    sag_model.eval().to(device)\n    \n    study_ids = df[\"study_id\"].unique()\n\n    for study_id in tqdm(study_ids, desc=\"Loading & Preprocessing\"):\n        try:\n            # ------------------ 1. Setup ------------------ #\n            row = df[df.study_id == study_id].iloc[0]\n            labels = torch.tensor(\n                [label_map[row[f\"spinal_canal_stenosis_{lvl.lower().replace('/', '_')}\"]] for lvl in levels],\n                dtype=torch.long\n            )\n\n            # --------- CHANGED PART: series lookup from df --------- #\n            try:\n                sag_sid = df.loc[\n                    (df.study_id == study_id) &\n                    (df.series_description == \"Sagittal T2/STIR\"),\n                    \"series_id\"\n                ].iloc[0]\n\n                ax_sid = df.loc[\n                    (df.study_id == study_id) &\n                    (df.series_description == \"Axial T2\"),\n                    \"series_id\"\n                ].iloc[0]\n            except IndexError:\n                continue\n            # ------------------------------------------------------- #\n\n            sag_path = os.path.join(base_img_path, str(study_id), str(sag_sid))\n            ax_path  = os.path.join(base_img_path, str(study_id), str(ax_sid))\n\n            sag_files = collect_dicom_files(sag_path)\n            ax_files  = collect_dicom_files(ax_path)\n\n            if not sag_files or not ax_files:\n                continue\n\n            # ------------------ 2. Sagittal Processing ------------------ #\n            mid_idx = len(sag_files) // 2\n            sag_stack = []\n\n            mid_slice_kp = None\n            mid_slice_dcm = None\n\n            for off in sag_offsets:\n                idx = max(0, min(mid_idx + off, len(sag_files) - 1))\n                dcm = pydicom.dcmread(sag_files[idx])\n                img = dcm.pixel_array\n\n                t_tensor, meta, orig_shape = preprocess_image(img)\n\n                with torch.no_grad():\n                    heatmaps = sag_model(t_tensor.to(device))\n                    coords = extract_coordinates(heatmaps, meta, orig_shape)\n\n                kp_dict = {lvl: (float(x), float(y)) for lvl, (x, y) in zip(levels, coords)}\n\n                if off == 0:\n                    mid_slice_kp = kp_dict\n                    mid_slice_dcm = dcm\n\n                img_norm = img.astype(float)\n                img_norm -= img_norm.min()\n                if img_norm.max() != 0:\n                    img_norm /= img_norm.max()\n                img_norm = (img_norm * 255).astype(np.uint8)\n\n                h, w = img_norm.shape\n                slice_crops = []\n\n                for lvl in levels:\n                    cx, cy = kp_dict[lvl]\n                    pad_h = int(0.09 * h)\n                    pad_w = int(0.09 * w)\n\n                    ymin, ymax = max(0, int(cy - pad_h)), min(h, int(cy + pad_h))\n                    xmin, xmax = max(0, int(cx - pad_w)), min(w, int(cx + pad_w))\n\n                    crop = img_norm[ymin:ymax, xmin:xmax]\n\n                    if crop.size == 0:\n                        crop = np.zeros((final_size, final_size), dtype=np.uint8)\n                    else:\n                        crop = cv2.resize(crop, (final_size, final_size), interpolation=cv2.INTER_LINEAR)\n\n                    slice_crops.append(torch.from_numpy(crop).unsqueeze(0))\n\n                sag_stack.append(torch.stack(slice_crops))\n\n            sagittal_tensor = torch.stack(sag_stack)  # [3, 5, 1, 384, 384]\n\n            # ------------------ 3. Axial Processing ------------------ #\n            ax_meta = []\n            for f in ax_files:\n                d_ax = pydicom.dcmread(f, stop_before_pixels=True)\n                ax_meta.append({\n                    \"path\": f,\n                    \"Z\": float(d_ax.ImagePositionPatient[2])\n                })\n\n            axial_stack = []\n\n            for lvl in levels:\n                cx, cy = mid_slice_kp[lvl]\n                z_target = pixel_to_patient_3d(mid_slice_dcm, cx, cy)[2]\n\n                ax_meta.sort(key=lambda a: abs(a[\"Z\"] - z_target))\n                selected_axials = ax_meta[:num_axial_slices]\n\n                lvl_axials = []\n                for meta in selected_axials:\n                    d_ax = pydicom.dcmread(meta[\"path\"])\n                    img = d_ax.pixel_array\n\n                    img_norm = img.astype(float)\n                    img_norm -= img_norm.min()\n                    if img_norm.max() != 0:\n                        img_norm /= img_norm.max()\n                    img_norm = (img_norm * 255).astype(np.uint8)\n\n                    h_ax, w_ax = img_norm.shape\n                    crop_size = 160\n                    cy_ax, cx_ax = h_ax // 2, w_ax // 2\n\n                    ymin, ymax = max(0, cy_ax - crop_size//2), min(h_ax, cy_ax + crop_size//2)\n                    xmin, xmax = max(0, cx_ax - crop_size//2), min(w_ax, cx_ax + crop_size//2)\n\n                    crop = img_norm[ymin:ymax, xmin:xmax]\n\n                    if crop.size == 0:\n                        crop = np.zeros((final_size, final_size), dtype=np.uint8)\n                    else:\n                        crop = cv2.resize(crop, (final_size, final_size), interpolation=cv2.INTER_LINEAR)\n\n                    lvl_axials.append(torch.from_numpy(crop).unsqueeze(0))\n\n                axial_stack.append(torch.stack(lvl_axials))\n\n            axial_tensor = torch.stack(axial_stack)  # [5, 3, 1, 384, 384]\n\n            # ------------------ 4. Store ------------------ #\n            processed.append({\n                \"sagittal\": sagittal_tensor.to(torch.uint8),\n                \"axial\": axial_tensor.to(torch.uint8),\n                \"label\": labels\n            })\n\n            if len(processed) % 50 == 0:\n                gc.collect()\n\n        except Exception:\n            continue\n\n    return processed","metadata":{"_uuid":"2df2ea39-1c3d-4cf0-a223-81151d3dc9d5","_cell_guid":"cbc291ac-7c0f-4227-8062-40644bbdb3ed","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:47:23.958940Z","iopub.execute_input":"2025-12-31T12:47:23.959522Z","iopub.status.idle":"2025-12-31T12:47:23.976413Z","shell.execute_reply.started":"2025-12-31T12:47:23.959498Z","shell.execute_reply":"2025-12-31T12:47:23.975585Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SpinalStenosisDataset(Dataset):\n    def __init__(self, data_list, transform=None):\n        self.data = data_list\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        sample = self.data[idx]\n\n        # Inputs are uint8 [Batch, Level, Channel, H, W]\n        sag = sample[\"sagittal\"] \n        ax = sample[\"axial\"]    \n        labels = sample[\"labels\"] if \"labels\" in sample else sample[\"label\"]\n\n        # --- Apply Augmentations ---\n        if self.transform:\n            # 1. Sagittal Processing\n            # Reshape to [N, H, W] for looping\n            B, L, C, H, W = sag.shape\n            sag_flat = sag.reshape(-1, H, W).numpy() \n            \n            sag_aug_list = []\n            for i in range(sag_flat.shape[0]):\n                # Add channel dimension for Albumentations: [H, W] -> [H, W, 1]\n                img = sag_flat[i][:, :, None] \n                \n                # Apply Transform -> Returns Tensor [C, H, W]\n                res = self.transform(image=img)[\"image\"] \n                sag_aug_list.append(res)\n            \n            # Stack: List of Tensors -> Tensor\n            sag = torch.stack(sag_aug_list).reshape(B, L, 1, 384, 384)\n\n            # 2. Axial Processing\n            L, B, C, H, W = ax.shape\n            ax_flat = ax.reshape(-1, H, W).numpy()\n            \n            ax_aug_list = []\n            for i in range(ax_flat.shape[0]):\n                img = ax_flat[i][:, :, None]\n                \n                # Apply Transform -> Returns Tensor [C, H, W]\n                res = self.transform(image=img)[\"image\"]\n                ax_aug_list.append(res)\n            \n            # Stack: List of Tensors -> Tensor\n            ax = torch.stack(ax_aug_list).reshape(L, B, 1, 384, 384)\n            \n        else:\n            # Fallback (Manual conversion if no transform provided)\n            # This path is unlikely used if you always pass transforms_val\n            sag = sag.float() / 255.0\n            ax = ax.float() / 255.0\n\n        return {\n            \"sagittal\": sag,\n            \"axial\": ax,\n            \"labels\": labels\n        }","metadata":{"_uuid":"00e5192c-858a-4e45-be71-40da5f1304b7","_cell_guid":"1f722a7d-b43c-4788-9cd8-564a4f72fb3f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T13:13:21.843724Z","iopub.execute_input":"2025-12-31T13:13:21.844321Z","iopub.status.idle":"2025-12-31T13:13:21.851513Z","shell.execute_reply.started":"2025-12-31T13:13:21.844299Z","shell.execute_reply":"2025-12-31T13:13:21.850775Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **MODEL**","metadata":{"_uuid":"cc6ec3ac-429c-46a8-a501-3de91d225bc2","_cell_guid":"1f6da345-64c8-4424-84d0-4fe068fba965","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"markdown","source":"**MODEL V3 (Removed OTSU & performing early fusion of level embeddings)**","metadata":{"_uuid":"25b1ec74-9e0c-474c-9219-1c607e7ded61","_cell_guid":"62a8057b-b78b-46e0-89de-c0e2b592e5e2","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torchvision.models as models\nimport timm\n\nclass CrossAttentionBlock(nn.Module):\n    def __init__(self, dim, num_heads=4):\n        super().__init__()\n        self.attn = nn.MultiheadAttention(\n            embed_dim=dim,\n            num_heads=num_heads,\n            batch_first=True\n        )\n        self.norm = nn.LayerNorm(dim)\n\n    def forward(self, query, key_value):\n        \"\"\"\n        query, key_value: [B, 5, F]\n        \"\"\"\n        out, _ = self.attn(query, key_value, key_value)\n        return self.norm(query + out)\n\nclass SpinalStenosisNet(nn.Module):\n    def __init__(self, feature_dim=512, num_classes=3, \n                 sag_pretrained_path=None, ax_pretrained_path=None):\n        super().__init__()\n        \n        # --- 1. Sagittal Backbone ---\n        print(\"Creating Sagittal Backbone...\")\n        self.sag_backbone = timm.create_model(\n            'efficientnetv2_rw_t.ra2_in1k', pretrained=True, in_chans=1, num_classes=0\n        )\n        if sag_pretrained_path:\n            self._load_weights(self.sag_backbone, sag_pretrained_path)\n            \n        for param in self.sag_backbone.parameters(): param.requires_grad = False\n            \n        # --- 2. Axial Backbone ---\n        print(\"Creating Axial Backbone...\")\n        self.ax_backbone = timm.create_model(\n            'efficientnetv2_rw_t.ra2_in1k', pretrained=True, in_chans=1, num_classes=0\n        )\n        if ax_pretrained_path:\n            self._load_weights(self.ax_backbone, ax_pretrained_path)\n        \n        for param in self.ax_backbone.parameters(): param.requires_grad = False\n\n        # Projections\n        bb_dim = self.sag_backbone.num_features # 1024\n        self.sag_proj = nn.Linear(bb_dim, feature_dim)\n        self.ax_proj = nn.Linear(bb_dim, feature_dim)\n\n        # RNN & Attention\n        self.level_gru = nn.GRU(feature_dim, feature_dim, batch_first=True, bidirectional=True)\n        self.ax_to_sag = CrossAttentionBlock(feature_dim * 2)\n        self.sag_to_ax = CrossAttentionBlock(feature_dim * 2)\n\n        # Heads\n        self.level_heads = nn.ModuleDict({\n            lvl: nn.Sequential(\n                nn.Linear(feature_dim * 4, 256),\n                nn.ReLU(),\n                nn.Dropout(0.4),\n                nn.Linear(256, num_classes)\n            ) for lvl in [\"L1/L2\", \"L2/L3\", \"L3/L4\", \"L4/L5\", \"L5/S1\"]\n        })\n\n    def _load_weights(self, model, path):\n        try:\n            state_dict = torch.load(path, map_location='cpu')\n            # Strip classifier keys\n            state_dict = {k: v for k, v in state_dict.items() if 'classifier' not in k}\n            model.load_state_dict(state_dict, strict=False)\n            print(f\"✅ Loaded weights from {path}\")\n        except Exception as e:\n            print(f\"⚠️ Failed load {path}: {e}\")\n\n    def forward(self, batch):\n        sag = batch[\"sagittal\"] # [B, 3, 5, 1, 384, 384]\n        ax = batch[\"axial\"]     # [B, 5, 3, 1, 384, 384]\n        B = sag.shape[0]\n\n        # --- Sagittal Features ---\n        sag_feats = []\n        for s in range(3):\n            # Flatten: [B*5, 1, 384, 384]\n            x = sag[:, s].reshape(B * 5, 1, 384, 384)\n            f = self.sag_backbone(x)\n            f = self.sag_proj(f).reshape(B, 5, -1)\n            sag_feats.append(f)\n        # Average over 3 slices\n        sag_feats = torch.stack(sag_feats).mean(dim=0) # [B, 5, F]\n        sag_ctx, _ = self.level_gru(sag_feats)         # [B, 5, 2F]\n\n        # --- Axial Features ---\n        ax_feats = []\n        for l in range(5):\n            # Flatten: [B*3, 1, 384, 384]\n            x = ax[:, l].reshape(B * 3, 1, 384, 384)\n            f = self.ax_backbone(x)\n            f = self.ax_proj(f).reshape(B, 3, -1).mean(dim=1) # Average over 3 slices\n            ax_feats.append(f)\n        \n        ax_feats = torch.stack(ax_feats, dim=1) # [B, 5, F]\n        ax_ctx, _ = self.level_gru(ax_feats)    # [B, 5, 2F]\n\n        # --- Cross Attention ---\n        sag_att = self.ax_to_sag(sag_ctx, ax_ctx)\n        ax_att = self.sag_to_ax(ax_ctx, sag_ctx)\n\n        # --- Prediction ---\n        outputs = {}\n        for i, lvl in enumerate([\"L1/L2\", \"L2/L3\", \"L3/L4\", \"L4/L5\", \"L5/S1\"]):\n            fused = torch.cat([sag_att[:, i], ax_att[:, i]], dim=1)\n            outputs[lvl] = self.level_heads[lvl](fused)\n\n        return outputs","metadata":{"_uuid":"3dc8a8d6-9ecd-4fc7-9717-1a8da36cbbd3","_cell_guid":"f44f67bf-9d95-46da-9376-29d82a063308","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:47:34.797661Z","iopub.execute_input":"2025-12-31T12:47:34.798285Z","iopub.status.idle":"2025-12-31T12:47:34.810475Z","shell.execute_reply.started":"2025-12-31T12:47:34.798264Z","shell.execute_reply":"2025-12-31T12:47:34.809790Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **TRAINING**","metadata":{"_uuid":"70735bc4-1a5f-4b1d-9d31-8a5b8cb26d14","_cell_guid":"3a5cda3d-9283-4b6b-9744-79bc94efa9c8","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"markdown","source":"**Training without OTSU & Loading into RAM**","metadata":{"_uuid":"ea8d6598-8e6f-4cb6-a324-19918dd40eb2","_cell_guid":"a275ed38-b88c-4929-83fc-16bf3fa30c9b","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torch.optim as optim\nimport copy\n\n# Configuration\nIMG_PATH = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images\"\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nLEVELS = [\"L1/L2\", \"L2/L3\", \"L3/L4\", \"L4/L5\", \"L5/S1\"]","metadata":{"_uuid":"4438aa02-92e8-4aea-bb2a-cf9b1f128734","_cell_guid":"1743271b-418f-4f81-a59a-95e365b68433","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:39:09.313963Z","iopub.execute_input":"2025-12-31T12:39:09.314485Z","iopub.status.idle":"2025-12-31T12:39:09.318629Z","shell.execute_reply.started":"2025-12-31T12:39:09.314457Z","shell.execute_reply":"2025-12-31T12:39:09.317894Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# Define Transforms with ToTensorV2\ntransforms_train = A.Compose([\n    A.RandomBrightnessContrast(brightness_limit=(-0.2, 0.2), contrast_limit=(-0.2, 0.2), p=1.0),\n    A.OneOf([\n        A.MotionBlur(blur_limit=5),\n        A.MedianBlur(blur_limit=5),\n        A.GaussianBlur(blur_limit=5),\n        A.GaussNoise(var_limit=(5.0, 30.0)),\n    ], p=0.9),\n    A.OneOf([\n        A.OpticalDistortion(distort_limit=1.0),\n        A.GridDistortion(num_steps=5, distort_limit=1.),\n        A.ElasticTransform(alpha=3),\n    ], p=0.6),\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, border_mode=0, p=1.0),\n    A.CoarseDropout(max_holes=16, max_height=64, max_width=64, min_holes=1, min_height=8, min_width=8, p=0.9),    \n    \n    # Normalization (Converts uint8 [0,255] to float [-1,1])\n    A.Normalize(mean=0.5, std=0.5),\n    \n    # Conversion to Tensor (HWC -> CHW)\n    ToTensorV2()\n])\n\ntransforms_val = A.Compose([\n    A.Normalize(mean=0.5, std=0.5),\n    ToTensorV2()\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T13:12:57.389292Z","iopub.execute_input":"2025-12-31T13:12:57.389579Z","iopub.status.idle":"2025-12-31T13:12:57.403014Z","shell.execute_reply.started":"2025-12-31T13:12:57.389528Z","shell.execute_reply":"2025-12-31T13:12:57.402289Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\ndf_train, df_val = train_test_split(df, test_size = 0.2, random_state = 11)\n\nprint(\"Preprocessing Train Data...\")\ntrain_data_list = load_data_into_memory(df_train, IMG_PATH, sagittal_model, DEVICE)\n\nprint(\"Preprocessing Val Data...\")\nval_data_list = load_data_into_memory(df_val, IMG_PATH, sagittal_model, DEVICE)","metadata":{"_uuid":"df33b64e-fcb2-4064-a302-047ab990cf14","_cell_guid":"373f373b-2ad8-4ea7-818b-3be9b45c9a0e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T12:51:56.921476Z","iopub.execute_input":"2025-12-31T12:51:56.922410Z","iopub.status.idle":"2025-12-31T13:07:56.635993Z","shell.execute_reply.started":"2025-12-31T12:51:56.922383Z","shell.execute_reply":"2025-12-31T13:07:56.634850Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize Datasets\ntrain_ds1 = SpinalStenosisDataset(train_data_list, transform = transforms_train)\ntrain_ds2 = SpinalStenosisDataset(train_data_list, transform = transforms_val)\n\ntrain_ds = ConcatDataset([train_ds1, train_ds2])\nval_ds = SpinalStenosisDataset(val_data_list, transform = transforms_val)\n\n# Initialize Dataloaders\ntrain_loader = DataLoader(train_ds, batch_size=32, shuffle=True, num_workers=0)\nval_loader = DataLoader(val_ds, batch_size=32, shuffle=False, num_workers=0)","metadata":{"_uuid":"00776dbb-5402-41ac-b681-80ae4808de28","_cell_guid":"7690162b-7677-4f25-9680-568eb56e889c","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-31T13:13:26.722232Z","iopub.execute_input":"2025-12-31T13:13:26.722503Z","iopub.status.idle":"2025-12-31T13:13:26.727378Z","shell.execute_reply.started":"2025-12-31T13:13:26.722482Z","shell.execute_reply":"2025-12-31T13:13:26.726752Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\nimport torch\n\n# 1. Get one batch\n# The loader returns a dictionary\nbatch = next(iter(train_loader))\n\n# 2. Select the first patient in the batch\npatient_idx = 0\n\n# Shapes Reminder:\n# Sagittal: [Batch, Stack(3), Level(5), Ch(1), H, W]\n# Axial:    [Batch, Level(5), Stack(3), Ch(1), H, W]\n# Labels:   [Batch, Level(5)]\n\nsag_imgs = batch['sagittal'][patient_idx] # [3, 5, 1, 384, 384]\nax_imgs  = batch['axial'][patient_idx]    # [5, 3, 1, 384, 384]\nlabels   = batch['labels'][patient_idx]   # [5]\n\n# 3. Setup Visualization\nlevels_names = [\"L1/L2\", \"L2/L3\", \"L3/L4\", \"L4/L5\", \"L5/S1\"]\nseverity_map = {0: \"Normal/Mild\", 1: \"Moderate\", 2: \"Severe\"}\n\nfig, axes = plt.subplots(2, 5, figsize=(20, 8))\nplt.subplots_adjust(hspace=0.3, wspace=0.1)\n\n# Function to un-normalize [-1, 1] -> [0, 1]\ndef unnorm(img_tensor):\n    img = img_tensor.cpu().numpy()\n    img = img * 0.5 + 0.5\n    return np.clip(img, 0, 1)\n\n# 4. Loop through the 5 Levels\nfor i in range(5):\n    lbl_text = severity_map[labels[i].item()]\n    \n    # --- Row 1: Sagittal (Middle Slice) ---\n    # Shape: [Stack, Level, Ch, H, W] -> Get Stack=1 (Middle)\n    sag_img = sag_imgs[1, i, 0] \n    axes[0, i].imshow(unnorm(sag_img), cmap='gray')\n    axes[0, i].set_title(f\"Sagittal {levels_names[i]}\\n{lbl_text}\", fontsize=10)\n    axes[0, i].axis('off')\n    \n    # --- Row 2: Axial (Middle Slice) ---\n    # Shape: [Level, Stack, Ch, H, W] -> Get Stack=1 (Middle)\n    ax_img = ax_imgs[i, 1, 0]\n    axes[1, i].imshow(unnorm(ax_img), cmap='gray')\n    axes[1, i].set_title(f\"Axial {levels_names[i]}\", fontsize=10)\n    axes[1, i].axis('off')\n\nplt.suptitle(f\"Training Sample (Patient {patient_idx}) - Middle Slices Only\", fontsize=16)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T13:13:27.602743Z","iopub.execute_input":"2025-12-31T13:13:27.603398Z","iopub.status.idle":"2025-12-31T13:13:36.766663Z","shell.execute_reply.started":"2025-12-31T13:13:27.603374Z","shell.execute_reply":"2025-12-31T13:13:36.765766Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import f1_score\nfrom transformers import get_cosine_schedule_with_warmup\nimport os\n\n# -------------------------------------------------------------------\n# CONFIGURATION (Matching Reference)\n# -------------------------------------------------------------------\nEPOCHS = 40\nLR = 2e-4\nWEIGHT_DECAY = 1e-2\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nLEVELS = [\"L1/L2\", \"L2/L3\", \"L3/L4\", \"L4/L5\", \"L5/S1\"]\n\n# Class Weights [1.0, 2.0, 4.0] for Normal, Moderate, Severe\nCE_CLASS_WEIGHTS = torch.tensor([1.0, 2.0, 4.0], dtype=torch.float32).to(DEVICE)\n\n# -------------------------------------------------------------------\n# MODEL & OPTIMIZER SETUP\n# -------------------------------------------------------------------\n# Initialize model (Ensure you provide paths if you have them, otherwise None)\n# Note: You should ideally have the 'Sagittal_PreTrain_EffNetV2.pth' from previous steps\nmodel = SpinalStenosisNet(\n    feature_dim=512,\n    num_classes=3,\n    sag_pretrained_path=\"/kaggle/input/lumbar-spine-keypoint-detection-models/Sagittal_PreTrain_EffNetV2.pth\", \n    ax_pretrained_path=\"/kaggle/input/lumbar-spine-keypoint-detection-models/Axial_PreTrain_EffNetV2.pth\" \n).to(DEVICE)\n\ncriterion = nn.CrossEntropyLoss(weight=CE_CLASS_WEIGHTS)\n\noptimizer = optim.AdamW(\n    filter(lambda p: p.requires_grad, model.parameters()),\n    lr=LR,\n    weight_decay=WEIGHT_DECAY\n)\n\nscaler = torch.cuda.amp.GradScaler()\n\n# Scheduler: Cosine with Warmup (Calculated per batch)\nnum_training_steps = EPOCHS * len(train_loader)\nnum_warmup_steps = int(0.1 * num_training_steps) # 10% Warmup\n\nscheduler = get_cosine_schedule_with_warmup(\n    optimizer,\n    num_warmup_steps=num_warmup_steps,\n    num_training_steps=num_training_steps\n)\n\n# -------------------------------------------------------------------\n# HELPER FUNCTIONS\n# -------------------------------------------------------------------\ndef move_batch_to_device(batch, device):\n    return {\n        \"sagittal\": batch[\"sagittal\"].to(device),\n        \"axial\": batch[\"axial\"].to(device),\n        \"labels\": batch[\"labels\"].to(device)\n    }\n\ndef ce_predict(logits):\n    return torch.argmax(logits, dim=1)\n\ndef validate(model, loader, device):\n    model.eval()\n\n    val_stats = {\n        lvl: {\"correct\": 0, \"total\": 0, \"y_true\": [], \"y_pred\": []}\n        for lvl in LEVELS\n    }\n\n    with torch.no_grad():\n        for batch in loader:\n            batch = move_batch_to_device(batch, device)\n            logits = model(batch)\n            labels = batch[\"labels\"]\n\n            for i, lvl in enumerate(LEVELS):\n                preds = ce_predict(logits[lvl])\n                gt = labels[:, i]\n\n                val_stats[lvl][\"correct\"] += (preds == gt).sum().item()\n                val_stats[lvl][\"total\"] += gt.size(0)\n                val_stats[lvl][\"y_true\"].append(gt.cpu().numpy())\n                val_stats[lvl][\"y_pred\"].append(preds.cpu().numpy())\n\n    overall_f1 = []\n\n    for lvl in LEVELS:\n        y_true = np.concatenate(val_stats[lvl][\"y_true\"])\n        y_pred = np.concatenate(val_stats[lvl][\"y_pred\"])\n        val_stats[lvl][\"f1_macro\"] = f1_score(\n            y_true, y_pred, average=\"macro\", zero_division=0\n        )\n        overall_f1.append(val_stats[lvl][\"f1_macro\"])\n\n    overall_acc = sum(val_stats[lvl][\"correct\"] for lvl in LEVELS) / \\\n                  sum(val_stats[lvl][\"total\"] for lvl in LEVELS)\n\n    return overall_acc, float(np.mean(overall_f1)), val_stats","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T13:19:35.511385Z","iopub.execute_input":"2025-12-31T13:19:35.512008Z","iopub.status.idle":"2025-12-31T13:19:40.457220Z","shell.execute_reply.started":"2025-12-31T13:19:35.511980Z","shell.execute_reply":"2025-12-31T13:19:40.456324Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# TRAINING LOOP\n# -------------------------------------------------------------------\ntrain_acc_steps = []\nval_metrics_history = [] # Stores (step, acc, f1)\nglobal_step = 0\n\nbest_val_f1 = 0.0\nbest_model_path = \"SpinalNet_Best_F1.pth\"\n\nprint(f\"Starting Training: {EPOCHS} Epochs, LR={LR}, Warmup={num_warmup_steps} steps\")\n\nfor epoch in range(EPOCHS):\n    model.train()\n    epoch_loss = 0\n\n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\")\n    \n    for batch in pbar:\n        batch = move_batch_to_device(batch, DEVICE)\n        labels = batch[\"labels\"]\n\n        optimizer.zero_grad(set_to_none=True)\n\n        with torch.cuda.amp.autocast():\n            logits = model(batch)\n            # Sum loss over all 5 levels\n            loss = sum(\n                criterion(logits[lvl], labels[:, i])\n                for i, lvl in enumerate(LEVELS)\n            )\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        scheduler.step() # Update LR every batch\n\n        epoch_loss += loss.item()\n\n        # Calculate Batch Accuracy for monitoring\n        correct, total = 0, 0\n        with torch.no_grad():\n            for i, lvl in enumerate(LEVELS):\n                preds = ce_predict(logits[lvl])\n                correct += (preds == labels[:, i]).sum().item()\n                total += labels.size(0)\n\n        acc = correct / total\n        train_acc_steps.append(acc)\n        global_step += 1\n\n        pbar.set_postfix(\n            loss=f\"{loss.item():.4f}\",\n            acc=f\"{acc:.4f}\",\n            lr=f\"{optimizer.param_groups[0]['lr']:.2e}\"\n        )\n\n    # ----------------- END OF EPOCH ----------------- #\n    avg_train_loss = epoch_loss / len(train_loader)\n    print(f\"\\nEpoch {epoch+1} Summary | Train Loss: {avg_train_loss:.4f}\")\n\n    # Validation\n    val_acc, val_f1, lvl_stats = validate(model, val_loader, DEVICE)\n    val_metrics_history.append((global_step, val_acc, val_f1))\n\n    print(f\"Validation Acc: {val_acc:.4f} | Macro F1: {val_f1:.4f}\")\n    \n    # Save Best Model\n    if val_f1 > best_val_f1:\n        print(f\"🔥 New Best F1! ({best_val_f1:.4f} -> {val_f1:.4f}). Saving model...\")\n        best_val_f1 = val_f1\n        torch.save(model.state_dict(), best_model_path)\n    \n    # Print per-level stats\n    print(\"-\" * 40)\n    for lvl in LEVELS:\n        print(f\"{lvl}: Acc={lvl_stats[lvl]['correct']/lvl_stats[lvl]['total']:.3f}, \"\n              f\"F1={lvl_stats[lvl]['f1_macro']:.3f}\")\n    print(\"-\" * 40)\n\n# -------------------------------------------------------------------\n# PLOTTING\n# -------------------------------------------------------------------\nsteps, v_accs, v_f1s = zip(*val_metrics_history)\ntrain_steps = range(len(train_acc_steps))\n\nplt.figure(figsize=(12, 6))\n\n# Plot Accuracy\nplt.subplot(1, 2, 1)\nplt.plot(train_steps, train_acc_steps, label=\"Train Acc\", alpha=0.3, color='blue')\nplt.plot(steps, v_accs, \"o-\", label=\"Val Acc\", color='orange', linewidth=2)\nplt.xlabel(\"Steps\")\nplt.ylabel(\"Accuracy\")\nplt.title(\"Training vs Validation Accuracy\")\nplt.legend()\nplt.grid(True)\n\n# Plot F1\nplt.subplot(1, 2, 2)\nplt.plot(steps, v_f1s, \"o-\", label=\"Val F1\", color='green', linewidth=2)\nplt.xlabel(\"Steps\")\nplt.ylabel(\"Macro F1 Score\")\nplt.title(\"Validation F1 Score\")\nplt.legend()\nplt.grid(True)\n\nplt.tight_layout()\nplt.show()\n\nprint(f\"Training Complete. Best F1: {best_val_f1:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T13:19:53.684226Z","iopub.execute_input":"2025-12-31T13:19:53.684922Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T13:24:33.902299Z","iopub.execute_input":"2025-12-31T13:24:33.903009Z","iopub.status.idle":"2025-12-31T13:24:33.955205Z","shell.execute_reply.started":"2025-12-31T13:24:33.902985Z","shell.execute_reply":"2025-12-31T13:24:33.954493Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Pre-Training feature extractor**","metadata":{"_uuid":"180e2cb1-4ecb-429a-95f4-69fca4a69ad6","_cell_guid":"ae9901e8-69f8-4b30-a8cc-aa84b0544e9a","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"markdown","source":"**Only pretraining for Axial, as Sagittal has been taken from the MSCAN paper**","metadata":{"_uuid":"5f21b226-3aca-4c9d-b9f1-3ad4e8cac2e7","_cell_guid":"03e74160-2726-4acd-846d-7b86899b7b1e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import os\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport pydicom\nfrom tqdm import tqdm\n\nLABEL_MAP = {'Normal/Mild' : 0, 'Moderate' : 1, 'Severe' : 2}\nTRAIN_IMG = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images\"\nCOLUMNS = ['study_id', 'spinal_canal_stenosis_l1_l2', 'spinal_canal_stenosis_l2_l3', 'spinal_canal_stenosis_l3_l4', 'spinal_canal_stenosis_l4_l5', 'spinal_canal_stenosis_l5_s1']","metadata":{"_uuid":"839ee612-aa6b-41d2-87b0-13ff22b0928a","_cell_guid":"edb99629-6f21-49c2-82ea-6888ebc9bd63","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv\")[COLUMNS].replace(LABEL_MAP)\ncoords_df = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv\")\ntrain_desc = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv\")","metadata":{"_uuid":"04e576d9-0b56-449d-8e00-f76d46b6637d","_cell_guid":"00d8715d-6bbd-4ec2-a3ad-23b3a1a8633b","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"coords_with_desc = coords_df.merge(train_desc, on=['study_id', 'series_id'], how='left')\ncoords = coords_with_desc[coords_with_desc['series_description'] == 'Axial T2']\nfinal = coords.merge(df, on=['study_id'], how='inner')","metadata":{"_uuid":"1e0b1c2a-12fc-4564-ad1e-38741dc077e8","_cell_guid":"4cef3150-be98-4ef2-8d23-b84cf92b8c33","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_dicom(path):\n    try:\n        dicom = pydicom.dcmread(path)\n        data = dicom.pixel_array\n        \n        # 1. Shift to 0\n        data = data - np.min(data)\n        \n        # 2. Scale to 0-1 using the Max value\n        if np.max(data) != 0:\n            data = data / np.max(data)\n        \n        # 3. Convert to uint8 (0-255)\n        # Albumentations expects this format for the best compatibility\n        data = (data * 255).astype(np.uint8)\n        return data\n    except Exception as e:\n        return np.zeros((256, 256)).astype(np.uint8)","metadata":{"_uuid":"3218e935-ecdd-404e-b5ea-579e558128a3","_cell_guid":"fcb46f2d-9800-47f5-a4da-197c86a9e53c","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_axial_data(df, base_img_path):\n    processed_samples = []\n    for _, row in tqdm(df.iterrows(), total=len(df), desc=\"Loading Axial Samples\"):\n        try:\n            if row[\"series_description\"] != \"Axial T2\":\n                continue\n            image_path = os.path.join(\n                base_img_path,\n                str(row[\"study_id\"]),\n                str(row[\"series_id\"]),\n                f\"{row['instance_number']}.dcm\"\n            )\n\n            if not os.path.exists(image_path):\n                continue\n                \n            image = load_dicom(image_path)\n\n            # This below logic for crop has been taken directly from MSCAN paper\n            x, y = row[\"x\"], row[\"y\"]\n            h, w = image.shape\n\n            pad_h = 0.09 * h\n            pad_w = 0.09 * w\n            \n            ymin = int(y - pad_h)\n            ymax = int(y + pad_h)\n            xmin = int(x - pad_w)\n            xmax = int(x + pad_w)\n\n            # Clamp boundaries to be within the image size\n            ymin = max(0, ymin)\n            ymax = min(h, ymax)\n            xmin = max(0, xmin)\n            xmax = min(w, xmax)\n\n            crop = image[ymin:ymax, xmin:xmax]\n\n            # If coordinates were way off and crop is empty, skip\n            if crop.size == 0 or crop.shape[0] == 0 or crop.shape[1] == 0:\n                continue\n\n            # Constructs column name like \"spinal_canal_stenosis_l4_l5\"\n            col = \"spinal_canal_stenosis_\" + row[\"level\"].replace(\"/\", \"_\").lower()\n            label = int(row[col])\n\n            # Store\n            processed_samples.append({\n                \"image\": crop,  \n                \"label\": label,\n            })\n\n        except Exception as e:\n            print(f\"[Axial Loader Error] {e}\") \n            continue\n            \n    return processed_samples","metadata":{"_uuid":"94fa1ab2-b716-408b-a2a0-5bf8db697795","_cell_guid":"d60237d9-dc99-48fa-9e5f-c842cd47a4a0","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_sagittal_data(df, base_img_path):\n    processed_samples = []\n\n    # Iterate through the dataframe\n    # We use tqdm to show a progress bar because loading thousands of DICOMs takes time\n    for _, row in tqdm(df.iterrows(), total=len(df), desc=\"Loading Sagittal Samples\"):\n        try:\n            # 1. Filter: Ensure we are only looking at Sagittal T2/STIR images\n            if row[\"series_description\"] != \"Sagittal T2/STIR\":\n                continue\n\n            # 2. Path Construction\n            image_path = os.path.join(\n                base_img_path,\n                str(row[\"study_id\"]),\n                str(row[\"series_id\"]),\n                f\"{row['instance_number']}.dcm\"\n            )\n\n            # 3. Load Image\n            if not os.path.exists(image_path):\n                # Only strictly necessary if your dataset might have missing files\n                continue\n                \n            image = load_dicom(image_path)\n\n            # 4. Crop Logic (MATCHING REFERENCE CODE EXACTLY)\n            # Reference: int(y - 0.09 * h) to int(y + 0.09 * h)\n            x, y = row[\"x\"], row[\"y\"]\n            h, w = image.shape\n\n            # Calculate boundaries\n            pad_h = 0.09 * h\n            pad_w = 0.09 * w\n            \n            ymin = int(y - pad_h)\n            ymax = int(y + pad_h)\n            xmin = int(x - pad_w)\n            xmax = int(x + pad_w)\n\n            # Clamp boundaries to be within the image size\n            ymin = max(0, ymin)\n            ymax = min(h, ymax)\n            xmin = max(0, xmin)\n            xmax = min(w, xmax)\n\n            # Perform the crop\n            crop = image[ymin:ymax, xmin:xmax]\n\n            # 5. Sanity Check\n            # If coordinates were way off and crop is empty, skip\n            if crop.size == 0 or crop.shape[0] == 0 or crop.shape[1] == 0:\n                continue\n\n            # 6. Label Extraction\n            # Constructs column name like \"spinal_canal_stenosis_l4_l5\"\n            col = \"spinal_canal_stenosis_\" + row[\"level\"].replace(\"/\", \"_\").lower()\n            label = int(row[col])\n\n            # 7. Store\n            processed_samples.append({\n                \"image\": crop,  # We store the crop (numpy array)\n                \"label\": label,\n                \"study_id\": row[\"study_id\"], # Optional: keep metadata for debugging\n                \"level\": row[\"level\"]\n            })\n\n        except Exception as e:\n            # If a specific image fails, print error but don't stop the whole loop\n            # print(f\"[Sagittal Loader Error] {e}\") \n            continue\n            \n    return processed_samples","metadata":{"_uuid":"2351dc71-11f9-4b35-bd64-82cfa80f0665","_cell_guid":"2594527f-c3e1-4d25-b69d-deb100fb5b26","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import albumentations as A\nimport torchvision.transforms as transforms\n\ntransforms_train = A.Compose([\n    A.RandomBrightnessContrast(brightness_limit=(-0.2, 0.2), contrast_limit=(-0.2, 0.2), p=1.0),\n    A.OneOf([\n        A.MotionBlur(blur_limit=5),\n        A.MedianBlur(blur_limit=5),\n        A.GaussianBlur(blur_limit=5),\n        A.GaussNoise(var_limit=(5.0, 30.0)),\n    ], p=1.0),\n\n    A.CoarseDropout(\n        max_holes=4,  # int | None\n        max_height=8,  # ScalarType | None\n        max_width=8,  # ScalarType | None\n        min_holes=2,  # int | None\n        min_height=None,  # ScalarType | None\n        min_width=None,  # ScalarType | None\n        fill_value=0,  # Union[float, Sequence[float]]\n        mask_fill_value=None,  # Union[float, Sequence[float], NoneType]\n        num_holes_range=(1, 1),  # tuple[int, int]\n        hole_height_range=(8, 8),  # tuple[ScalarType, ScalarType]\n        hole_width_range=(8, 8),  # tuple[ScalarType, ScalarType]\n        always_apply=None,  # bool | None\n        p=0.7,  # float\n    ),\n    A.Defocus(\n    radius=(3, 5),  # ScaleIntType\n    alias_blur=(0.1, 0.5),  # ScaleFloatType\n    always_apply=None,  # bool | None\n    p=0.4,  # float\n    ),\n\n\n    A.Perspective(\n    scale=(0.05, 0.1),  # ScaleFloatType\n    keep_size=True,  # bool\n    pad_mode=0,  # int\n    pad_val=0,  # ColorType\n    mask_pad_val=0,  # ColorType\n    fit_output=False,  # bool\n    interpolation=1,  # <class 'int'>\n    always_apply=None,  # bool | None\n    p=0.2,  # float\n    ),\n    A.PixelDropout(\n    dropout_prob=0.01,  # float\n    per_channel=False,  # bool\n    drop_value=0,  # ScaleFloatType | None\n    mask_drop_value=None,  # ScaleFloatType | None\n    always_apply=None,  # bool | None\n    p=0.5,  # float\n    ),\n\n    A.RandomToneCurve(\n    scale=0.1,  # float\n    per_channel=False,  # bool\n    always_apply=None,  # bool | None\n    p=0.4,  # float\n    ),\n    A.RandomGamma(\n    gamma_limit=(80, 120),  # ScaleIntType\n    always_apply=None,  # bool | None\n    p=0.5,  # float\n    ),\n    A.ShiftScaleRotate(\n    shift_limit=(-0.0625, 0.0625),  # ScaleFloatType\n    scale_limit=(-0.1, 0.1),  # ScaleFloatType\n    rotate_limit=(-15, 15),  # ScaleFloatType\n    interpolation=1,  # <class 'int'>\n    border_mode=4,  # int\n    value=0,  # ColorType\n    mask_value=0,  # ColorType\n    shift_limit_x=None,  # ScaleFloatType | None\n    shift_limit_y=None,  # ScaleFloatType | None\n    rotate_method=\"largest_box\",  # Literal['largest_box', 'ellipse']\n    always_apply=None,  # bool | None\n    p=1.0,  # float\n    ),\n\n\n    #grey scale 3 channel\n    A.CLAHE(clip_limit=4.0, tile_grid_size=(8, 8), always_apply=True),\n    A.Normalize(mean=0.5, std=0.5),\n    A.Resize(384, 384),\n\n    \n])\n\ntransforms_val = A.Compose([\n    A.CLAHE(clip_limit=4.0, tile_grid_size=(8, 8), always_apply=True),\n    A.Normalize(mean=0.5, std=0.5),\n    A.Resize(384, 384)\n])","metadata":{"_uuid":"698c6286-70ca-4cb4-8420-ba3fa644fe80","_cell_guid":"afedcc32-6d42-4a3f-ae05-6ea7de973653","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PretrainingDataset(Dataset):\n    def __init__(self, data_list, transform):\n        self.data = data_list\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        item = self.data[idx]\n        img = item[\"image\"]          # numpy array, uint8, H x W\n        label = item[\"label\"]\n        \n        # FIX: Keep as uint8. \n        # Albumentations expects uint8 inputs for proper noise/blur application.\n        # A.Normalize will handle the float conversion later.\n        \n        # Albumentations expects H x W or H x W x C\n        if img.ndim == 2:\n            img = img[..., None]     # H x W x 1\n\n        if self.transform is not None:\n            augmented = self.transform(image=img)\n            img = augmented[\"image\"]  # numpy array, float32, H x W x C (normalized)\n\n        # Convert to torch tensor, Channel-First (C, H, W)\n        img = torch.from_numpy(img).permute(2, 0, 1).float()\n\n        return img, torch.tensor(label, dtype=torch.long)","metadata":{"_uuid":"6076688e-12ce-4735-bdb8-24c95bc273f5","_cell_guid":"8408d6e1-f6a0-4975-8c36-b68d1d0291fe","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader, ConcatDataset\nfrom sklearn.model_selection import train_test_split\n\nsamples = load_axial_data(\n    df=final,\n    base_img_path=TRAIN_IMG\n)","metadata":{"_uuid":"a6ce2fd1-4f9c-465a-ad3d-0efd143b8f3b","_cell_guid":"8446b8dc-3605-4407-ad9e-d2b3d5314768","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Total samples loaded: {len(samples)}\")\n\nlabels = [s[\"label\"] for s in samples]\n\ntrain_samples, val_samples = train_test_split(\n    samples,\n    test_size=0.2,    # 10% for validation (matches reference ~0.1 split)\n    stratify=labels,\n    random_state=42\n)\n\nprint(f\"Training samples (Before concat): {len(train_samples)}\")\nprint(f\"Validation samples: {len(val_samples)}\")\n\n# (Uses transforms_train: Noise, Dropout, Distortions)\ntrain_ds_aug = PretrainingDataset(\n    data_list=train_samples,\n    transform=transforms_train \n)\n# (Uses transforms_val: Only Resize & Normalize)\n# This prevents the model from forgetting what a real spine looks like.\ntrain_ds_clean = PretrainingDataset(\n    data_list=train_samples,\n    transform=transforms_val \n)\n# This effectively doubles your epoch size (50% clean, 50% augmented)\ntrain_ds_final = ConcatDataset([train_ds_aug, train_ds_clean])\n\nval_ds = PretrainingDataset(\n    data_list=val_samples,\n    transform=transforms_val\n)\n\n\nbatch_size = 32\ntrain_dl = DataLoader(\n    train_ds_final,\n    batch_size=batch_size,\n    shuffle=True,       # CRITICAL: Mixes clean and dirty images in every batch\n    num_workers=4,      # Adjust based on CPU cores\n    pin_memory=True,    # Faster transfer to GPU\n    drop_last=True      # Drops incomplete batch at the end\n)\n\nval_dl = DataLoader(\n    val_ds,\n    batch_size=batch_size,\n    shuffle=False,      # No need to shuffle validation\n    num_workers=4,\n    pin_memory=True\n)\n\nprint(f\"Final Training DataLoader length: {len(train_dl)} batches\")\nprint(f\"Final Validation DataLoader length: {len(val_dl)} batches\")","metadata":{"_uuid":"96dfe6d4-7bd2-48d5-aa83-b4a6ce8bbc2b","_cell_guid":"b866bd86-c132-4e53-990e-34b17824e33e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\n# Reverse label map for display\nINV_SEVERITY_MAP = {\n    0: \"Normal/Mild\",\n    1: \"Moderate\",\n    2: \"Severe\"\n}\n\n# 1. Get one batch\nimages, labels = next(iter(train_dl))\n\n# 2. Move to CPU\nimages = images.cpu()\nlabels = labels.cpu()\n\n# 3. Setup Plot for 16 images\n# Ensure we don't crash if the batch size is smaller than 16\nN = min(16, images.size(0)) \n\n# Create a 4x4 grid (or smaller if N < 16)\nrows = 4\ncols = 4\nplt.figure(figsize=(16, 16))\n\nfor i in range(N):\n    plt.subplot(rows, cols, i + 1)\n    \n    # Extract image: [1, H, W] -> [H, W]\n    img = images[i, 0].numpy()\n    \n    # --- UN-NORMALIZE ---\n    # Reverses A.Normalize(mean=0.5, std=0.5) => [-1, 1] to [0, 1]\n    img = img * 0.5 + 0.5\n    img = np.clip(img, 0, 1)\n    \n    plt.imshow(img, cmap=\"gray\")\n    plt.title(f\"{INV_SEVERITY_MAP[int(labels[i])]}\")\n    plt.axis(\"off\")\n\nplt.suptitle(f\"Sagittal ROI Batch (Showing {N} images)\\nLook for the mix of 'Clean' vs 'Augmented'\", fontsize=16)\nplt.tight_layout()\nplt.show()","metadata":{"_uuid":"41a91d14-3c22-473d-83bf-ff65d178b03b","_cell_guid":"ca4f7426-def9-4076-b86f-7b9e328e47a9","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport timm\nfrom torch.optim import AdamW\nfrom tqdm import tqdm\nimport sys\n\n# --------------------------------------------------\n# Model\n# --------------------------------------------------\nmodel_name = 'efficientnetv2_rw_t.ra2_in1k'\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nprint(f\"Creating model: {model_name}\")\nmodel = timm.create_model(\n    model_name,\n    pretrained=True,\n    in_chans=1,\n    num_classes=3\n).to(device)\n\n# --------------------------------------------------\n# Loss (Weighted for Class Imbalance)\n# --------------------------------------------------\n# Reference weights: 1.0 (Normal), 2.0 (Moderate), 4.0 (Severe)\nweights = torch.tensor([1.0, 2.0, 4.0], device=device)\ncriterion = nn.CrossEntropyLoss(weight=weights)\n\n# --------------------------------------------------\n# Optimizer\n# --------------------------------------------------\n# Reference LR: 0.00005\noptimizer = AdamW(model.parameters(), lr=0.00005)\n\n# --------------------------------------------------\n# Scheduler\n# --------------------------------------------------\n# Reference: ReduceLROnPlateau, mode='max' (Maximize Accuracy), factor=0.1, patience=2\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode='max',      # Changed to MAX because we are tracking Accuracy\n    factor=0.1,\n    patience=2,\n    verbose=True\n)\n\n# --------------------------------------------------\n# Training config\n# --------------------------------------------------\nepochs = 20  # Reference used 100\nbest_acc = 0.0 # Reference tracks Best Accuracy, not Loss\nearly_stopping_patience = 5\nearly_stopping_counter = early_stopping_patience\nsave_path = \"Axial_PreTrain_EffNetV2.pth\"\n\n# --------------------------------------------------\n# Training Loop\n# --------------------------------------------------\nprint(f\"Starting training on {device}...\")\n\nfor epoch in range(epochs):\n    print(f\"\\nEpoch {epoch + 1}/{epochs}\")\n    print(\"-\" * 30)\n\n    for phase in [\"train\", \"val\"]:\n        is_train = phase == \"train\"\n        model.train() if is_train else model.eval()\n\n        dataloader = train_dl if is_train else val_dl\n        dataset_size = len(dataloader.dataset)\n\n        running_loss = 0.0\n        running_corrects = 0\n        \n        # Determine dataset length for the progress bar\n        # (len(dataloader) gives number of batches)\n        pbar = tqdm(dataloader, desc=phase.upper(), leave=True)\n\n        with torch.set_grad_enabled(is_train):\n            for inputs, labels in pbar:\n                inputs = inputs.to(device, non_blocking=True)\n                labels = labels.to(device, non_blocking=True)\n\n                optimizer.zero_grad(set_to_none=True)\n\n                outputs = model(inputs)\n                _, preds = torch.max(outputs, 1)\n                loss = criterion(outputs, labels)\n\n                if is_train:\n                    loss.backward()\n                    optimizer.step()\n\n                # Statistics\n                batch_size = inputs.size(0)\n                running_loss += loss.item() * batch_size\n                running_corrects += torch.sum(preds == labels.data)\n                \n                # Update progress bar with current batch loss\n                pbar.set_postfix({\"loss\": f\"{loss.item():.4f}\"})\n\n        # Epoch Statistics\n        epoch_loss = running_loss / dataset_size\n        epoch_acc = running_corrects.double() / dataset_size\n\n        print(f\"{phase.capitalize()} Loss: {epoch_loss:.4f} | Acc: {epoch_acc:.4f}\")\n\n        # --------------------------------------------------\n        # Scheduler & Checkpointing\n        # --------------------------------------------------\n        \n        # Step Scheduler on Validation Accuracy (Maximize)\n        if phase == \"val\":\n            scheduler.step(epoch_acc)\n\n            # Save if Accuracy Improves\n            if epoch_acc > best_acc:\n                best_acc = epoch_acc\n                early_stopping_counter = early_stopping_patience\n                torch.save(model.state_dict(), save_path)\n                print(f\"✅ Best model saved! (Acc: {best_acc:.4f})\")\n            else:\n                early_stopping_counter -= 1\n                print(f\"⏳ Early stopping counter: {early_stopping_counter}/{early_stopping_patience}\")\n                \n    # Early Stopping Trigger (checked after val phase)\n    if early_stopping_counter == 0:\n        print(\"\\n🛑 Early stopping triggered. Training finished.\")\n        break","metadata":{"_uuid":"1c4e01f5-ed93-4da9-a71d-55c3e439b4bb","_cell_guid":"2c8bdad4-fbb6-46ba-9eb4-7094755c453b","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"_uuid":"77198957-8084-4951-9e07-ee0f52b6c9b7","_cell_guid":"f8c9acdd-705b-4a05-8a4f-63b9627faaac","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import f1_score, confusion_matrix, ConfusionMatrixDisplay\nimport matplotlib.pyplot as plt\nimport torch\nimport numpy as np\nfrom tqdm import tqdm\n\n# 1. Setup\nmodel.eval()\ny_pred = []\ny_true = []\n\n# 2. Inference Loop\nprint(\"Running inference on validation set...\")\nwith torch.no_grad():\n    for inputs, labels in tqdm(val_dl, desc=\"Validating\"):\n        inputs = inputs.to(device)\n        \n        # Forward pass\n        outputs = model(inputs)\n        preds = torch.argmax(outputs, dim=1)\n        \n        # Store results\n        y_pred.extend(preds.cpu().numpy())\n        y_true.extend(labels.cpu().numpy())\n\n# 3. Calculate Macro F1\nmacro_f1 = f1_score(y_true, y_pred, average='macro')\nprint(f\"\\n🔥 Validation Macro F1 Score: {macro_f1:.4f}\")\n\n# 4. Generate & Plot Confusion Matrix\ncm = confusion_matrix(y_true, y_pred)\ntarget_names = ['Normal/Mild', 'Moderate', 'Severe']\n\n# Create the plot\nfig, ax = plt.subplots(figsize=(8, 8))\ndisp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=target_names)\ndisp.plot(cmap='Blues', ax=ax, values_format='d')\n\nplt.title(f'Confusion Matrix (Macro F1: {macro_f1:.4f})', fontsize=14)\nplt.show()\n\n# 5. Print Per-Class Accuracy (Optional but helpful)\n# Normalize CM by row (true label) to see recall per class\ncm_normalized = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]\nprint(\"\\nPer-Class Recall (Accuracy):\")\nfor i, name in enumerate(target_names):\n    print(f\"{name}: {cm_normalized[i, i]*100:.2f}%\")","metadata":{"_uuid":"3bd2cbaa-5891-4786-8320-b42653f5e75f","_cell_guid":"1f71f5fa-cb36-4e67-bb7b-5b5125fcc172","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"4061a00e-731f-432d-8155-90ee39c8a5c5","_cell_guid":"da5f4bdc-7de7-45d9-85d1-94b864aa4b90","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}