{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"9a676ffc","cell_type":"markdown","source":"# Two-Branch VGG16 PyTorch Notebook\n\nTrain a two-branch VGG16 on Axial T2 and Sagittal T2/STIR lumbar spine DICOM images for severity classification.","metadata":{}},{"id":"5374820a","cell_type":"code","source":"# === 1) Imports & Utilities ===\nimport os\nimport numpy as np\nimport pandas as pd\nimport pydicom\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torchvision import transforms, models\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import precision_recall_fscore_support, confusion_matrix, classification_report\nimport matplotlib.pyplot as plt\n\n# Function: load DICOM and convert to PIL image\ndef load_dicom_image(path):\n    ds = pydicom.dcmread(path)\n    img = ds.pixel_array.astype(np.float32)\n    img = img - img.min()\n    img = img / (img.max() + 1e-6)\n    img = (img * 255).astype(np.uint8)\n    # create PIL and convert to 3-channel RGB\n    return Image.fromarray(img).convert('RGB')\n\n\n# in your transforms definitions:\n\ncommon_norm = transforms.Normalize(mean=[0.485,0.456,0.406],\n                                  std =[0.229,0.224,0.225])\n\naxial_transforms = transforms.Compose([\n    transforms.Resize(256),\n    transforms.CenterCrop(224),\n    transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    common_norm\n])\nsag_transforms = transforms.Compose([\n    transforms.Resize(256),\n    transforms.CenterCrop(224),\n    transforms.RandomAffine(15, translate=(0.1,0.1)),\n    transforms.ToTensor(),\n    common_norm\n])\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T17:55:14.441650Z","iopub.execute_input":"2025-05-25T17:55:14.441900Z","iopub.status.idle":"2025-05-25T17:55:23.632902Z","shell.execute_reply.started":"2025-05-25T17:55:14.441881Z","shell.execute_reply":"2025-05-25T17:55:23.632345Z"}},"outputs":[],"execution_count":null},{"id":"2a40c861","cell_type":"code","source":"# === 2) Load CSVs & Generate Image Paths ===\n# Update train_path to your data directory\ntrain_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\ntrain_df      = pd.read_csv(os.path.join(train_path, 'train.csv'))\nlabel_df      = pd.read_csv(os.path.join(train_path, 'train_label_coordinates.csv'))\ntrain_desc_df = pd.read_csv(os.path.join(train_path, 'train_series_descriptions.csv'))\n\n# ==== Rebuild paths_df correctly ====\ndef generate_image_paths(desc_df, image_root):\n    records = []\n    for _, row in desc_df.iterrows():\n        study_id           = row['study_id']\n        series_id          = row['series_id']\n        series_description = row['series_description']  # use the row’s own description\n        series_dir = os.path.join(image_root, str(study_id), str(series_id))\n        if not os.path.isdir(series_dir):\n            continue\n        for fname in os.listdir(series_dir):\n            # (optionally filter for .dcm)\n            records.append({\n                'study_id':            study_id,\n                'series_id':           series_id,\n                'series_description':  series_description,\n                'image_path':          os.path.join(series_dir, fname)\n            })\n    return pd.DataFrame(records)\n\n# Usage: point `train_images_dir` at wherever your DICOM folders live\ntrain_images_dir = os.path.join(train_path, 'train_images')\npaths_df = generate_image_paths(train_desc_df, train_images_dir)\n\nprint(\"✔ paths_df:\", paths_df.shape)\nprint(paths_df['series_description'].value_counts())\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T17:55:23.633925Z","iopub.execute_input":"2025-05-25T17:55:23.634314Z","iopub.status.idle":"2025-05-25T17:56:55.988277Z","shell.execute_reply.started":"2025-05-25T17:55:23.634276Z","shell.execute_reply":"2025-05-25T17:56:55.987580Z"}},"outputs":[],"execution_count":null},{"id":"589c9a33-e0a7-4fdd-afe0-0f6e4b10b17b","cell_type":"code","source":"# ==== Helper: unpivot the original train.csv into (study,condition,level,severity) rows ====\ndef reshape_row(row):\n    data = {'study_id': [], 'condition': [], 'level': [], 'severity': []}\n    # skip the non‐severity columns\n    for col, val in row.items():\n        if col in ['study_id','series_id','instance_number','x','y','series_description']:\n            continue\n        parts = col.split('_')\n        # reconstruct the condition name (all but last two tokens)\n        condition = ' '.join([w.capitalize() for w in parts[:-2]])\n        # reconstruct the level as e.g. 'L1/L2'\n        level = parts[-2].upper() + '/' + parts[-1].upper()\n        data['study_id'].append(row['study_id'])\n        data['condition'].append(condition)\n        data['level'].append(level)\n        data['severity'].append(val)\n    return pd.DataFrame(data)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T18:00:26.148576Z","iopub.execute_input":"2025-05-25T18:00:26.149159Z","iopub.status.idle":"2025-05-25T18:00:26.156604Z","shell.execute_reply.started":"2025-05-25T18:00:26.149130Z","shell.execute_reply":"2025-05-25T18:00:26.155933Z"}},"outputs":[],"execution_count":null},{"id":"b70d6f43","cell_type":"code","source":"# ==== DEBUGGED MERGE SECTION ====\n\n# 1) Build the flat label table\nnew_train_df = pd.concat([reshape_row(row) for _, row in train_df.iterrows()],\n                         ignore_index=True)\nprint(\"➜ new_train_df shape:\",      new_train_df.shape)\nprint(new_train_df.head(3))\n\n# 2) Normalize text keys so merges won’t fail on caps/slashes/spaces\nnew_train_df['condition'] = (\n    new_train_df['condition']\n    .str.lower()\n    .str.replace(' ', '_')\n)\nnew_train_df['level'] = (\n    new_train_df['level']\n    .str.lower()\n    .str.replace('/', '_')\n)\nlabel_df['condition'] = (\n    label_df['condition']\n    .str.lower()\n    .str.replace(' ', '_')\n)\nlabel_df['level'] = (\n    label_df['level']\n    .str.lower()\n    .str.replace('/', '_')\n)\n\n# 3) Merge1: labels → coordinates\nmerged1 = pd.merge(\n    new_train_df,\n    label_df,\n    on=['study_id', 'condition', 'level'],\n    how='inner'\n)\nprint(\"➜ after merge1:\", merged1.shape)\nprint(merged1[['study_id','condition','level']].drop_duplicates().head(3))\n\n# 4) Confirm `series_id` alignment\nprint(\"➜ merged1 columns:\", merged1.columns.tolist())\nprint(\"➜ paths_df columns:\", paths_df.columns.tolist())\nprint(\"➜ dtype(series_id):\", merged1['series_id'].dtype,\n      paths_df['series_id'].dtype)\n\n# 5) Merge2: bring in series_description & image_path\nmerged2 = merged1.merge(\n    paths_df[['study_id','series_id','series_description','image_path']],\n    on=['study_id','series_id'],\n    how='inner'\n)\nprint(\"➜ after merge2:\", merged2.shape)\nprint(\"➜ unique descriptions:\", merged2['series_description'].unique())\n\n# 6) Filter to only the two T2 views you want\nmerged2 = merged2[merged2['series_description']\n                  .isin(['Axial T2', 'Sagittal T2/STIR'])].copy()\nprint(\"➜ after filtering modalities:\", merged2.shape)\n\n# 7) Normalize severity text, then map to ints\nmerged2['severity'] = (\n    merged2['severity']\n    .str.lower()\n    .str.replace('/', '_')\n)\nseverity_map = {'normal_mild':0, 'moderate':1, 'severe':2}\nmerged2['severity'] = merged2['severity'].map(severity_map)\n\n# 8) Drop any rows that failed to map or lost their path\nmerged2 = merged2.dropna(subset=['image_path','severity']).reset_index(drop=True)\nprint(\"➜ final merged2 shape:\", merged2.shape)\n\n# 9) Ready for splitting\ntrain_data = merged2.copy()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T18:00:34.414898Z","iopub.execute_input":"2025-05-25T18:00:34.415447Z","iopub.status.idle":"2025-05-25T18:00:36.515673Z","shell.execute_reply.started":"2025-05-25T18:00:34.415425Z","shell.execute_reply":"2025-05-25T18:00:36.514873Z"}},"outputs":[],"execution_count":null},{"id":"ab4848a5","cell_type":"code","source":"# ==== 4) Dataset, Transforms & DataLoader (classification-only) ====\n\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\n\n# 4a) Build per-study map of the two T2 series IDs\nseries_map = (\n    train_desc_df[\n        train_desc_df['series_description']\n        .isin(['Axial T2', 'Sagittal T2/STIR'])\n    ]\n    .drop_duplicates(subset=['study_id','series_description'])\n    .pivot(index='study_id',\n           columns='series_description',\n           values='series_id')\n    .reset_index()\n)\nseries_map.columns = ['study_id','ax_series_id','sag_series_id']\nprint(\"→ series_map:\", series_map.shape)\n\n# 4b) Merge classification labels with that map\nclf_map = pd.merge(new_train_df, series_map, on='study_id', how='inner')\nprint(\"→ classification merge:\", clf_map.shape)\n\n# 4c) Grab one image path per series_id\nax_paths = (\n    paths_df[paths_df['series_description']=='Axial T2']\n    .drop_duplicates(subset=['series_id'], keep='first')\n    [['series_id','image_path']]\n    .rename(columns={'series_id':'ax_series_id',\n                     'image_path':'axial_t2_path'})\n)\nsag_paths = (\n    paths_df[paths_df['series_description']=='Sagittal T2/STIR']\n    .drop_duplicates(subset=['series_id'], keep='first')\n    [['series_id','image_path']]\n    .rename(columns={'series_id':'sag_series_id',\n                     'image_path':'sagittal_t2stir_path'})\n)\n\n# 4d) Attach file paths\ndataset_df = (\n    clf_map\n    .merge(ax_paths, on='ax_series_id', how='inner')\n    .merge(sag_paths, on='sag_series_id', how='inner')\n)\nprint(\"→ before cleanup:\", dataset_df.shape)\n\n# 4e) Keep only what you need and robustly map severity → 0/1/2\ndataset_df = dataset_df[['axial_t2_path','sagittal_t2stir_path','severity']].copy()\n\n# 4e.1) Normalize the text\ndataset_df['severity_norm'] = (\n    dataset_df['severity']\n    .str.lower()\n    .str.replace('/', '_')\n    .str.strip()\n)\n\n# 4e.2) Map to integers\nseverity_map = {'normal_mild':0, 'moderate':1, 'severe':2}\ndataset_df['severity_mapped'] = dataset_df['severity_norm'].map(severity_map)\n\n# 4e.3) Inspect any unmapped rows (optional)\nunmapped = dataset_df['severity_mapped'].isna().sum()\nprint(f\"Unmapped severity rows: {unmapped}\")\n\n# 4e.4) Drop rows we couldn’t map\ndataset_df = dataset_df.dropna(subset=['severity_mapped']).copy()\n\n# 4e.5) Finalize the column\ndataset_df['severity'] = dataset_df['severity_mapped'].astype(int)\ndataset_df = dataset_df[['axial_t2_path','sagittal_t2stir_path','severity']]\n\nprint(\"→ final dataset_df (post‐drop):\", dataset_df.shape)\n\n# 4f) Transforms for each branch\naxial_transforms = transforms.Compose([\n    transforms.Resize(256), transforms.CenterCrop(224),\n    transforms.RandomHorizontalFlip(), transforms.ToTensor(),\n    transforms.Normalize([0.485],[0.229])\n])\nsag_transforms = transforms.Compose([\n    transforms.Resize(256), transforms.CenterCrop(224),\n    transforms.RandomAffine(15, translate=(0.1,0.1)),\n    transforms.ToTensor(), transforms.Normalize([0.485],[0.229])\n])\n\n# 4g) Two-branch Dataset\nclass SpineT2PairDataset(Dataset):\n    def __init__(self, df, axial_transform, sag_transform):\n        self.df = df\n        self.axial_transform = axial_transform\n        self.sag_transform   = sag_transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i):\n        rec = self.df.iloc[i]\n        img_ax  = load_dicom_image(rec['axial_t2_path'])\n        img_sag = load_dicom_image(rec['sagittal_t2stir_path'])\n        if self.axial_transform: img_ax  = self.axial_transform(img_ax)\n        if self.sag_transform:   img_sag = self.sag_transform(img_sag)\n        return img_ax, img_sag, int(rec['severity'])\n\n# 4h) Stratified split & sampler\ntrain_df, val_df = train_test_split(\n    dataset_df,\n    stratify=dataset_df['severity'],\n    test_size=0.3,\n    random_state=42\n)\nprint(\"Train/Val:\", train_df.shape, \"/\", val_df.shape)\n\ncounts = train_df['severity'].value_counts().sort_index().values\nclass_weights = 1.0 / counts\nsample_weights = train_df['severity'].map(lambda x: class_weights[x]).values\nsampler = WeightedRandomSampler(\n    sample_weights,\n    num_samples=len(sample_weights),\n    replacement=True\n)\n\n# 4i) DataLoaders\ntrain_ds = SpineT2PairDataset(train_df, axial_transforms, sag_transforms)\nval_ds   = SpineT2PairDataset(val_df,   axial_transforms, sag_transforms)\n\ntrain_loader = DataLoader(\n    train_ds, batch_size=16, sampler=sampler,\n    num_workers=4, pin_memory=True\n)\nval_loader = DataLoader(\n    val_ds, batch_size=16, shuffle=False,\n    num_workers=4, pin_memory=True\n)\n\nprint(\"Batches/epoch:\", len(train_loader), \"/\", len(val_loader))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T18:00:40.601493Z","iopub.execute_input":"2025-05-25T18:00:40.601775Z","iopub.status.idle":"2025-05-25T18:00:40.785789Z","shell.execute_reply.started":"2025-05-25T18:00:40.601749Z","shell.execute_reply":"2025-05-25T18:00:40.785039Z"}},"outputs":[],"execution_count":null},{"id":"ce7c854d-a315-4f26-87a5-9a407b0ee7d1","cell_type":"code","source":"print(\"→ dataset_df shape:\", dataset_df.shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T18:00:46.098249Z","iopub.execute_input":"2025-05-25T18:00:46.098551Z","iopub.status.idle":"2025-05-25T18:00:46.102615Z","shell.execute_reply.started":"2025-05-25T18:00:46.098532Z","shell.execute_reply":"2025-05-25T18:00:46.101779Z"}},"outputs":[],"execution_count":null},{"id":"14cac909-d420-484d-bea8-01e7e11558a4","cell_type":"code","source":"# ==== 5) Model Definition, Loss & Optimizer ====\n\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torchvision import models\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(\"Using device:\", device)\n\n# Two-branch VGG16\nclass TwoBranchVGG16(nn.Module):\n    def __init__(self, num_classes=3, pretrained=True):\n        super().__init__()\n        base = models.vgg16(pretrained=pretrained)\n        # shared feature extractor\n        self.features = base.features\n        self.avgpool  = base.avgpool\n        feat_dim = base.classifier[0].in_features  # typically 25088\n\n        # fusion + classifier\n        self.classifier = nn.Sequential(\n            nn.Linear(feat_dim*2, 512),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.5),\n            nn.Linear(512, num_classes)\n        )\n\n    def forward(self, x_ax, x_sag):\n        fa = self.features(x_ax)\n        fa = self.avgpool(fa)\n        fa = fa.flatten(1)\n        fs = self.features(x_sag)\n        fs = self.avgpool(fs)\n        fs = fs.flatten(1)\n        f  = torch.cat([fa, fs], dim=1)\n        return self.classifier(f)\n\n# Instantiate model\nmodel = TwoBranchVGG16(num_classes=3).to(device)\n\n# Use class_weights from your sampler setup\nclass_weights_tensor = torch.tensor(class_weights, dtype=torch.float, device=device)\n\ncriterion = nn.CrossEntropyLoss(weight=class_weights_tensor)\noptimizer = optim.Adam(model.parameters(), lr=1e-4)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=3)\n\nprint(model)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T18:00:48.744783Z","iopub.execute_input":"2025-05-25T18:00:48.745028Z","iopub.status.idle":"2025-05-25T18:00:53.870908Z","shell.execute_reply.started":"2025-05-25T18:00:48.745004Z","shell.execute_reply":"2025-05-25T18:00:53.870278Z"}},"outputs":[],"execution_count":null},{"id":"5ec9aa56-2a2f-4efc-99e2-633649a2c884","cell_type":"code","source":"# ==== 6) Training & Evaluation Functions ====\n\nfrom sklearn.metrics import precision_recall_fscore_support, confusion_matrix\n\ndef train_epoch(loader):\n    model.train()\n    total_loss, total_correct, total_samples = 0., 0, 0\n    for i, (ax, sag, labels) in enumerate(loader, 1):\n        # every 200 batches, show how far along we are\n        if i % 200 == 0:\n            print(f\"  processed {i}/{len(loader)} batches\", end=\"\\r\", flush=True)\n\n        ax, sag, labels = ax.to(device), sag.to(device), labels.to(device)\n        optimizer.zero_grad()\n        outputs = model(ax, sag)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        total_loss   += loss.item() * labels.size(0)\n        preds         = outputs.argmax(dim=1)\n        total_correct += (preds == labels).sum().item()\n        total_samples += labels.size(0)\n\n    # print a newline to clear the \"\\r\" on the last batch\n    print()\n    avg_loss = total_loss / total_samples\n    accuracy = total_correct / total_samples\n    return avg_loss, accuracy\n\ndef eval_epoch(loader):\n    model.eval()\n    total_loss, total_correct, total_samples = 0., 0, 0\n    all_preds, all_labels = [], []\n    with torch.no_grad():\n        for ax, sag, labels in loader:\n            ax, sag, labels = ax.to(device), sag.to(device), labels.to(device)\n            outputs = model(ax, sag)\n            loss = criterion(outputs, labels)\n\n            total_loss += loss.item() * labels.size(0)\n            preds = outputs.argmax(dim=1)\n            total_correct += (preds == labels).sum().item()\n            total_samples += labels.size(0)\n\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n\n    avg_loss = total_loss / total_samples\n    accuracy = total_correct / total_samples\n    p, r, f1, _ = precision_recall_fscore_support(all_labels, all_preds,\n                                                   labels=[0,1,2], zero_division=0)\n    cm = confusion_matrix(all_labels, all_preds)\n    return avg_loss, accuracy, p, r, f1, cm\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T18:00:57.011974Z","iopub.execute_input":"2025-05-25T18:00:57.012505Z","iopub.status.idle":"2025-05-25T18:00:57.023532Z","shell.execute_reply.started":"2025-05-25T18:00:57.012477Z","shell.execute_reply":"2025-05-25T18:00:57.022882Z"}},"outputs":[],"execution_count":null},{"id":"ac6610d1-7001-492f-9a23-abd141653eed","cell_type":"code","source":"import os\nimport torch\n\n# Where to save checkpoints\ncheckpoint_dir = \"checkpoints\"\nos.makedirs(checkpoint_dir, exist_ok=True)\n\n# Try to resume from best checkpoint if it exists\nbest_val_loss = float('inf')\nstart_epoch  = 1\nbest_ckpt    = os.path.join(checkpoint_dir, \"best.pt\")\nif os.path.isfile(best_ckpt):\n    ckpt = torch.load(best_ckpt, map_location=device)\n    model.load_state_dict(ckpt['model_state_dict'])\n    optimizer.load_state_dict(ckpt['optimizer_state_dict'])\n    best_val_loss = ckpt['val_loss']\n    start_epoch   = ckpt['epoch'] + 1\n    print(f\"Resuming from epoch {start_epoch} with val_loss={best_val_loss:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T18:01:00.059141Z","iopub.execute_input":"2025-05-25T18:01:00.059662Z","iopub.status.idle":"2025-05-25T18:01:00.064949Z","shell.execute_reply.started":"2025-05-25T18:01:00.059641Z","shell.execute_reply":"2025-05-25T18:01:00.064194Z"}},"outputs":[],"execution_count":null},{"id":"0b2af739-4d82-4098-82a9-a8011aac7a2d","cell_type":"code","source":"# ==== 7) Training Loop ====\n\nnum_epochs = 10\nhistory = {\n    'train_loss': [], 'train_acc': [],\n    'val_loss':   [], 'val_acc':   []\n}\n\nfor epoch in range(start_epoch, num_epochs+1):\n    print(f\"\\n Starting epoch {epoch}/{num_epochs}\", flush=True)\n    \n    tr_loss, tr_acc = train_epoch(train_loader)\n    print()\n    \n    val_loss, val_acc, p, r, f1, cm = eval_epoch(val_loader)\n    scheduler.step(val_loss)\n\n    history['train_loss'].append(tr_loss)\n    history['train_acc'].append(tr_acc)\n    history['val_loss'].append(val_loss)\n    history['val_acc'].append(val_acc)\n\n    # ---- save epoch checkpoint ----\n    epoch_ckpt = {\n        'epoch': epoch,\n        'model_state_dict': model.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'val_loss': val_loss\n    }\n    torch.save(epoch_ckpt, os.path.join(checkpoint_dir, f\"epoch_{epoch}.pt\"))\n\n    # ---- save best checkpoint ----\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        torch.save(epoch_ckpt, best_ckpt)\n        print(f\"  ✔ New best model saved (val_loss={val_loss:.4f})\")\n\n    # ---- print metrics ----\n    print(f\"  Precision: {p}, Recall: {r}, F1: {f1}\")\n    print(f\"  Confusion Matrix:\\n{cm}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T18:01:01.935275Z","iopub.execute_input":"2025-05-25T18:01:01.935532Z"}},"outputs":[],"execution_count":null},{"id":"dda7ec71-bcac-4735-bfed-7d1ee42ed18c","cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}