{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"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":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":9187072,"sourceType":"datasetVersion","datasetId":5504483}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This notebook has taken reference from [this](https://www.kaggle.com/competitions/rsna-2024-lumbar-spine-degenerative-classification/writeups/ngel-jacinto-s-nchez-ruiz-27th-place-solution-algo) submission in RSNA 2024 competition. I am using this for stage 1 of my framework to classify Lumbar Spine degeneration severity.","metadata":{}},{"cell_type":"code","source":"!pip install -q git+https://github.com/qubvel/segmentation_models.pytorch","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nfrom PIL import Image, ImageFilter\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nimport segmentation_models_pytorch as smp\nfrom torch.utils.data import Dataset, DataLoader\nimport os\nimport copy\nimport pydicom \nimport matplotlib.pyplot as plt\nimport random\nimport torchvision.transforms.functional as TF\nimport torchvision.transforms as T\nimport torchvision\nfrom fastai.vision.all import *\nimport math\nimport glob\nimport gc","metadata":{"_uuid":"75cfffbe-a831-492a-8c5c-8bb55ff70ffc","_cell_guid":"8aeb694e-77b7-4f74-9a9c-33ce009e355d","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n\nH = W = 512\nLEVELS = [\"L1L2\", \"L2L3\", \"L3L4\", \"L4L5\", \"L5S1\"]\nK = 5\nSIGMA = 5 # std deviation - spread of the heatmap keypoints\nANGLE = 30\nLR = 5e-4\nSIGMA = torch.as_tensor(SIGMA)\nA = -1/(2*(SIGMA**2)).to(DEVICE)\n\nEPOCHS = 2\nWEIGHT_DECAY = 1e-4\nTH = 0.5\nS2 = 64\nS2 = torch.as_tensor(S2)\nBATCH_SIZE = 16\n\nRSNA_TRAIN_IMGS = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images'\nRSNA_TEST_IMGS = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images'\n\ntorch_resize = torchvision.transforms.Resize((H,W),antialias=True)\n\nx_map = torch.stack([torch.arange(W)]*H).float()\ny_map = torch.stack([torch.arange(H)]*W).float()\nIDX_MAP = torch.stack([x_map,y_map.T]).view(1,2,H,W).to(DEVICE) # will be used to create the ground truth heatmaps","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **SAGITTAL T1**","metadata":{"_uuid":"f3549eb4-1e10-46a2-9366-ec20a09570fa","_cell_guid":"39f9f655-8cd0-40d6-b14d-7e06906061b8","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"markdown","source":"**Data preparation**","metadata":{"_uuid":"e3920502-4cc4-406b-baa1-9d5200b2f5f2","_cell_guid":"20e98ada-0319-4c2e-a402-d40915c6838a","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"### RSNA DATASET ####\n\nrsna = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv\").dropna()\nfn = rsna[rsna['condition'].isin(['Left Neural Foraminal Narrowing', 'Right Neural Foraminal Narrowing'])].drop(columns=['condition']).sort_values(['study_id','series_id','level']).reset_index(drop=True)\nprint(f\"Number of studies with unlabeled foraminal discs keypoints : {(fn.series_id.value_counts() < 10).sum()}\\n\")","metadata":{"_uuid":"c0922f69-e5cc-438f-9f24-0f8a3de5e122","_cell_guid":"24cfd0ea-502a-4a30-a309-a4166c25d3f9","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"coordinates = {}\nfor i in range(len(fn)):\n    row = fn.iloc[i]\n    coordinates[row['study_id']] = {}\nfor i in range(len(fn)):\n    row = fn.iloc[i]\n    coordinates[row['study_id']][row['series_id']] = {}\n\nfor i in range(len(fn)):\n    row = fn.iloc[i]\n    coordinates[row['study_id']][row['series_id']][row['instance_number']] = {\n        'L1/L2':{\n            'x':torch.nan,\n            'y':torch.nan\n        },\n        'L2/L3':{\n            'x':torch.nan,\n            'y':torch.nan\n        },\n        'L3/L4':{\n            'x':torch.nan,\n            'y':torch.nan\n        },\n        'L4/L5':{\n            'x':torch.nan,\n            'y':torch.nan\n        },\n        'L5/S1':{\n            'x':torch.nan,\n            'y':torch.nan\n        }\n    }\n\nfor i in range(len(fn)):\n    row = fn.iloc[i]\n    coordinates[row['study_id']][row['series_id']][row['instance_number']][row['level']]['x'] = row['x']\n    coordinates[row['study_id']][row['series_id']][row['instance_number']][row['level']]['y'] = row['y']","metadata":{"_uuid":"71919893-fcd4-4c33-ab07-c655f654bb7d","_cell_guid":"7d9831e6-36a6-4761-b9bb-39a53baa2960","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Create a dataframe where each row has a unique combination of study_id, series_id & instance_number\nfn = fn[['study_id', 'series_id', 'instance_number']].groupby(['study_id', 'series_id', 'instance_number']).count().reset_index()\n\n## We will only have a subset of keypoints for any instance numbers, rest will be torch.nan --> Have to impute\nv = np.zeros((len(fn),10))\nfrom tqdm import tqdm\nfor i in tqdm(range(len(fn))):\n    row = fn.iloc[i]\n    k = 0\n    for level in coordinates[row['study_id']][row['series_id']][row['instance_number']]:\n        v[i,k:k+2] = list(coordinates[row['study_id']][row['series_id']][row['instance_number']][level].values())\n        k += 2\n\ncoor = ['x_L1L2','y_L1L2','x_L2L3','y_L2L3','x_L3L4','y_L3L4','x_L4L5','y_L4L5','x_L5S1','y_L5S1']\nfn[coor] = v\n\nprint(f\"% of rows with NaN : {(fn.isna().any(axis = 1).sum() / len(fn))*100}\")","metadata":{"_uuid":"78568fb2-70af-4143-aaa7-c8ab97895cab","_cell_guid":"74d1be12-f42e-40f2-b22a-1398dfc94780","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## For every combination of study_id + series_id, augment the data for instance_numbers not present in the current dataframe\n### eg : instance_numbers for some (study_id, series_id) -> [3, 5, 8, 12]; missing ones are 4, 6, 7, 9, 10, 11\n\nrows_to_add = []\n\nfor (study_id, series_id), df in tqdm(fn.groupby(['study_id', 'series_id'])):\n    existing_instances = sorted(df['instance_number'].unique())\n    min_inst, max_inst = min(existing_instances), max(existing_instances)\n    full_range = list(range(min_inst, max_inst + 1))\n    \n    # Find missing instance numbers\n    missing = sorted(set(full_range) - set(existing_instances))\n    \n    # Append missing rows (NaN coordinates)\n    for inst in missing:\n        rows_to_add.append({\n            'study_id': int(study_id),\n            'series_id': int(series_id),\n            'instance_number': inst,\n            **{c: torch.nan for c in coor}\n        })\n\n# Append to fn\nif rows_to_add:\n    fn = pd.concat([fn, pd.DataFrame(rows_to_add)], ignore_index=True)\n\nfn[['study_id', 'series_id', 'instance_number']] = fn[['study_id', 'series_id', 'instance_number']].astype(np.int64)\nfn = fn.sort_values(['study_id', 'series_id', 'instance_number']).reset_index(drop=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for col in coor:\n    fn[col] = fn.groupby(['study_id', 'series_id'])[col].transform(\n        lambda x: x.fillna(x.mean())\n    )\n\ndata_T1 = fn\n# fn_mean = fn.groupby(['study_id', 'series_id']).mean()\n# fn_mean = fn_mean.loc[[(study_id,series_id) for study_id,series_id in fn[['study_id','series_id']].values]][coor].values # Create a fn_mean table with same shape as fn_values\n# fn_values = fn[coor].values # To be imputed\n# mask = fn[coor].isna() # True if value is NaN\n# fn_values[mask] = fn_mean[mask] # imputation step\n# fn[coor] = fn_values\n\n# fn['filename'] = (\n#     RSNA_TRAIN + '/' \n#     + fn['study_id'].astype(str) + '/' \n#     + fn['series_id'].astype(str) + '/' \n#     + fn['instance_number'].astype(str) + '.dcm'\n# )\n# fn = fn.drop(columns = ['study_id', 'series_id', 'instance_number'])\n\n# # FInal Dataset\n# data_T1 = fn.reset_index(drop=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_T1","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def augment_image_and_centers(image, centers, center=(H/2, W/2)):\n    angle = torch.as_tensor(random.uniform(-ANGLE, ANGLE)) # -30 to 30 degress -> 0 degree meaning no rotation can also be applied\n    image = torchvision.transforms.functional.rotate(\n        image, angle.item(),\n        interpolation=torchvision.transforms.InterpolationMode.BILINEAR,\n        center=center\n    )\n\n    angle = -angle * math.pi / 180\n    s, c = torch.sin(angle), torch.cos(angle)\n    rot = torch.stack([\n        torch.stack([c, s]),\n        torch.stack([-s, c])\n    ]).to(centers.device)\n\n    center = torch.as_tensor(center, device=centers.device).float()\n    centers = ((centers - center) @ rot) + center\n\n    return image, centers\n\n\nclass Sagittal_T1(Dataset):\n    def __init__(self, df, VALID=False, alpha=0):\n        self.data = df\n        self.VALID = VALID\n        self.alpha = alpha\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, index):\n        row = self.data.iloc[index]\n        centers = torch.as_tensor([x for x in row[coor]]).view(5,2).float()\n\n        sample = f\"{RSNA_TRAIN_IMGS}/{int(row['study_id'])}/{int(row['series_id'])}/{int(row['instance_number'])}.dcm\"\n        image = pydicom.dcmread(sample).pixel_array\n        height, width = image.shape\n\n        if height > width:\n            d = width\n            h = int((height - d)*(.5 + self.alpha*(.5 - np.random.rand()))) if not self.VALID else (height - d)//2\n            image = image[h:h+d]\n            centers[:,1] -= h\n            height = width\n        elif height < width:\n            d = height\n            w = int((width - d)*(.5 + self.alpha*(.5 - np.random.rand()))) if not self.VALID else (width - d)//2\n            image = image[:,w:w+d]\n            centers[:,0] -= w\n            width = height\n\n        image = torch_resize(torch.as_tensor((image / np.max(image)).astype(np.float32)).unsqueeze(0))\n        image = image.float().to(DEVICE)\n\n        centers[:,0] = centers[:,0]*W/width\n        centers[:,1] = centers[:,1]*H/height\n\n        if not self.VALID:\n            image, centers = augment_image_and_centers(image, centers)\n\n        return image, centers","metadata":{"_uuid":"e50c77fd-12f9-40c1-9e95-e5bbe1618d71","_cell_guid":"3ea65e46-a15e-451c-9553-965bc350f366","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Visualizing the dataset**","metadata":{}},{"cell_type":"code","source":"tds = Sagittal_T1(fn)\n\nfor k in range(5):\n    image,centers = tds.__getitem__(np.random.randint(len(tds)))\n    centers = centers[centers.isnan().sum(1) == 0]\n    mask = IDX_MAP - centers.view(len(centers),2,1,1).to(DEVICE)\n    mask = (mask*mask).sum(1)\n    mask = torch.exp(A*mask)\n    mask = mask.sum(0)\n    plt.imshow(image.cpu()[0] + .5*(mask.cpu() > TH))\n    plt.show()\ndel tds","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"vds = Sagittal_T1(fn,VALID=True)\n\nfor k in range(5):\n    image,centers = vds.__getitem__(np.random.randint(len(vds)))\n    centers = centers[centers.isnan().sum(1) == 0]\n    mask = IDX_MAP - centers.view(len(centers),2,1,1).to(DEVICE)\n    mask = (mask*mask).sum(1)\n    mask = torch.exp(A*mask)\n    mask = mask.sum(0)\n    plt.imshow(image.cpu()[0] + .5*(mask.cpu() > TH))\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self, classes = K):\n        super().__init__()\n        self.classes = classes\n        self.UNet = smp.Unet(\n            encoder_name=\"resnet50\",\n            classes=classes,\n            in_channels=1\n        ).to(DEVICE)\n\n    def forward(self,X):\n        height, width = X.shape[-2:]\n        x = self.UNet(X.view(-1,1,height,width)).view(-1,height*width)\n        # Min-Max scaling of the predicted pixels\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        \n        return x.view(-1, self.classes, height, width)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class LossFunction(nn.Module):\n    def __init__(self,alpha=.5, smooth = 1e-6):\n        super().__init__()\n        self.alpha = alpha\n        self.smooth = smooth\n\n    def clone(self):\n        return LossFunction(self.alpha)\n\n    def forward(self, heatmaps, centers):\n        h, w = heatmaps.shape[-2:]\n        heatmaps = heatmaps.view(-1,h*w)\n        centers = centers.view(-1,2)\n        m = centers.isnan().sum(1) == 0\n        heatmaps = heatmaps[m]\n        centers = centers[m]\n\n        # Create ground truth heatmaps using the centers\n        mask = IDX_MAP - centers.view(len(centers),2,1,1).to(DEVICE)\n        mask = (mask*mask).sum(1)\n        mask = torch.exp(A*mask)\n        mask = mask.view(-1,h*w)\n\n        # Loss formula\n        D = 1 - ((mask*heatmaps).sum(-1))**2/((mask*mask).sum(-1)*(heatmaps*heatmaps).sum(-1)+self.smooth)\n        return D.mean()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def nt(nmin,nmax,tcur,tmax):\n    return (nmax - .5*(nmax-nmin)*(1+np.cos(tcur*np.pi/tmax))).astype(np.float32)\n\ndef cb(self):\n    alpha = torch.as_tensor(nt(.25,1,learn.train_iter,EPOCHS*n_iter))\n    learn.dls.train_ds.alpha = alpha\nalpha_cb = Callback(before_batch=cb)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import KFold\nkf = KFold(n_splits=5, shuffle=True, random_state=2003)\n\ndata_T1['fold'] = -1\nfor fold, (_, val_idx) in enumerate(kf.split(data_T1)):\n    data_T1.loc[val_idx, 'fold'] = fold","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Training**","metadata":{}},{"cell_type":"code","source":"FOLDS = [0, 1, 2, 3, 4]\n\nfor f in FOLDS:\n    model = Model()\n    tdf = data_T1[data_T1['fold'] != f]\n    vdf = data_T1[data_T1['fold'] == f]\n\n    tds = Sagittal_T1(tdf)\n    vds = Sagittal_T1(vdf, VALID = True)\n\n    tdl = DataLoader(tds, batch_size=BATCH_SIZE, shuffle=True, drop_last=True)\n    vdl = DataLoader(vds, batch_size=BATCH_SIZE, shuffle=False)\n\n    dataloaders = DataLoaders(tdl, vdl)\n\n    n_iter = len(tds) // BATCH_SIZE\n    learn = Learner(\n        dataloaders,\n        model, \n        lr = LR, \n        loss_func = LossFunction(alpha = 0.5),\n        cbs=[\n            ShowGraphCallback(),\n            alpha_cb\n        ]\n    )\n    learn.fit_one_cycle(2) # 2 epochs\n    torch.save(model, 'keypoint_detector_sagittal_T1_'+str(f))\n    del tdl, vdl, dataloaders, model, learn\n    gc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **SAGITTAL T2**","metadata":{}},{"cell_type":"code","source":"!pip install -q git+https://github.com/qubvel/segmentation_models.pytorch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:25:58.117569Z","iopub.execute_input":"2025-12-13T10:25:58.118448Z","iopub.status.idle":"2025-12-13T10:26:07.037558Z","shell.execute_reply.started":"2025-12-13T10:25:58.118422Z","shell.execute_reply":"2025-12-13T10:26:07.036579Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport pydicom\nimport numpy as np\nimport os\nimport glob\nfrom tqdm import tqdm\nimport gc\n\nimport torchvision\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset\nfrom fastai.vision.all import *\nimport segmentation_models_pytorch as smp\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:07.039228Z","iopub.execute_input":"2025-12-13T10:26:07.039524Z","iopub.status.idle":"2025-12-13T10:26:12.409857Z","shell.execute_reply.started":"2025-12-13T10:26:07.039501Z","shell.execute_reply":"2025-12-13T10:26:12.408953Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 1337\nFOLDS = [1,2,3,4,5]\nTRAIN_PATH = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\nENCODER_NAME = \"resnet18\"\nPATCH_H = 512\nPATCH_W = 512\nANGLE = 30\nS2 = 64\nBS = 16\nLR = 1e-4\nEPOCHS = 1\nTH = .5\n\nS2 = torch.as_tensor(S2)\nA = -1/(2*S2).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:12.410756Z","iopub.execute_input":"2025-12-13T10:26:12.411011Z","iopub.status.idle":"2025-12-13T10:26:12.551571Z","shell.execute_reply.started":"2025-12-13T10:26:12.410985Z","shell.execute_reply":"2025-12-13T10:26:12.550875Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_coor = pd.read_csv(TRAIN_PATH + 'train_label_coordinates.csv')\ndf_coor.sample(10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:12.553890Z","iopub.execute_input":"2025-12-13T10:26:12.554229Z","iopub.status.idle":"2025-12-13T10:26:12.628789Z","shell.execute_reply.started":"2025-12-13T10:26:12.554211Z","shell.execute_reply":"2025-12-13T10:26:12.627897Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = df_coor[\n    df_coor['condition'] == 'Spinal Canal Stenosis'\n].sort_values([\n    'study_id',\n    'series_id',\n    'level'\n]).reset_index(drop=True)\nS.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:12.629614Z","iopub.execute_input":"2025-12-13T10:26:12.629922Z","iopub.status.idle":"2025-12-13T10:26:12.648490Z","shell.execute_reply.started":"2025-12-13T10:26:12.629896Z","shell.execute_reply":"2025-12-13T10:26:12.647847Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S['x_mean_fraction'] = S['x']/S.groupby(['study_id','series_id'])['x'].mean().loc[[(study_id,series_id) for study_id,series_id in S[['study_id','series_id']].values]].values\nS.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:12.649476Z","iopub.execute_input":"2025-12-13T10:26:12.649793Z","iopub.status.idle":"2025-12-13T10:26:12.716415Z","shell.execute_reply.started":"2025-12-13T10:26:12.649764Z","shell.execute_reply":"2025-12-13T10:26:12.715652Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.boxplot(S['x_mean_fraction'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:12.717264Z","iopub.execute_input":"2025-12-13T10:26:12.717511Z","iopub.status.idle":"2025-12-13T10:26:12.848172Z","shell.execute_reply.started":"2025-12-13T10:26:12.717493Z","shell.execute_reply":"2025-12-13T10:26:12.847377Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S[S['x_mean_fraction'] < .8]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:12.848923Z","iopub.execute_input":"2025-12-13T10:26:12.849177Z","iopub.status.idle":"2025-12-13T10:26:12.862145Z","shell.execute_reply.started":"2025-12-13T10:26:12.849155Z","shell.execute_reply":"2025-12-13T10:26:12.861358Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = S[S['x_mean_fraction'] > .8]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:12.863289Z","iopub.execute_input":"2025-12-13T10:26:12.864560Z","iopub.status.idle":"2025-12-13T10:26:12.876218Z","shell.execute_reply.started":"2025-12-13T10:26:12.864522Z","shell.execute_reply":"2025-12-13T10:26:12.875416Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"coordinates = {}\nfor i in range(len(S)):\n    row = S.iloc[i]\n    coordinates[row['study_id']] = {}\nfor i in range(len(S)):\n    row = S.iloc[i]\n    coordinates[row['study_id']][row['series_id']] = {}\nfor i in range(len(S)):\n    row = S.iloc[i]\n    coordinates[row['study_id']][row['series_id']][row['instance_number']] = {\n        'L1/L2':{\n            'x':torch.nan,\n            'y':torch.nan\n        },\n        'L2/L3':{\n            'x':torch.nan,\n            'y':torch.nan\n        },\n        'L3/L4':{\n            'x':torch.nan,\n            'y':torch.nan\n        },\n        'L4/L5':{\n            'x':torch.nan,\n            'y':torch.nan\n        },\n        'L5/S1':{\n            'x':torch.nan,\n            'y':torch.nan\n        }\n    }\nfor i in range(len(S)):\n    row = S.iloc[i]\n    coordinates[row['study_id']][row['series_id']][row['instance_number']][row['level']]['x'] = row['x']\n    coordinates[row['study_id']][row['series_id']][row['instance_number']][row['level']]['y'] = row['y']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:12.879104Z","iopub.execute_input":"2025-12-13T10:26:12.879326Z","iopub.status.idle":"2025-12-13T10:26:14.646160Z","shell.execute_reply.started":"2025-12-13T10:26:12.879310Z","shell.execute_reply":"2025-12-13T10:26:14.645392Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S =  S[[\n    'study_id',\n    'series_id',\n    'instance_number'\n]].groupby([\n    'study_id',\n    'series_id',\n    'instance_number'\n]).count().reset_index()\nS.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:14.647089Z","iopub.execute_input":"2025-12-13T10:26:14.647363Z","iopub.status.idle":"2025-12-13T10:26:14.661108Z","shell.execute_reply.started":"2025-12-13T10:26:14.647341Z","shell.execute_reply":"2025-12-13T10:26:14.660284Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"v = np.zeros((len(S),10))\nfor i in tqdm(range(len(S))):\n    row = S.iloc[i]\n    k = 0\n    for level in coordinates[row['study_id']][row['series_id']][row['instance_number']]:\n        v[i,k:k+2] = list(coordinates[row['study_id']][row['series_id']][row['instance_number']][level].values())\n        k += 2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:14.661887Z","iopub.execute_input":"2025-12-13T10:26:14.662131Z","iopub.status.idle":"2025-12-13T10:26:14.847968Z","shell.execute_reply.started":"2025-12-13T10:26:14.662107Z","shell.execute_reply":"2025-12-13T10:26:14.847235Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"coor = [\n    'x_L1L2',\n    'y_L1L2',\n    'x_L2L3',\n    'y_L2L3',\n    'x_L3L4',\n    'y_L3L4',\n    'x_L4L5',\n    'y_L4L5',\n    'x_L5S1',\n    'y_L5S1'    \n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:14.848775Z","iopub.execute_input":"2025-12-13T10:26:14.848982Z","iopub.status.idle":"2025-12-13T10:26:14.853869Z","shell.execute_reply.started":"2025-12-13T10:26:14.848966Z","shell.execute_reply":"2025-12-13T10:26:14.853045Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S[coor] = v\nS.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:14.854592Z","iopub.execute_input":"2025-12-13T10:26:14.855008Z","iopub.status.idle":"2025-12-13T10:26:14.884465Z","shell.execute_reply.started":"2025-12-13T10:26:14.854991Z","shell.execute_reply":"2025-12-13T10:26:14.883660Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for (study_id,series_id),df in tqdm(S.groupby(['study_id','series_id'])):\n    sample = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/\" + str(study_id) + '/' + str(series_id)\n    instance_numbers = [int(x.replace('\\\\','/').split('/')[-1].replace('.dcm','')) for x in glob.glob(sample+'/*.dcm')]\n    instance_numbers.sort()\n    instance_numbers = np.array(instance_numbers)\n    D = len(instance_numbers)\n    L = D//3\n    FIRST = int(np.arange(D)[instance_numbers == df['instance_number'].min()])\n    LAST = int(np.arange(D)[instance_numbers == df['instance_number'].max()])\n    M = (FIRST + LAST)//2\n    START = max([0,M - L//2])\n    END = min([D,M+L-L//2+1])\n    new = instance_numbers[START:END].tolist()\n    if FIRST > 0: new.append(instance_numbers[FIRST - 1])\n    if FIRST > 1: new.append(instance_numbers[FIRST - 2])\n    if LAST < D - 1: new.append(instance_numbers[LAST + 1])\n    if LAST < D - 2: new.append(instance_numbers[LAST + 2])\n    L = len(new)\n    S = pd.concat([\n            S,\n            pd.DataFrame({\n                'study_id':[int(study_id)]*L,\n                'series_id':[int(series_id)]*L,\n                'instance_number':new,\n                'x_L1L2':[torch.nan]*L,\n                'y_L1L2':[torch.nan]*L,\n                'x_L2L3':[torch.nan]*L,\n                'y_L2L3':[torch.nan]*L,\n                'x_L3L4':[torch.nan]*L,\n                'y_L3L4':[torch.nan]*L,\n                'x_L4L5':[torch.nan]*L,\n                'y_L4L5':[torch.nan]*L,\n                'x_L5S1':[torch.nan]*L,\n                'y_L5S1':[torch.nan]*L\n            })\n        ])\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:14.885375Z","iopub.execute_input":"2025-12-13T10:26:14.885667Z","iopub.status.idle":"2025-12-13T10:26:20.443660Z","shell.execute_reply.started":"2025-12-13T10:26:14.885649Z","shell.execute_reply":"2025-12-13T10:26:20.442782Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = S.reset_index(drop=True)\nS[['study_id','series_id','instance_number']] = S[['study_id','series_id','instance_number']].astype(np.int64)\nS.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.444469Z","iopub.execute_input":"2025-12-13T10:26:20.444826Z","iopub.status.idle":"2025-12-13T10:26:20.469146Z","shell.execute_reply.started":"2025-12-13T10:26:20.444784Z","shell.execute_reply":"2025-12-13T10:26:20.468425Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S_mean = S.groupby(['study_id','series_id']).mean()\nS_mean.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.470040Z","iopub.execute_input":"2025-12-13T10:26:20.470294Z","iopub.status.idle":"2025-12-13T10:26:20.496320Z","shell.execute_reply.started":"2025-12-13T10:26:20.470274Z","shell.execute_reply":"2025-12-13T10:26:20.495588Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S_mean = S_mean.loc[[(study_id,series_id) for study_id,series_id in S[['study_id','series_id']].values]][coor].values\nS_mean.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.497977Z","iopub.execute_input":"2025-12-13T10:26:20.498269Z","iopub.status.idle":"2025-12-13T10:26:20.589280Z","shell.execute_reply.started":"2025-12-13T10:26:20.498240Z","shell.execute_reply":"2025-12-13T10:26:20.588495Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S_values = S[coor].values\nS_values.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.590071Z","iopub.execute_input":"2025-12-13T10:26:20.590352Z","iopub.status.idle":"2025-12-13T10:26:20.596768Z","shell.execute_reply.started":"2025-12-13T10:26:20.590329Z","shell.execute_reply":"2025-12-13T10:26:20.596044Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mask = S[coor].isna()\nmask.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.597700Z","iopub.execute_input":"2025-12-13T10:26:20.598046Z","iopub.status.idle":"2025-12-13T10:26:20.612497Z","shell.execute_reply.started":"2025-12-13T10:26:20.598025Z","shell.execute_reply":"2025-12-13T10:26:20.611771Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## This only imputes the NaN values determined by the mask, with the mean\n\nS_values[mask] = S_mean[mask]\nS[coor] = S_values\nS.tail(10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.613984Z","iopub.execute_input":"2025-12-13T10:26:20.614244Z","iopub.status.idle":"2025-12-13T10:26:20.637737Z","shell.execute_reply.started":"2025-12-13T10:26:20.614224Z","shell.execute_reply":"2025-12-13T10:26:20.636991Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S[S.isna().sum(1) > 0].reset_index(drop=True).tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.638624Z","iopub.execute_input":"2025-12-13T10:26:20.638905Z","iopub.status.idle":"2025-12-13T10:26:20.663240Z","shell.execute_reply.started":"2025-12-13T10:26:20.638883Z","shell.execute_reply":"2025-12-13T10:26:20.662588Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_meta_f = pd.read_csv(TRAIN_PATH + 'train_series_descriptions.csv')\ndf_meta_f.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.664253Z","iopub.execute_input":"2025-12-13T10:26:20.664588Z","iopub.status.idle":"2025-12-13T10:26:20.678886Z","shell.execute_reply.started":"2025-12-13T10:26:20.664560Z","shell.execute_reply":"2025-12-13T10:26:20.678201Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = S.merge(df_meta_f[['series_id','series_description']], left_on='series_id', right_on='series_id')\nS.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.679796Z","iopub.execute_input":"2025-12-13T10:26:20.680071Z","iopub.status.idle":"2025-12-13T10:26:20.708314Z","shell.execute_reply.started":"2025-12-13T10:26:20.680050Z","shell.execute_reply":"2025-12-13T10:26:20.707601Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S.groupby('series_description').count()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.709300Z","iopub.execute_input":"2025-12-13T10:26:20.709851Z","iopub.status.idle":"2025-12-13T10:26:20.727980Z","shell.execute_reply.started":"2025-12-13T10:26:20.709825Z","shell.execute_reply":"2025-12-13T10:26:20.727287Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = S[S.series_description == 'Sagittal T2/STIR'].reset_index(drop=True)\nS.groupby('series_description').count()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.728881Z","iopub.execute_input":"2025-12-13T10:26:20.729311Z","iopub.status.idle":"2025-12-13T10:26:20.747771Z","shell.execute_reply.started":"2025-12-13T10:26:20.729287Z","shell.execute_reply":"2025-12-13T10:26:20.746911Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## For the below code, we need a precomputed fold data - train.csv which contains the final classifications\n# S = S.merge(train[['study_id','fold']],left_on='study_id',right_on='study_id')\n# S.tail()\n\nS['fold'] = np.random.randint(1, 6, size=len(S))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.748716Z","iopub.execute_input":"2025-12-13T10:26:20.748970Z","iopub.status.idle":"2025-12-13T10:26:20.754468Z","shell.execute_reply.started":"2025-12-13T10:26:20.748953Z","shell.execute_reply":"2025-12-13T10:26:20.753488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S.groupby('fold').count()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.758132Z","iopub.execute_input":"2025-12-13T10:26:20.758444Z","iopub.status.idle":"2025-12-13T10:26:20.777488Z","shell.execute_reply.started":"2025-12-13T10:26:20.758415Z","shell.execute_reply":"2025-12-13T10:26:20.776765Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def augment_image_and_centers(image,centers,center=(PATCH_H/2,PATCH_W/2)):\n    angle = torch.as_tensor(random.uniform(-ANGLE, ANGLE))\n    image = torchvision.transforms.functional.rotate(\n        image,angle.item(),\n        interpolation=torchvision.transforms.InterpolationMode.BILINEAR,\n        center=center\n    )\n    angle = -angle*math.pi/180\n    s = torch.sin(angle)\n    c = torch.cos(angle)\n    rot = torch.stack([\n        torch.stack([c, s]),\n        torch.stack([-s, c])\n    ])\n    center = torch.as_tensor(center).float()\n    centers = ((centers.cpu() - center) @ rot) + center\n\n    return image,centers\ntorch_resize = torchvision.transforms.Resize((PATCH_H,PATCH_W),antialias=True)\n\nx_map = torch.stack([torch.arange(PATCH_W)]*PATCH_H).float()\ny_map = torch.stack([torch.arange(PATCH_H)]*PATCH_W).float()\nidx_map = torch.stack([x_map,y_map.T]).view(1,2,PATCH_H,PATCH_W).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.778316Z","iopub.execute_input":"2025-12-13T10:26:20.778539Z","iopub.status.idle":"2025-12-13T10:26:20.794947Z","shell.execute_reply.started":"2025-12-13T10:26:20.778522Z","shell.execute_reply":"2025-12-13T10:26:20.794051Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Sagittal_T2_Dataset(Dataset):\n    def __init__(self, df, VALID=False, alpha=0):\n        self.data = df\n        self.VALID = VALID\n        self.alpha = alpha\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, index):\n        row = self.data.iloc[index]\n\n        centers = torch.as_tensor([x for x in row[coor]]).view(5,2).float()\n        \n        sample = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/\" + str(row['study_id']) + '/' + str(row['series_id']) + '/'+str(row['instance_number']) + '.dcm'\n        \n        image = pydicom.dcmread(sample).pixel_array\n        H,W = image.shape\n        # Perform image crops\n        if H > W:\n            d = W\n            if not self.VALID:\n                h = int((H - d)*(.5 + self.alpha*(.5 - np.random.rand())))\n            else:\n                h = (H - d)//2\n            image = image[h:h+d]\n            centers[:,1] -= h\n            H = W\n        elif H < W:\n            d = H\n            if not self.VALID:\n                w = int((W - d)*(.5 + self.alpha*(.5 - np.random.rand())))\n            else:\n                w = (W - d)//2\n            image = image[:,w:w+d]\n            centers[:,0] -= w\n            W = H\n        image = torch_resize(torch.as_tensor((image/np.max(image)).astype(np.float32)).unsqueeze(0))\n        image = image.float().to(device)\n        \n        centers[:,0] = centers[:,0]*PATCH_W/W\n        centers[:,1] = centers[:,1]*PATCH_H/H\n\n        if not self.VALID: image,centers = augment_image_and_centers(image,centers)\n\n        return image,centers","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.796032Z","iopub.execute_input":"2025-12-13T10:26:20.797701Z","iopub.status.idle":"2025-12-13T10:26:20.812277Z","shell.execute_reply.started":"2025-12-13T10:26:20.797673Z","shell.execute_reply":"2025-12-13T10:26:20.811527Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tds = Sagittal_T2_Dataset(S)\nfor k in range(5):\n    image,centers = tds.__getitem__(np.random.randint(len(tds)))\n    centers = centers[centers.isnan().sum(1) == 0]\n    mask = idx_map - centers.view(len(centers),2,1,1).to(device)\n    mask = (mask*mask).sum(1)\n    mask = torch.exp(A*mask)\n    mask = mask.sum(0)\n    plt.imshow(image.cpu()[0] + .5*(mask.cpu() > TH))\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:20.813123Z","iopub.execute_input":"2025-12-13T10:26:20.813419Z","iopub.status.idle":"2025-12-13T10:26:22.044399Z","shell.execute_reply.started":"2025-12-13T10:26:20.813394Z","shell.execute_reply":"2025-12-13T10:26:22.043454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del tds\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:22.045437Z","iopub.execute_input":"2025-12-13T10:26:22.045742Z","iopub.status.idle":"2025-12-13T10:26:22.287759Z","shell.execute_reply.started":"2025-12-13T10:26:22.045719Z","shell.execute_reply":"2025-12-13T10:26:22.286993Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:22.288609Z","iopub.execute_input":"2025-12-13T10:26:22.288886Z","iopub.status.idle":"2025-12-13T10:26:22.301302Z","shell.execute_reply.started":"2025-12-13T10:26:22.288860Z","shell.execute_reply":"2025-12-13T10:26:22.300525Z"}},"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        # Min-Max normalization\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        \n        return x.view(-1,self.classes,H,W)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:22.302364Z","iopub.execute_input":"2025-12-13T10:26:22.302879Z","iopub.status.idle":"2025-12-13T10:26:22.321689Z","shell.execute_reply.started":"2025-12-13T10:26:22.302853Z","shell.execute_reply":"2025-12-13T10:26:22.320913Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myLoss(nn.Module):\n    def __init__(\n            self,\n            alpha=.5,\n            smooth = 1e-6\n        ):\n        super().__init__()\n        self.alpha = alpha\n        self.smooth = smooth\n\n    def clone(self):\n        return myLoss(self.alpha)\n\n    def forward(\n            self,\n            heatmaps,# Predictions\n            centers # Targets\n        ):\n        H,W = heatmaps.shape[-2:]\n        heatmaps = heatmaps.view(-1,H*W)\n        centers = centers.view(-1,2)\n        m = centers.isnan().sum(1) == 0\n        heatmaps = heatmaps[m]\n        centers = centers[m]\n#       Ideal heatmaps\n        mask = idx_map - centers.view(len(centers),2,1,1).to(device)\n        mask = (mask*mask).sum(1)\n        mask = torch.exp(A*mask)\n        mask = mask.view(-1,H*W)\n#       Distance\n        D = 1 - ((mask*heatmaps).sum(-1))**2/((mask*mask).sum(-1)*(heatmaps*heatmaps).sum(-1)+self.smooth)\n        return D.mean()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:22.323794Z","iopub.execute_input":"2025-12-13T10:26:22.324030Z","iopub.status.idle":"2025-12-13T10:26:22.332740Z","shell.execute_reply.started":"2025-12-13T10:26:22.324013Z","shell.execute_reply":"2025-12-13T10:26:22.332019Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CosineAnnealingAlpha\ndef nt(nmin,nmax,tcur,tmax):\n    return (nmax - .5*(nmax-nmin)*(1+np.cos(tcur*np.pi/tmax))).astype(np.float32)\n\ndef cb(self):\n    alpha = torch.as_tensor(nt(.25,1,learn.train_iter,EPOCHS*n_iter))\n    learn.dls.train_ds.alpha = alpha\nalpha_cb = Callback(before_batch=cb)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:22.333521Z","iopub.execute_input":"2025-12-13T10:26:22.333818Z","iopub.status.idle":"2025-12-13T10:26:22.342435Z","shell.execute_reply.started":"2025-12-13T10:26:22.333797Z","shell.execute_reply":"2025-12-13T10:26:22.341692Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for f in FOLDS:\n    seed_everything(SEED)\n    model = myUNet(5)\n    # model = torch.load(PATH + 'Sagittal_T1/level_segmentation/Sagittal_T1_sagittal_level_segmentation_'+str(f))\n    \n    tdf = S[S['fold'] != f]\n    vdf = S[S['fold'] == f]\n\n    tds = Sagittal_T2_Dataset(tdf)\n    vds = Sagittal_T2_Dataset(vdf,VALID=True)\n    \n    tdl = torch.utils.data.DataLoader(tds, batch_size=BS, shuffle=True, drop_last=True)\n    vdl = torch.utils.data.DataLoader(vds, batch_size=BS, shuffle=False)\n\n    dls = DataLoaders(tdl,vdl)\n\n    n_iter = len(tds)//BS\n\n    learn = Learner(\n        dls,\n        model,\n        lr=LR,\n        loss_func=myLoss(alpha=0.5),\n        cbs=[\n            ShowGraphCallback(),\n            alpha_cb\n        ]\n    )\n    learn.fit_one_cycle(3)\n    torch.save(model,'Sagittal_T2_sagittal_level_segmentation_'+str(f))\n    del tdl,vdl,dls,model,learn\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:26:27.666048Z","iopub.execute_input":"2025-12-13T10:26:27.666717Z","iopub.status.idle":"2025-12-13T11:50:40.038285Z","shell.execute_reply.started":"2025-12-13T10:26:27.666688Z","shell.execute_reply":"2025-12-13T11:50:40.036859Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Sample model testing","metadata":{}},{"cell_type":"code","source":"model = torch.load(\"/kaggle/working/Sagittal_T2_sagittal_level_segmentation_2\", map_location='cuda', weights_only = False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T11:51:36.163917Z","iopub.execute_input":"2025-12-13T11:51:36.164479Z","iopub.status.idle":"2025-12-13T11:51:36.238424Z","shell.execute_reply.started":"2025-12-13T11:51:36.164455Z","shell.execute_reply":"2025-12-13T11:51:36.237796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\nimport random\n\nidx = random.randint(0, len(vds) - 1)\nimage, gt_centers = vds[idx]\n\nwith torch.no_grad():\n    pred_heatmaps = model(image.unsqueeze(0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T11:54:32.276029Z","iopub.execute_input":"2025-12-13T11:54:32.276698Z","iopub.status.idle":"2025-12-13T11:54:32.322500Z","shell.execute_reply.started":"2025-12-13T11:54:32.276675Z","shell.execute_reply":"2025-12-13T11:54:32.321897Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def heatmaps_to_keypoints(heatmaps):\n    \"\"\"\n    heatmaps: Tensor [1, C, H, W]\n    returns: list of (x, y)\n    \"\"\"\n    heatmaps = heatmaps[0].cpu().numpy()\n    coords = []\n\n    for c in range(heatmaps.shape[0]):\n        y, x = np.unravel_index(\n            np.argmax(heatmaps[c]),\n            heatmaps[c].shape\n        )\n        coords.append((x, y))\n\n    return coords\n\npred_centers = heatmaps_to_keypoints(pred_heatmaps)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T11:54:34.329403Z","iopub.execute_input":"2025-12-13T11:54:34.329954Z","iopub.status.idle":"2025-12-13T11:54:34.336075Z","shell.execute_reply.started":"2025-12-13T11:54:34.329930Z","shell.execute_reply":"2025-12-13T11:54:34.335423Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nimg = image.cpu()[0]  # [H, W]\n\nplt.figure(figsize=(6, 6))\nplt.imshow(img, cmap='gray')\nplt.axis('off')\n\n# Ground truth (green)\ngt = gt_centers[gt_centers.isnan().sum(1) == 0]\nfor (x, y) in gt:\n    plt.scatter(x, y, c='lime', s=40, label='GT')\n\n# Prediction (red)\nfor (x, y) in pred_centers:\n    plt.scatter(x, y, c='red', s=40, label='Pred')\n\nplt.title(\"Sagittal T2 – GT (green) vs Prediction (red)\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T11:54:44.021228Z","iopub.execute_input":"2025-12-13T11:54:44.021716Z","iopub.status.idle":"2025-12-13T11:54:44.286753Z","shell.execute_reply.started":"2025-12-13T11:54:44.021693Z","shell.execute_reply":"2025-12-13T11:54:44.285798Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **AXIAL T2**","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport pydicom\nimport numpy as np\nimport os\nimport glob\nfrom tqdm import tqdm\nimport gc\n\nimport torchvision\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset\nfrom fastai.vision.all import *\nimport segmentation_models_pytorch as smp\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 1337\nFOLDS = [1,2,3,4,5]\nPATH = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\nTRAIN_PATH = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/'\nENCODER_NAME = \"resnet18\"\nANGLE = 180\nS2 = 64\nPATCH_H = 512\nPATCH_W = 512\nBS = 24\nLR = 2.5e-4\nEPOCHS = 1\nTH = .5\n\nS2 = torch.as_tensor(S2)\nA = -1/(2*S2).to(device)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_coor = pd.read_csv(PATH + 'train_label_coordinates.csv')\ndf_coor.tail()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = df_coor[\n    df_coor['condition'].isin([\n        'Left Subarticular Stenosis',\n        'Right Subarticular Stenosis'\n    ])\n].sort_values([\n    'study_id',\n    'series_id',\n    'level',\n    'condition'\n]).reset_index(drop=True)\nS.tail()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"centers = {}\nfor i in range(len(S)):\n    row = S.iloc[i]\n    centers[row['study_id']]={}\nfor i in range(len(S)):\n    row = S.iloc[i]\n    centers[row['study_id']][row['series_id']]={}\nfor i in range(len(S)):\n    row = S.iloc[i]\n    centers[row['study_id']][row['series_id']][row['instance_number']]={'L':[], 'R':[], 'level':[]}\nfor i in range(len(S)):\n    row = S.iloc[i]\n    centers[row['study_id']][row['series_id']][row['instance_number']][row['condition'][0]].append([row['x'],row['y']])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = S[[\n    'study_id',\n    'series_id',\n    'instance_number'\n]].groupby([\n    'study_id',\n    'series_id',\n    'instance_number'\n]).count().reset_index()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"coordinates = np.zeros((len(S),4))\ncoordinates[:] = np.nan\nfor i in range(len(S)):\n    row = S.iloc[i]\n    for side in ['L','R']:\n            if len(centers[row['study_id']][row['series_id']][row['instance_number']][side]) > 0:\n                c = np.array(centers[row['study_id']][row['series_id']][row['instance_number']][side]).mean(0)\n                coordinates[\n                    i,\n                    {'L':0, 'R':2}[side]:{'L':0, 'R':2}[side]+2\n                ] = c","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"coor = [\n    'x_L',\n    'y_L',\n    'x_R',\n    'y_R'    \n]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S[coor] = coordinates\nS.tail()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S[S[coor].isna().sum(1) > 0].reset_index(drop=True).tail()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# I suppose those are centered inside levels and I'll can safely assign neighbor slices to the same level\ncentroids = S[S[coor].isna().sum(1) == 0].reset_index(drop=True)\ncentroids.tail()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for (study_id,series_id),df in tqdm(centroids.groupby(['study_id','series_id'])):\n    sample = TRAIN_PATH + str(study_id) + '/' + str(series_id)\n    instance_numbers = [int(x.replace('\\\\','/').split('/')[-1].replace('.dcm','')) for x in glob.glob(sample+'/*.dcm')]\n    instance_numbers.sort()\n    instance_numbers = np.array(instance_numbers)\n    D = len(instance_numbers)\n    for i in range(len(df)):\n        row = df[i:i+1]\n        instance_number = row['instance_number'].values\n        instance_number_index = int(np.arange(D)[instance_numbers == instance_number])\n        if instance_number_index > 0:\n            new = row.copy()\n            new['instance_number'] = instance_numbers[instance_number_index - 1]\n            centroids = pd.concat([\n                centroids,\n                new\n            ])\n        if instance_number_index < D - 1:\n            new = row.copy()\n            new['instance_number'] = instance_numbers[instance_number_index + 1]\n            centroids = pd.concat([\n                centroids,\n                new\n            ])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = pd.concat([S[S[coor].isna().sum(1) > 0],centroids]).groupby(['study_id','series_id','instance_number']).mean().reset_index()\nS.tail()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S[S[coor].isna().sum(1) > 0].reset_index(drop=True).tail()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S['flip'] = False\nfS = S.copy()\nfS['flip'] = True\nS = pd.concat([S,fS]).reset_index(drop=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S['fold'] = np.random.randint(1, 6, size=len(S))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_meta_f = pd.read_csv(PATH + 'train_series_descriptions.csv')\ndf_meta_f.tail()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S = S.merge(df_meta_f[['series_id','series_description']], left_on='series_id', right_on='series_id')\nS.tail()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S.groupby('series_description').count()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"S.groupby('fold').count()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def augment_image_and_centers(image,centers,center=(PATCH_H/2,PATCH_W/2)):\n    # Randomly rotate the image.\n    angle = torch.as_tensor(random.uniform(-ANGLE, ANGLE))\n    image = torchvision.transforms.functional.rotate(\n        image,angle.item(),\n        interpolation=torchvision.transforms.InterpolationMode.BILINEAR,\n        center=center\n    )\n    angle = -angle*math.pi/180\n    s = torch.sin(angle)\n    c = torch.cos(angle)\n    rot = torch.stack([\n        torch.stack([c, s]),\n        torch.stack([-s, c])\n    ])\n    center = torch.as_tensor(center).float()\n    centers = ((centers.cpu() - center) @ rot) + center\n\n    return image,centers\n\ntorch_resize = torchvision.transforms.Resize((PATCH_H,PATCH_W),antialias=True)\n\nx_map = torch.stack([torch.arange(PATCH_W)]*PATCH_H).float()\ny_map = torch.stack([torch.arange(PATCH_H)]*PATCH_W).float()\nidx_map = torch.stack([x_map,y_map.T]).view(1,2,PATCH_H,PATCH_W).to(device)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Axial_T2_axial_side_Dataset(Dataset):\n    def __init__(self, df, VALID=False, alpha=0):\n        self.data = df\n        self.VALID = VALID\n        self.alpha = alpha\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, index):\n        row = self.data.iloc[index]\n\n        centers = torch.as_tensor([x for x in row[coor]]).view(2,2).float().to(device)\n        \n        sample = TRAIN_PATH + str(row['study_id']) + '/'+ str(row['series_id']) + '/'+ str(row['instance_number']) + '.dcm'\n        \n        image = pydicom.dcmread(sample).pixel_array\n        H,W = image.shape\n        if H > W:\n            d = W\n            if not self.VALID:\n                h = int((H - d)*(.5 + self.alpha*(.5 - np.random.rand())))\n            else:\n                h = (H - d)//2\n            image = image[h:h+d]\n            centers[:,1] -= h\n            H = W\n        elif H < W:\n            d = H\n            if not self.VALID:\n                w = int((W - d)*(.5 + self.alpha*(.5 - np.random.rand())))\n            else:\n                w = (W - d)//2\n            image = image[:,w:w+d]\n            centers[:,0] -= w\n            W = H\n        image = torch_resize(torch.as_tensor((image/np.max(image)).astype(np.float32)).unsqueeze(0))\n        image = image.float().to(device)\n        \n        centers[:,0] = centers[:,0]*PATCH_W/W\n        centers[:,1] = centers[:,1]*PATCH_H/H\n\n        if not self.VALID: image,centers = augment_image_and_centers(image,centers,centers.nanmean(0).tolist())\n\n        if row['flip']:\n            image = image.flip(2)\n            centers = centers.flip(0)\n            centers[:,0] = PATCH_W - centers[:,0] - 1\n\n        return image,centers","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tds = Axial_T2_axial_side_Dataset(S)\nfor k in range(10):\n    image,centers = tds.__getitem__(np.random.randint(len(tds)))\n    centers = centers[centers.isnan().sum(1) == 0]\n    mask = idx_map - centers.view(len(centers),2,1,1).to(device)\n    mask = (mask*mask).sum(1)\n    mask = torch.exp(A*mask)\n    mask = mask.sum(0)\n    plt.imshow(image.cpu()[0] + .5*(mask.cpu() > TH))\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del tds\ngc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myUNet(nn.Module):\n    def __init__(\n        self,\n        classes\n        ):\n        super(myUNet, self).__init__()\n\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        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        \n        return x.view(-1,self.classes,H,W)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myLoss(nn.Module):\n    def __init__(\n            self,\n            alpha=.5,\n            smooth = 1e-6\n        ):\n        super().__init__()\n        self.alpha = alpha\n        self.smooth = smooth\n\n    def clone(self):\n        return myLoss(self.alpha)\n\n    def forward(\n            self,\n            heatmaps,# Predictions\n            centers # Targets\n        ):\n        H,W = heatmaps.shape[-2:]\n        heatmaps = heatmaps.view(-1,H*W)\n        centers = centers.view(-1,2)\n        m = centers.isnan().sum(1) == 0\n        heatmaps = heatmaps[m]\n        centers = centers[m]\n        # Ideal heatmaps\n        mask = idx_map - centers.view(len(centers),2,1,1).to(device)\n        mask = (mask*mask).sum(1)\n        mask = torch.exp(A*mask)\n        mask = mask.view(-1,H*W)\n        # loss\n        D = 1 - ((mask*heatmaps).sum(-1))**2/((mask*mask).sum(-1)*(heatmaps*heatmaps).sum(-1)+self.smooth)\n        \n        return D.mean()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CosineAnnealingAlpha\ndef nt(nmin,nmax,tcur,tmax):\n    return (nmax - .5*(nmax-nmin)*(1+np.cos(tcur*np.pi/tmax))).astype(np.float32)\n\n# callback to update alpha during training\ndef cb(self):\n    alpha = torch.as_tensor(nt(.25,1,learn.train_iter,EPOCHS*n_iter))\n    learn.dls.train_ds.alpha = alpha\nalpha_cb = Callback(before_batch=cb)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for f in FOLDS:\n    seed_everything(SEED)\n    model = myUNet(2)\n    \n    tdf = S[S['fold'] != f]\n    vdf = S[S['fold'] == f]\n\n    tds = Axial_T2_axial_side_Dataset(tdf)\n    vds = Axial_T2_axial_side_Dataset(vdf,VALID=True)\n    \n    tdl = torch.utils.data.DataLoader(tds, batch_size=BS, shuffle=True, drop_last=True)\n    vdl = torch.utils.data.DataLoader(vds, batch_size=BS, shuffle=False)\n\n    dls = DataLoaders(tdl,vdl)\n\n    n_iter = len(tds)//BS\n\n    learn = Learner(\n        dls,\n        model,\n        lr=LR,\n        loss_func=myLoss(alpha=0.5),\n        cbs=[\n            ShowGraphCallback(),\n            alpha_cb\n        ]\n    )\n    learn.fit_one_cycle(2)\n    torch.save(model,'Axial_T2_axial_side_segmentation_'+str(f))\n    del tdl,vdl,dls,model,learn\n    gc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}