{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","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,"isSourceIdPinned":false}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ============================================================\n# IMPORTS\n# ============================================================\n\nimport os\nimport gc\nimport cv2\nimport glob\nimport random\nimport warnings\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nimport pydicom\n\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom torch.utils.data import (\n    Dataset,\n    DataLoader,\n    WeightedRandomSampler\n)\n\nfrom torchvision import transforms, models\n\nfrom sklearn.model_selection import train_test_split\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    precision_score,\n    recall_score,\n    f1_score,\n    roc_auc_score,\n    confusion_matrix\n)\n\n# ============================================================\n# WARNING CONTROL\n# ============================================================\n\nwarnings.filterwarnings(\"ignore\")\n\n# ============================================================\n# REPRODUCIBILITY\n# ============================================================\n\nSEED = 42\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\n\ntorch.manual_seed(SEED)\n\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n\n# ============================================================\n# DEVICE CONFIGURATION\n# ============================================================\n\ndevice = torch.device(\n    \"cuda\" if torch.cuda.is_available() else \"cpu\"\n)\n\nprint(\"=\"*60)\nprint(\"DEVICE CONFIGURATION\")\nprint(\"=\"*60)\n\nprint(f\"Using Device : {device}\")\n\nif torch.cuda.is_available():\n\n    print(f\"GPU Name : {torch.cuda.get_device_name(0)}\")\n\n    print(\n        f\"GPU Memory : \"\n        f\"{torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB\"\n    )\n\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:10:54.955907Z","iopub.execute_input":"2026-05-14T17:10:54.956572Z","iopub.status.idle":"2026-05-14T17:11:06.150350Z","shell.execute_reply.started":"2026-05-14T17:10:54.956541Z","shell.execute_reply":"2026-05-14T17:11:06.149525Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# DATASET PATHS\n# ============================================================\n\nROOT_PATH = \"/kaggle/input/competitions/rsna-2024-lumbar-spine-degenerative-classification\"\n\nTRAIN_CSV = os.path.join(ROOT_PATH, \"train.csv\")\n\nTRAIN_DESC_CSV = os.path.join(\n    ROOT_PATH,\n    \"train_series_descriptions.csv\"\n)\n\nTRAIN_IMAGES_DIR = os.path.join(\n    ROOT_PATH,\n    \"train_images\"\n)\n\n# ============================================================\n# FILE EXISTENCE CHECK\n# ============================================================\n\nrequired_files = [\n    TRAIN_CSV,\n    TRAIN_DESC_CSV,\n    TRAIN_IMAGES_DIR\n]\n\nprint(\"=\"*60)\nprint(\"DATASET VALIDATION\")\nprint(\"=\"*60)\n\nfor file_path in required_files:\n\n    if os.path.exists(file_path):\n\n        print(f\"[FOUND] {file_path}\")\n\n    else:\n\n        raise FileNotFoundError(\n            f\"Missing File : {file_path}\"\n        )\n\nprint(\"=\"*60)\n\n# ============================================================\n# LOAD CSV FILES\n# ============================================================\n\ntry:\n\n    train_df = pd.read_csv(TRAIN_CSV)\n\n    series_df = pd.read_csv(TRAIN_DESC_CSV)\n\n    print(\"\\nCSV FILES LOADED SUCCESSFULLY\")\n\nexcept Exception as e:\n\n    raise RuntimeError(\n        f\"CSV Loading Failed : {e}\"\n    )\n\n# ============================================================\n# BASIC DATASET INFO\n# ============================================================\n\nprint(\"\\nTRAIN DATA SHAPE :\", train_df.shape)\n\nprint(\"SERIES DATA SHAPE :\", series_df.shape)\n\nprint(\"\\nTRAIN COLUMNS\\n\")\n\nprint(train_df.columns.tolist())","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:11:20.841215Z","iopub.execute_input":"2026-05-14T17:11:20.842225Z","iopub.status.idle":"2026-05-14T17:11:20.898314Z","shell.execute_reply.started":"2026-05-14T17:11:20.842192Z","shell.execute_reply":"2026-05-14T17:11:20.897557Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# ALL LUMBAR CONDITIONS\n# ============================================================\n\nall_conditions = [\n\n    # Spinal Canal Stenosis\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    # Left Neural Foraminal Narrowing\n    \"left_neural_foraminal_narrowing_l1_l2\",\n    \"left_neural_foraminal_narrowing_l2_l3\",\n    \"left_neural_foraminal_narrowing_l3_l4\",\n    \"left_neural_foraminal_narrowing_l4_l5\",\n    \"left_neural_foraminal_narrowing_l5_s1\",\n\n    # Right Neural Foraminal Narrowing\n    \"right_neural_foraminal_narrowing_l1_l2\",\n    \"right_neural_foraminal_narrowing_l2_l3\",\n    \"right_neural_foraminal_narrowing_l3_l4\",\n    \"right_neural_foraminal_narrowing_l4_l5\",\n    \"right_neural_foraminal_narrowing_l5_s1\",\n\n    # Left Subarticular Stenosis\n    \"left_subarticular_stenosis_l1_l2\",\n    \"left_subarticular_stenosis_l2_l3\",\n    \"left_subarticular_stenosis_l3_l4\",\n    \"left_subarticular_stenosis_l4_l5\",\n    \"left_subarticular_stenosis_l5_s1\",\n\n    # Right Subarticular Stenosis\n    \"right_subarticular_stenosis_l1_l2\",\n    \"right_subarticular_stenosis_l2_l3\",\n    \"right_subarticular_stenosis_l3_l4\",\n    \"right_subarticular_stenosis_l4_l5\",\n    \"right_subarticular_stenosis_l5_s1\"\n]\n\n# ============================================================\n# LABEL CREATION\n# ============================================================\n\ndef create_label(row):\n\n    try:\n\n        for col in all_conditions:\n\n            value = row[col]\n\n            if value in [\"Moderate\", \"Severe\"]:\n\n                return 1\n\n        return 0\n\n    except Exception:\n\n        return 0\n\n# ============================================================\n# GENERATE LABELS\n# ============================================================\n\ndf = train_df[[\"study_id\"] + all_conditions].copy()\n\ndf[\"label\"] = df.apply(\n    create_label,\n    axis=1\n)\n\ndf = df[[\"study_id\", \"label\"]]\n\ndf = df.reset_index(drop=True)\n\n# ============================================================\n# LABEL DISTRIBUTION\n# ============================================================\n\nprint(\"=\"*60)\nprint(\"LABEL DISTRIBUTION\")\nprint(\"=\"*60)\n\nprint(df[\"label\"].value_counts())\n\nprint(\"\\nPERCENTAGE DISTRIBUTION\\n\")\n\nprint(\n    df[\"label\"].value_counts(normalize=True) * 100\n)\n\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:11:39.301516Z","iopub.execute_input":"2026-05-14T17:11:39.302567Z","iopub.status.idle":"2026-05-14T17:11:39.378356Z","shell.execute_reply.started":"2026-05-14T17:11:39.302529Z","shell.execute_reply":"2026-05-14T17:11:39.377322Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# FILTER SAGITTAL T2/STIR\n# ============================================================\n\nsagittal_df = series_df[\n    series_df[\"series_description\"] == \"Sagittal T2/STIR\"\n].copy()\n\nsagittal_df = sagittal_df.reset_index(drop=True)\n\nprint(\"=\"*60)\nprint(\"SAGITTAL MRI SERIES\")\nprint(\"=\"*60)\n\nprint(\"Total Sagittal Studies :\", len(sagittal_df))\n\n# ============================================================\n# MERGE LABELS + SERIES INFO\n# ============================================================\n\nmerged_df = pd.merge(\n    df,\n    sagittal_df,\n    on=\"study_id\"\n)\n\nmerged_df = merged_df.reset_index(drop=True)\n\nprint(\"\\nMERGED DATASET SIZE :\", len(merged_df))\n\n# ============================================================\n# SAFE DICOM PATH EXTRACTION\n# ============================================================\n\ndef get_dicom_paths(row):\n\n    try:\n\n        study_id = row[\"study_id\"]\n\n        series_id = row[\"series_id\"]\n\n        folder_path = os.path.join(\n            TRAIN_IMAGES_DIR,\n            str(study_id),\n            str(series_id)\n        )\n\n        dicom_files = sorted(\n            glob.glob(\n                os.path.join(folder_path, \"*.dcm\")\n            )\n        )\n\n        if len(dicom_files) == 0:\n\n            return None\n\n        return dicom_files\n\n    except Exception:\n\n        return None\n\n# ============================================================\n# CREATE DICOM PATHS\n# ============================================================\n\nmerged_df[\"dicom_paths\"] = merged_df.apply(\n    get_dicom_paths,\n    axis=1\n)\n\n# ============================================================\n# REMOVE INVALID STUDIES\n# ============================================================\n\nmerged_df = merged_df[\n    merged_df[\"dicom_paths\"].notnull()\n]\n\nmerged_df = merged_df.reset_index(drop=True)\n\nprint(\"\\nVALID MRI STUDIES :\", len(merged_df))\n\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:11:56.265950Z","iopub.execute_input":"2026-05-14T17:11:56.266644Z","iopub.status.idle":"2026-05-14T17:12:36.943580Z","shell.execute_reply.started":"2026-05-14T17:11:56.266612Z","shell.execute_reply":"2026-05-14T17:12:36.942708Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# TRAIN VALIDATION SPLIT\n# ============================================================\n\ntrain_data, val_data = train_test_split(\n\n    merged_df,\n\n    test_size=0.2,\n\n    stratify=merged_df[\"label\"],\n\n    random_state=SEED\n)\n\ntrain_data = train_data.reset_index(drop=True)\n\nval_data = val_data.reset_index(drop=True)\n\nprint(\"=\"*60)\nprint(\"TRAIN VALIDATION SPLIT\")\nprint(\"=\"*60)\n\nprint(f\"Train Samples      : {len(train_data)}\")\nprint(f\"Validation Samples : {len(val_data)}\")\n\n# ============================================================\n# CLASS DISTRIBUTION\n# ============================================================\n\nprint(\"\\nTRAIN LABEL DISTRIBUTION\\n\")\n\nprint(train_data[\"label\"].value_counts())\n\nprint(\"\\nVALID LABEL DISTRIBUTION\\n\")\n\nprint(val_data[\"label\"].value_counts())\n\n# ============================================================\n# WEIGHTED RANDOM SAMPLER\n# ============================================================\n\ntrain_labels = train_data[\"label\"].values\n\nclass_sample_count = np.array([\n    len(np.where(train_labels == t)[0])\n    for t in np.unique(train_labels)\n])\n\nweights = 1. / class_sample_count\n\nsamples_weight = np.array([\n    weights[t]\n    for t in train_labels\n])\n\nsamples_weight = torch.from_numpy(\n    samples_weight\n).double()\n\nsampler = WeightedRandomSampler(\n    samples_weight,\n    len(samples_weight),\n    replacement=True\n)\n\nprint(\"\\nWeighted Sampler Created Successfully\")\n\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:12:56.430717Z","iopub.execute_input":"2026-05-14T17:12:56.431198Z","iopub.status.idle":"2026-05-14T17:12:56.460391Z","shell.execute_reply.started":"2026-05-14T17:12:56.431168Z","shell.execute_reply":"2026-05-14T17:12:56.459625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# MRI IMAGE TRANSFORMS\n# ============================================================\n\ntrain_transform = transforms.Compose([\n\n    transforms.ToPILImage(),\n\n    transforms.Resize((224, 224)),\n\n    transforms.RandomHorizontalFlip(p=0.5),\n\n    transforms.RandomRotation(5),\n\n    transforms.ColorJitter(\n        brightness=0.1,\n        contrast=0.1\n    ),\n\n    transforms.ToTensor(),\n\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\nval_transform = transforms.Compose([\n\n    transforms.ToPILImage(),\n\n    transforms.Resize((224, 224)),\n\n    transforms.ToTensor(),\n\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\nprint(\"=\"*60)\nprint(\"TRANSFORM PIPELINE READY\")\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:13:07.275708Z","iopub.execute_input":"2026-05-14T17:13:07.276091Z","iopub.status.idle":"2026-05-14T17:13:07.282789Z","shell.execute_reply.started":"2026-05-14T17:13:07.276063Z","shell.execute_reply":"2026-05-14T17:13:07.282014Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CUSTOM MRI DATASET\n# ============================================================\n\nclass LumbarDataset(Dataset):\n\n    def __init__(self, dataframe, transform=None):\n\n        self.dataframe = dataframe.reset_index(drop=True)\n\n        self.transform = transform\n\n    def __len__(self):\n\n        return len(self.dataframe)\n\n    # ========================================================\n    # SAFE DICOM LOADER\n    # ========================================================\n\n    def load_dicom(self, path):\n\n        try:\n\n            dicom = pydicom.dcmread(path)\n\n            image = dicom.pixel_array.astype(np.float32)\n\n            image = cv2.normalize(\n                image,\n                None,\n                0,\n                255,\n                cv2.NORM_MINMAX\n            )\n\n            image = image.astype(np.uint8)\n\n            return image\n\n        except Exception:\n\n            return np.zeros((512, 512), dtype=np.uint8)\n\n    # ========================================================\n    # GET ITEM\n    # ========================================================\n\n    def __getitem__(self, idx):\n\n        try:\n\n            row = self.dataframe.iloc[idx]\n\n            dicom_paths = row[\"dicom_paths\"]\n\n            label = row[\"label\"]\n\n            # =================================================\n            # MULTI-SLICE EXTRACTION\n            # =================================================\n\n            middle_idx = len(dicom_paths) // 2\n\n            prev_idx = max(middle_idx - 1, 0)\n\n            next_idx = min(\n                middle_idx + 1,\n                len(dicom_paths) - 1\n            )\n\n            slice_1 = self.load_dicom(\n                dicom_paths[prev_idx]\n            )\n\n            slice_2 = self.load_dicom(\n                dicom_paths[middle_idx]\n            )\n\n            slice_3 = self.load_dicom(\n                dicom_paths[next_idx]\n            )\n\n            # =================================================\n            # CREATE 3-CHANNEL MRI IMAGE\n            # =================================================\n\n            image = np.stack(\n                [slice_1, slice_2, slice_3],\n                axis=-1\n            )\n\n            # =================================================\n            # APPLY TRANSFORMS\n            # =================================================\n\n            if self.transform:\n\n                image = self.transform(image)\n\n            return image, torch.tensor(label).long()\n\n        except Exception as e:\n\n            print(f\"Dataset Error : {e}\")\n\n            dummy_image = torch.zeros((3, 224, 224))\n\n            dummy_label = torch.tensor(0).long()\n\n            return dummy_image, dummy_label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:13:16.485449Z","iopub.execute_input":"2026-05-14T17:13:16.486113Z","iopub.status.idle":"2026-05-14T17:13:16.495069Z","shell.execute_reply.started":"2026-05-14T17:13:16.486086Z","shell.execute_reply":"2026-05-14T17:13:16.494116Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# DATASET OBJECTS\n# ============================================================\n\ntrain_dataset = LumbarDataset(\n    train_data,\n    transform=train_transform\n)\n\nval_dataset = LumbarDataset(\n    val_data,\n    transform=val_transform\n)\n\nprint(\"=\"*60)\nprint(\"DATASETS CREATED\")\nprint(\"=\"*60)\n\nprint(f\"Train Dataset Size : {len(train_dataset)}\")\n\nprint(f\"Validation Dataset Size : {len(val_dataset)}\")\n\n# ============================================================\n# DATALOADERS\n# ============================================================\n\nBATCH_SIZE = 8\n\ntrain_loader = DataLoader(\n\n    train_dataset,\n\n    batch_size=BATCH_SIZE,\n\n    sampler=sampler,\n\n    num_workers=2,\n\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n\n    val_dataset,\n\n    batch_size=BATCH_SIZE,\n\n    shuffle=False,\n\n    num_workers=2,\n\n    pin_memory=True\n)\n\nprint(\"\\nDATALOADERS READY\")\n\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:13:28.116003Z","iopub.execute_input":"2026-05-14T17:13:28.116972Z","iopub.status.idle":"2026-05-14T17:13:28.124093Z","shell.execute_reply.started":"2026-05-14T17:13:28.116946Z","shell.execute_reply":"2026-05-14T17:13:28.123184Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# VISUALIZE MRI SAMPLES\n# ============================================================\n\nimages, labels = next(iter(train_loader))\n\nplt.figure(figsize=(15, 8))\n\nfor i in range(6):\n\n    image = images[i].permute(1, 2, 0).cpu().numpy()\n\n    # ========================================================\n    # DE-NORMALIZE IMAGE\n    # ========================================================\n\n    image = image * np.array(\n        [0.229, 0.224, 0.225]\n    ) + np.array(\n        [0.485, 0.456, 0.406]\n    )\n\n    image = np.clip(image, 0, 1)\n\n    plt.subplot(2, 3, i + 1)\n\n    plt.imshow(image, cmap=\"gray\")\n\n    plt.title(\n        f\"Label : {labels[i].item()}\"\n    )\n\n    plt.axis(\"off\")\n\nplt.tight_layout()\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:14:08.981546Z","iopub.execute_input":"2026-05-14T17:14:08.982493Z","iopub.status.idle":"2026-05-14T17:14:11.375410Z","shell.execute_reply.started":"2026-05-14T17:14:08.982454Z","shell.execute_reply":"2026-05-14T17:14:11.374274Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# LOAD PRETRAINED RESNET50\n# ============================================================\n\nmodel = models.densenet121(\n    weights=models.DenseNet121_Weights.DEFAULT\n)\n\nprint(\"=\"*60)\nprint(\"PRETRAINED DENSENET121 LOADED\")\nprint(\"=\"*60)\n\n# ============================================================\n# FREEZE EARLY LAYERS\n# ============================================================\n\nfor param in model.parameters():\n\n    param.requires_grad = False\n\n# ============================================================\n# UNFREEZE HIGH-LEVEL FEATURES\n# ============================================================\n\nfor param in model.features.denseblock4.parameters():\n\n    param.requires_grad = True\n\n# ============================================================\n# MODIFY FINAL CLASSIFIER\n# ============================================================\n\nnum_features = model.classifier.in_features\n\nmodel.classifier = nn.Sequential(\n\n    nn.Linear(num_features, 512),\n\n    nn.ReLU(),\n\n    nn.BatchNorm1d(512),\n\n    nn.Dropout(0.4),\n\n    nn.Linear(512, 128),\n\n    nn.ReLU(),\n\n    nn.Dropout(0.3),\n\n    nn.Linear(128, 2)\n)\n\nmodel = model.to(device)\n\nprint(model.classifier)\n\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:14:33.941905Z","iopub.execute_input":"2026-05-14T17:14:33.942483Z","iopub.status.idle":"2026-05-14T17:14:35.300594Z","shell.execute_reply.started":"2026-05-14T17:14:33.942450Z","shell.execute_reply":"2026-05-14T17:14:35.299900Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# LOSS FUNCTION\n# ============================================================\n\ncriterion = nn.CrossEntropyLoss()\n\n# ============================================================\n# OPTIMIZER\n# ============================================================\n\noptimizer = optim.AdamW(\n\n    model.parameters(),\n\n    lr=1e-4,\n\n    weight_decay=1e-4\n)\n\n# ============================================================\n# LEARNING RATE SCHEDULER\n# ============================================================\n\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n\n    optimizer,\n\n    mode='min',\n\n    factor=0.5,\n\n    patience=2\n)\n\nprint(\"=\"*60)\nprint(\"TRAINING CONFIGURATION READY\")\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:14:47.750539Z","iopub.execute_input":"2026-05-14T17:14:47.751402Z","iopub.status.idle":"2026-05-14T17:14:47.757933Z","shell.execute_reply.started":"2026-05-14T17:14:47.751360Z","shell.execute_reply":"2026-05-14T17:14:47.757091Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# TRAINING FUNCTION\n# ============================================================\n\ndef train_one_epoch(model, loader):\n\n    model.train()\n\n    running_loss = 0\n\n    all_preds = []\n\n    all_labels = []\n\n    for images, labels in tqdm(loader):\n\n        images = images.to(device)\n\n        labels = labels.to(device)\n\n        optimizer.zero_grad()\n\n        # ====================================================\n        # FORWARD PASS\n        # ====================================================\n\n        outputs = model(images)\n\n        loss = criterion(outputs, labels)\n\n        # ====================================================\n        # BACKPROPAGATION\n        # ====================================================\n\n        loss.backward()\n\n        optimizer.step()\n\n        running_loss += loss.item()\n\n        _, preds = torch.max(outputs, 1)\n\n        all_preds.extend(\n            preds.cpu().numpy()\n        )\n\n        all_labels.extend(\n            labels.cpu().numpy()\n        )\n\n    # ========================================================\n    # METRICS\n    # ========================================================\n\n    epoch_loss = running_loss / len(loader)\n\n    epoch_acc = accuracy_score(\n        all_labels,\n        all_preds\n    )\n\n    epoch_precision = precision_score(\n        all_labels,\n        all_preds,\n        zero_division=0\n    )\n\n    epoch_recall = recall_score(\n        all_labels,\n        all_preds,\n        zero_division=0\n    )\n\n    epoch_f1 = f1_score(\n        all_labels,\n        all_preds,\n        zero_division=0\n    )\n\n    return (\n        epoch_loss,\n        epoch_acc,\n        epoch_precision,\n        epoch_recall,\n        epoch_f1\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:15:06.546154Z","iopub.execute_input":"2026-05-14T17:15:06.546761Z","iopub.status.idle":"2026-05-14T17:15:06.555871Z","shell.execute_reply.started":"2026-05-14T17:15:06.546734Z","shell.execute_reply":"2026-05-14T17:15:06.554800Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# VALIDATION FUNCTION\n# ============================================================\n\ndef validate(model, loader):\n\n    model.eval()\n\n    running_loss = 0\n\n    all_preds = []\n\n    all_labels = []\n\n    all_probs = []\n\n    with torch.no_grad():\n\n        for images, labels in tqdm(loader):\n\n            images = images.to(device)\n\n            labels = labels.to(device)\n\n            # =================================================\n            # FORWARD PASS\n            # =================================================\n\n            outputs = model(images)\n\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item()\n\n            probs = torch.softmax(\n                outputs,\n                dim=1\n            )\n\n            _, preds = torch.max(outputs, 1)\n\n            all_preds.extend(\n                preds.cpu().numpy()\n            )\n\n            all_labels.extend(\n                labels.cpu().numpy()\n            )\n\n            all_probs.extend(\n                probs[:,1].cpu().numpy()\n            )\n\n    # ========================================================\n    # METRICS\n    # ========================================================\n\n    epoch_loss = running_loss / len(loader)\n\n    epoch_acc = accuracy_score(\n        all_labels,\n        all_preds\n    )\n\n    epoch_precision = precision_score(\n        all_labels,\n        all_preds,\n        zero_division=0\n    )\n\n    epoch_recall = recall_score(\n        all_labels,\n        all_preds,\n        zero_division=0\n    )\n\n    epoch_f1 = f1_score(\n        all_labels,\n        all_preds,\n        zero_division=0\n    )\n\n    epoch_auc = roc_auc_score(\n        all_labels,\n        all_probs\n    )\n\n    return (\n\n        epoch_loss,\n\n        epoch_acc,\n\n        epoch_precision,\n\n        epoch_recall,\n\n        epoch_f1,\n\n        epoch_auc,\n\n        all_labels,\n\n        all_preds\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:15:23.390221Z","iopub.execute_input":"2026-05-14T17:15:23.390524Z","iopub.status.idle":"2026-05-14T17:15:23.398047Z","shell.execute_reply.started":"2026-05-14T17:15:23.390502Z","shell.execute_reply":"2026-05-14T17:15:23.397231Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# TRAINING LOOP\n# ============================================================\n\nEPOCHS = 10\n\nbest_val_f1 = 0\n\nhistory = {\n\n    \"train_loss\": [],\n    \"train_acc\": [],\n    \"train_f1\": [],\n\n    \"val_loss\": [],\n    \"val_acc\": [],\n    \"val_f1\": [],\n    \"val_auc\": []\n}\n\nprint(\"=\"*60)\nprint(\"STARTING TRAINING\")\nprint(\"=\"*60)\n\nfor epoch in range(EPOCHS):\n\n    print(f\"\\nEpoch [{epoch+1}/{EPOCHS}]\")\n\n    print(\"=\"*60)\n\n    # ========================================================\n    # TRAINING\n    # ========================================================\n\n    train_results = train_one_epoch(\n        model,\n        train_loader\n    )\n\n    (\n        train_loss,\n        train_acc,\n        train_precision,\n        train_recall,\n        train_f1\n    ) = train_results\n\n    # ========================================================\n    # VALIDATION\n    # ========================================================\n\n    val_results = validate(\n        model,\n        val_loader\n    )\n\n    (\n        val_loss,\n        val_acc,\n        val_precision,\n        val_recall,\n        val_f1,\n        val_auc,\n        val_labels,\n        val_preds\n    ) = val_results\n\n    # ========================================================\n    # SCHEDULER STEP\n    # ========================================================\n\n    scheduler.step(val_loss)\n\n    # ========================================================\n    # SAVE HISTORY\n    # ========================================================\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"train_acc\"].append(train_acc)\n    history[\"train_f1\"].append(train_f1)\n\n    history[\"val_loss\"].append(val_loss)\n    history[\"val_acc\"].append(val_acc)\n    history[\"val_f1\"].append(val_f1)\n    history[\"val_auc\"].append(val_auc)\n\n    # ========================================================\n    # PRINT RESULTS\n    # ========================================================\n\n    print(\"\\nTRAIN RESULTS\")\n    print(\"-\"*40)\n\n    print(f\"Loss       : {train_loss:.4f}\")\n    print(f\"Accuracy   : {train_acc:.4f}\")\n    print(f\"Precision  : {train_precision:.4f}\")\n    print(f\"Recall     : {train_recall:.4f}\")\n    print(f\"F1 Score   : {train_f1:.4f}\")\n\n    print(\"\\nVALIDATION RESULTS\")\n    print(\"-\"*40)\n\n    print(f\"Loss       : {val_loss:.4f}\")\n    print(f\"Accuracy   : {val_acc:.4f}\")\n    print(f\"Precision  : {val_precision:.4f}\")\n    print(f\"Recall     : {val_recall:.4f}\")\n    print(f\"F1 Score   : {val_f1:.4f}\")\n    print(f\"ROC-AUC    : {val_auc:.4f}\")\n\n    # ========================================================\n    # SAVE BEST MODEL\n    # ========================================================\n\n    if val_f1 > best_val_f1:\n\n        best_val_f1 = val_f1\n\n        torch.save(\n            model.state_dict(),\n            \"best_densenet121_lumbar.pth\"\n        )\n\n        print(\"\\nBest Model Saved!\")\n\nprint(\"\\nTRAINING COMPLETED\")\n\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:16:00.276940Z","iopub.execute_input":"2026-05-14T17:16:00.277908Z","iopub.status.idle":"2026-05-14T17:22:01.754255Z","shell.execute_reply.started":"2026-05-14T17:16:00.277879Z","shell.execute_reply":"2026-05-14T17:22:01.753164Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# LOAD BEST MODEL\n# ============================================================\n\nmodel.load_state_dict(\n    torch.load(\n        \"best_densenet121_lumbar.pth\"\n    )\n)\n\nmodel.eval()\n\nprint(\"=\"*60)\nprint(\"BEST MODEL LOADED SUCCESSFULLY\")\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:22:42.440895Z","iopub.execute_input":"2026-05-14T17:22:42.441227Z","iopub.status.idle":"2026-05-14T17:22:42.589928Z","shell.execute_reply.started":"2026-05-14T17:22:42.441197Z","shell.execute_reply":"2026-05-14T17:22:42.589176Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# FINAL VALIDATION\n# ============================================================\n\nfinal_results = validate(\n    model,\n    val_loader\n)\n\n(\n    final_loss,\n    final_acc,\n    final_precision,\n    final_recall,\n    final_f1,\n    final_auc,\n    final_labels,\n    final_preds\n) = final_results\n\nprint(\"=\"*60)\nprint(\"FINAL VALIDATION RESULTS\")\nprint(\"=\"*60)\n\nprint(f\"Loss       : {final_loss:.4f}\")\nprint(f\"Accuracy   : {final_acc:.4f}\")\nprint(f\"Precision  : {final_precision:.4f}\")\nprint(f\"Recall     : {final_recall:.4f}\")\nprint(f\"F1 Score   : {final_f1:.4f}\")\nprint(f\"ROC-AUC    : {final_auc:.4f}\")\n\nprint(\"=\"*60)\n\n# ============================================================\n# CONFUSION MATRIX\n# ============================================================\n\ncm = confusion_matrix(\n    final_labels,\n    final_preds\n)\n\nprint(\"\\nCONFUSION MATRIX\\n\")\n\nprint(cm)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:22:52.301037Z","iopub.execute_input":"2026-05-14T17:22:52.301325Z","iopub.status.idle":"2026-05-14T17:22:59.154453Z","shell.execute_reply.started":"2026-05-14T17:22:52.301303Z","shell.execute_reply":"2026-05-14T17:22:59.153408Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CLASSIFICATION REPORT\n# ============================================================\n\nfrom sklearn.metrics import classification_report\n\nreport = classification_report(\n\n    final_labels,\n\n    final_preds,\n\n    target_names=[\n        \"Normal\",\n        \"Abnormal\"\n    ],\n\n    zero_division=0\n)\n\nprint(\"=\"*60)\nprint(\"CLASSIFICATION REPORT\")\nprint(\"=\"*60)\n\nprint(report)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:23:19.880750Z","iopub.execute_input":"2026-05-14T17:23:19.881226Z","iopub.status.idle":"2026-05-14T17:23:19.897724Z","shell.execute_reply.started":"2026-05-14T17:23:19.881195Z","shell.execute_reply":"2026-05-14T17:23:19.896874Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# TRAINING CURVES\n# ============================================================\n\nplt.figure(figsize=(14,5))\n\n# ============================================================\n# LOSS CURVE\n# ============================================================\n\nplt.subplot(1,2,1)\n\nplt.plot(\n    history[\"train_loss\"],\n    label=\"Train Loss\"\n)\n\nplt.plot(\n    history[\"val_loss\"],\n    label=\"Validation Loss\"\n)\n\nplt.xlabel(\"Epoch\")\n\nplt.ylabel(\"Loss\")\n\nplt.title(\"Loss Curve\")\n\nplt.legend()\n\n# ============================================================\n# F1 CURVE\n# ============================================================\n\nplt.subplot(1,2,2)\n\nplt.plot(\n    history[\"train_f1\"],\n    label=\"Train F1\"\n)\n\nplt.plot(\n    history[\"val_f1\"],\n    label=\"Validation F1\"\n)\n\nplt.xlabel(\"Epoch\")\n\nplt.ylabel(\"F1 Score\")\n\nplt.title(\"F1 Curve\")\n\nplt.legend()\n\nplt.tight_layout()\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:23:34.396581Z","iopub.execute_input":"2026-05-14T17:23:34.397307Z","iopub.status.idle":"2026-05-14T17:23:34.692216Z","shell.execute_reply.started":"2026-05-14T17:23:34.397281Z","shell.execute_reply":"2026-05-14T17:23:34.691316Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install grad-cam -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:28:32.207585Z","iopub.execute_input":"2026-05-14T17:28:32.208074Z","iopub.status.idle":"2026-05-14T17:28:43.557159Z","shell.execute_reply.started":"2026-05-14T17:28:32.208045Z","shell.execute_reply":"2026-05-14T17:28:43.556412Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# GRAD-CAM IMPORTS\n# ============================================================\n\nfrom pytorch_grad_cam import GradCAM\n\nfrom pytorch_grad_cam.utils.image import (\n    show_cam_on_image\n)\n\nfrom pytorch_grad_cam.utils.model_targets import (\n    ClassifierOutputTarget\n)\n\nprint(\"=\"*60)\nprint(\"GRAD-CAM LIBRARY LOADED\")\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:28:47.980921Z","iopub.execute_input":"2026-05-14T17:28:47.981561Z","iopub.status.idle":"2026-05-14T17:28:48.197165Z","shell.execute_reply.started":"2026-05-14T17:28:47.981526Z","shell.execute_reply":"2026-05-14T17:28:48.196263Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.features.norm5","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:29:01.571885Z","iopub.execute_input":"2026-05-14T17:29:01.572688Z","iopub.status.idle":"2026-05-14T17:29:01.578880Z","shell.execute_reply.started":"2026-05-14T17:29:01.572657Z","shell.execute_reply":"2026-05-14T17:29:01.578125Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# TARGET LAYER FOR GRAD-CAM\n# ============================================================\n\ntarget_layers = [model.features.norm5]\n\nprint(\"=\"*60)\nprint(\"TARGET LAYER SELECTED\")\nprint(\"=\"*60)\n\nprint(target_layers)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:29:15.495895Z","iopub.execute_input":"2026-05-14T17:29:15.496664Z","iopub.status.idle":"2026-05-14T17:29:15.501920Z","shell.execute_reply.started":"2026-05-14T17:29:15.496638Z","shell.execute_reply":"2026-05-14T17:29:15.500861Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# GRAD-CAM VISUALIZATION FUNCTION\n# ============================================================\n\ndef visualize_gradcam(model, loader, num_images=6):\n\n    model.eval()\n\n    images, labels = next(iter(loader))\n\n    images = images.to(device)\n\n    labels = labels.to(device)\n\n    # ========================================================\n    # INITIALIZE GRAD-CAM\n    # ========================================================\n\n    cam = GradCAM(\n        model=model,\n        target_layers=target_layers\n    )\n\n    # ========================================================\n    # MODEL PREDICTIONS\n    # ========================================================\n\n    outputs = model(images)\n\n    probs = torch.softmax(outputs, dim=1)\n\n    preds = torch.argmax(outputs, dim=1)\n\n    # ========================================================\n    # VISUALIZATION\n    # ========================================================\n\n    plt.figure(figsize=(18, 10))\n\n    for i in range(num_images):\n\n        # ====================================================\n        # IMAGE PREPARATION\n        # ====================================================\n\n        image_tensor = images[i]\n\n        image_np = image_tensor.permute(\n            1, 2, 0\n        ).detach().cpu().numpy()\n\n        # ====================================================\n        # DE-NORMALIZE\n        # ====================================================\n\n        image_np = image_np * np.array(\n            [0.229, 0.224, 0.225]\n        ) + np.array(\n            [0.485, 0.456, 0.406]\n        )\n\n        image_np = np.clip(image_np, 0, 1)\n\n        # ====================================================\n        # TARGET CLASS\n        # ====================================================\n\n        target_category = preds[i].item()\n\n        targets = [\n            ClassifierOutputTarget(\n                target_category\n            )\n        ]\n\n        # ====================================================\n        # GENERATE CAM\n        # ====================================================\n\n        grayscale_cam = cam(\n            input_tensor=image_tensor.unsqueeze(0),\n            targets=targets\n        )\n\n        grayscale_cam = grayscale_cam[0]\n\n        # ====================================================\n        # OVERLAY HEATMAP\n        # ====================================================\n\n        visualization = show_cam_on_image(\n\n            image_np,\n\n            grayscale_cam,\n\n            use_rgb=True\n        )\n\n        # ====================================================\n        # DISPLAY\n        # ====================================================\n\n        plt.subplot(2, 3, i + 1)\n\n        plt.imshow(visualization)\n\n        plt.axis(\"off\")\n\n        true_label = labels[i].item()\n\n        pred_label = preds[i].item()\n\n        confidence = probs[i][pred_label].item()\n\n        plt.title(\n\n            f\"True : {true_label}\\n\"\n            f\"Pred : {pred_label}\\n\"\n            f\"Conf : {confidence:.2f}\"\n        )\n\n    plt.tight_layout()\n\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:29:29.586211Z","iopub.execute_input":"2026-05-14T17:29:29.586753Z","iopub.status.idle":"2026-05-14T17:29:29.596725Z","shell.execute_reply.started":"2026-05-14T17:29:29.586721Z","shell.execute_reply":"2026-05-14T17:29:29.595846Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# RUN GRAD-CAM\n# ============================================================\n\nvisualize_gradcam(\n\n    model,\n\n    val_loader,\n\n    num_images=6\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T17:29:44.110945Z","iopub.execute_input":"2026-05-14T17:29:44.111422Z","iopub.status.idle":"2026-05-14T17:29:45.967475Z","shell.execute_reply.started":"2026-05-14T17:29:44.111392Z","shell.execute_reply":"2026-05-14T17:29:45.966465Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}