{"metadata":{"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9539917,"sourceType":"datasetVersion","datasetId":5726703}],"dockerImageVersionId":30787,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"papermill":{"default_parameters":{},"duration":17947.044951,"end_time":"2024-07-21T19:59:57.358595","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-07-21T15:00:50.313644","version":"2.5.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"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\n!pip install segmentation_models_pytorch\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":{"papermill":{"duration":31.433451,"end_time":"2024-07-21T15:01:24.463598","exception":false,"start_time":"2024-07-21T15:00:53.030147","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-10-24T11:46:30.846186Z","iopub.execute_input":"2024-10-24T11:46:30.847106Z","iopub.status.idle":"2024-10-24T11:46:42.973671Z","shell.execute_reply.started":"2024-10-24T11:46:30.847052Z","shell.execute_reply":"2024-10-24T11:46:42.972503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 1337\nFOLDS = [1,2,3,4,5]\nPATH = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'# Main path\nTRAIN_PATH = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/'# Training images folder\nENCODER_NAME = \"resnet18\"\nPATCH_H = 512\nPATCH_W = 512\nANGLE = 30\nS2 = 64\nBS = 16\nLR = 5e-4\nEPOCHS = 2\nTH = .5","metadata":{"papermill":{"duration":0.021232,"end_time":"2024-07-21T15:01:24.498504","exception":false,"start_time":"2024-07-21T15:01:24.477272","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-10-24T11:46:42.976624Z","iopub.execute_input":"2024-10-24T11:46:42.977094Z","iopub.status.idle":"2024-10-24T11:46:42.984470Z","shell.execute_reply.started":"2024-10-24T11:46:42.977043Z","shell.execute_reply":"2024-10-24T11:46:42.983406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S2 = torch.as_tensor(S2)\nA = -1/(2*S2).to(device)","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:42.985782Z","iopub.execute_input":"2024-10-24T11:46:42.986144Z","iopub.status.idle":"2024-10-24T11:46:42.996376Z","shell.execute_reply.started":"2024-10-24T11:46:42.986101Z","shell.execute_reply":"2024-10-24T11:46:42.995426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/sagittal-t1/train_split.csv')\ntrain.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:42.997739Z","iopub.execute_input":"2024-10-24T11:46:42.998130Z","iopub.status.idle":"2024-10-24T11:46:43.043644Z","shell.execute_reply.started":"2024-10-24T11:46:42.998087Z","shell.execute_reply":"2024-10-24T11:46:43.042665Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_coor = pd.read_csv(PATH + 'train_label_coordinates.csv')\ndf_coor.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:43.046313Z","iopub.execute_input":"2024-10-24T11:46:43.046665Z","iopub.status.idle":"2024-10-24T11:46:43.141208Z","shell.execute_reply.started":"2024-10-24T11:46:43.046627Z","shell.execute_reply":"2024-10-24T11:46:43.140028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F = df_coor[\n    df_coor['condition'].isin([\n        'Left Neural Foraminal Narrowing',\n        'Right Neural Foraminal Narrowing'\n    ])\n].drop(columns=['condition']).sort_values([\n    'study_id',\n    'series_id',\n    'level'\n]).reset_index(drop=True)\nF.tail()","metadata":{"papermill":{"duration":0.051267,"end_time":"2024-07-21T15:01:24.918000","exception":false,"start_time":"2024-07-21T15:01:24.866733","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-10-24T11:46:43.142519Z","iopub.execute_input":"2024-10-24T11:46:43.142940Z","iopub.status.idle":"2024-10-24T11:46:43.168935Z","shell.execute_reply.started":"2024-10-24T11:46:43.142902Z","shell.execute_reply":"2024-10-24T11:46:43.167994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coordinates = {}\nfor i in range(len(F)):\n    row = F.iloc[i]\n    coordinates[row['study_id']] = {}\nfor i in range(len(F)):\n    row = F.iloc[i]\n    coordinates[row['study_id']][row['series_id']] = {}\nfor i in range(len(F)):\n    row = F.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(F)):\n    row = F.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":{"papermill":{"duration":2.138078,"end_time":"2024-07-21T15:01:27.440513","exception":false,"start_time":"2024-07-21T15:01:25.302435","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-10-24T11:46:43.170166Z","iopub.execute_input":"2024-10-24T11:46:43.170552Z","iopub.status.idle":"2024-10-24T11:46:49.043222Z","shell.execute_reply.started":"2024-10-24T11:46:43.170509Z","shell.execute_reply":"2024-10-24T11:46:49.042202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F =  F[[\n    'study_id',\n    'series_id',\n    'instance_number'\n]].groupby([\n    'study_id',\n    'series_id',\n    'instance_number'\n]).count().reset_index()\nF.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:49.044304Z","iopub.execute_input":"2024-10-24T11:46:49.044622Z","iopub.status.idle":"2024-10-24T11:46:49.063538Z","shell.execute_reply.started":"2024-10-24T11:46:49.044589Z","shell.execute_reply":"2024-10-24T11:46:49.062634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"v = np.zeros((len(F),10))\nfor i in tqdm(range(len(F))):\n    row = F.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":{"papermill":{"duration":2.758526,"end_time":"2024-07-21T15:01:30.215104","exception":false,"start_time":"2024-07-21T15:01:27.456578","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-10-24T11:46:49.064895Z","iopub.execute_input":"2024-10-24T11:46:49.065280Z","iopub.status.idle":"2024-10-24T11:46:50.258666Z","shell.execute_reply.started":"2024-10-24T11:46:49.065238Z","shell.execute_reply":"2024-10-24T11:46:50.257735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-24T11:46:50.259975Z","iopub.execute_input":"2024-10-24T11:46:50.260287Z","iopub.status.idle":"2024-10-24T11:46:50.264892Z","shell.execute_reply.started":"2024-10-24T11:46:50.260253Z","shell.execute_reply":"2024-10-24T11:46:50.263955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F[coor] = v\nF.tail()","metadata":{"papermill":{"duration":0.050977,"end_time":"2024-07-21T15:01:30.282179","exception":false,"start_time":"2024-07-21T15:01:30.231202","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-10-24T11:46:50.266153Z","iopub.execute_input":"2024-10-24T11:46:50.266469Z","iopub.status.idle":"2024-10-24T11:46:50.292300Z","shell.execute_reply.started":"2024-10-24T11:46:50.266420Z","shell.execute_reply":"2024-10-24T11:46:50.291361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for (study_id,series_id),df in tqdm(F.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    FIRST = int(np.arange(D)[instance_numbers == df['instance_number'].min()])\n    LAST = int(np.arange(D)[instance_numbers == df['instance_number'].max()])\n    new = instance_numbers[FIRST+1:LAST].tolist()\n    if FIRST > 0: new.append(instance_numbers[FIRST - 1])\n    if LAST < D - 1: new.append(instance_numbers[LAST + 1])\n    L = len(new)\n    F = pd.concat([\n        F,\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":{"execution":{"iopub.status.busy":"2024-10-24T11:46:50.293688Z","iopub.execute_input":"2024-10-24T11:46:50.294034Z","iopub.status.idle":"2024-10-24T11:46:57.629007Z","shell.execute_reply.started":"2024-10-24T11:46:50.293997Z","shell.execute_reply":"2024-10-24T11:46:57.627833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F = F.reset_index(drop=True)\nF[['study_id','series_id','instance_number']] = F[['study_id','series_id','instance_number']].astype(np.int64)\nF.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:57.630542Z","iopub.execute_input":"2024-10-24T11:46:57.630965Z","iopub.status.idle":"2024-10-24T11:46:57.658940Z","shell.execute_reply.started":"2024-10-24T11:46:57.630920Z","shell.execute_reply":"2024-10-24T11:46:57.657827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F_mean = F.groupby(['study_id','series_id']).mean()\nF_mean.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:57.663673Z","iopub.execute_input":"2024-10-24T11:46:57.664078Z","iopub.status.idle":"2024-10-24T11:46:57.689061Z","shell.execute_reply.started":"2024-10-24T11:46:57.664040Z","shell.execute_reply":"2024-10-24T11:46:57.688059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F_mean = F_mean.loc[[(study_id,series_id) for study_id,series_id in F[['study_id','series_id']].values]][coor].values\nF_mean.shape","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:57.690411Z","iopub.execute_input":"2024-10-24T11:46:57.691053Z","iopub.status.idle":"2024-10-24T11:46:57.863946Z","shell.execute_reply.started":"2024-10-24T11:46:57.691009Z","shell.execute_reply":"2024-10-24T11:46:57.862810Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F_values = F[coor].values\nF_values.shape","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:57.865573Z","iopub.execute_input":"2024-10-24T11:46:57.866127Z","iopub.status.idle":"2024-10-24T11:46:57.874578Z","shell.execute_reply.started":"2024-10-24T11:46:57.866077Z","shell.execute_reply":"2024-10-24T11:46:57.873468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask = F[coor].isna()\nmask.shape","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:57.876152Z","iopub.execute_input":"2024-10-24T11:46:57.877179Z","iopub.status.idle":"2024-10-24T11:46:57.887131Z","shell.execute_reply.started":"2024-10-24T11:46:57.877135Z","shell.execute_reply":"2024-10-24T11:46:57.886219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F_values[mask] = F_mean[mask]\nF[coor] = F_values\nF.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:57.888963Z","iopub.execute_input":"2024-10-24T11:46:57.889406Z","iopub.status.idle":"2024-10-24T11:46:57.913772Z","shell.execute_reply.started":"2024-10-24T11:46:57.889349Z","shell.execute_reply":"2024-10-24T11:46:57.912570Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F[F.isna().sum(1) > 0]","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:57.915231Z","iopub.execute_input":"2024-10-24T11:46:57.916331Z","iopub.status.idle":"2024-10-24T11:46:57.936818Z","shell.execute_reply.started":"2024-10-24T11:46:57.916278Z","shell.execute_reply":"2024-10-24T11:46:57.935801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_meta_f = pd.read_csv(PATH + 'train_series_descriptions.csv')\ndf_meta_f.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:57.938405Z","iopub.execute_input":"2024-10-24T11:46:57.939115Z","iopub.status.idle":"2024-10-24T11:46:57.959532Z","shell.execute_reply.started":"2024-10-24T11:46:57.939066Z","shell.execute_reply":"2024-10-24T11:46:57.958641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F = F.merge(df_meta_f[['series_id','series_description']], left_on='series_id', right_on='series_id')\nF.tail()","metadata":{"papermill":{"duration":0.097784,"end_time":"2024-07-21T15:01:30.669426","exception":false,"start_time":"2024-07-21T15:01:30.571642","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-10-24T11:46:57.961341Z","iopub.execute_input":"2024-10-24T11:46:57.961881Z","iopub.status.idle":"2024-10-24T11:46:57.987297Z","shell.execute_reply.started":"2024-10-24T11:46:57.961847Z","shell.execute_reply":"2024-10-24T11:46:57.986170Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F.groupby('series_description').count()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:57.989182Z","iopub.execute_input":"2024-10-24T11:46:57.989698Z","iopub.status.idle":"2024-10-24T11:46:58.008008Z","shell.execute_reply.started":"2024-10-24T11:46:57.989637Z","shell.execute_reply":"2024-10-24T11:46:58.007011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F = F.merge(train[['study_id','fold']],left_on='study_id',right_on='study_id')\nF.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:58.009652Z","iopub.execute_input":"2024-10-24T11:46:58.010278Z","iopub.status.idle":"2024-10-24T11:46:58.034108Z","shell.execute_reply.started":"2024-10-24T11:46:58.010231Z","shell.execute_reply":"2024-10-24T11:46:58.033180Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"F.groupby('fold').count()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:58.035403Z","iopub.execute_input":"2024-10-24T11:46:58.035846Z","iopub.status.idle":"2024-10-24T11:46:58.056963Z","shell.execute_reply.started":"2024-10-24T11:46:58.035797Z","shell.execute_reply":"2024-10-24T11:46:58.055828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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    # https://discuss.pytorch.org/t/rotation-matrix/128260\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":{"papermill":{"duration":0.029285,"end_time":"2024-07-21T15:01:31.222075","exception":false,"start_time":"2024-07-21T15:01:31.192790","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-10-24T11:46:58.058726Z","iopub.execute_input":"2024-10-24T11:46:58.059312Z","iopub.status.idle":"2024-10-24T11:46:58.072249Z","shell.execute_reply.started":"2024-10-24T11:46:58.059267Z","shell.execute_reply":"2024-10-24T11:46:58.071111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Sagittal_T1_sagittal_level_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 = 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#       By plane resizing I've been distorting the proportions\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":{"papermill":{"duration":0.034916,"end_time":"2024-07-21T15:01:31.320342","exception":false,"start_time":"2024-07-21T15:01:31.285426","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-10-24T11:46:58.074567Z","iopub.execute_input":"2024-10-24T11:46:58.075656Z","iopub.status.idle":"2024-10-24T11:46:58.089596Z","shell.execute_reply.started":"2024-10-24T11:46:58.075590Z","shell.execute_reply":"2024-10-24T11:46:58.088733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tds = Sagittal_T1_sagittal_level_Dataset(F)","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:58.091041Z","iopub.execute_input":"2024-10-24T11:46:58.092071Z","iopub.status.idle":"2024-10-24T11:46:58.101974Z","shell.execute_reply.started":"2024-10-24T11:46:58.092027Z","shell.execute_reply":"2024-10-24T11:46:58.100915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for k in range(5):\n    image,centers = tds.__getitem__(np.random.randint(len(tds)))\n    centers = centers[centers.isnan().sum(1) == 0]\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.sum(0)\n    plt.imshow(image.cpu()[0] + .5*(mask.cpu() > TH))\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:58.103298Z","iopub.execute_input":"2024-10-24T11:46:58.104132Z","iopub.status.idle":"2024-10-24T11:46:59.703840Z","shell.execute_reply.started":"2024-10-24T11:46:58.104081Z","shell.execute_reply":"2024-10-24T11:46:59.702807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vds = Sagittal_T1_sagittal_level_Dataset(F,VALID=True)","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:59.705052Z","iopub.execute_input":"2024-10-24T11:46:59.705363Z","iopub.status.idle":"2024-10-24T11:46:59.709830Z","shell.execute_reply.started":"2024-10-24T11:46:59.705329Z","shell.execute_reply":"2024-10-24T11:46:59.708802Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for k in range(5):\n    image,centers = vds.__getitem__(np.random.randint(len(vds)))\n    centers = centers[centers.isnan().sum(1) == 0]\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.sum(0)\n    plt.imshow(image.cpu()[0] + .5*(mask.cpu() > TH))\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:46:59.710727Z","iopub.execute_input":"2024-10-24T11:46:59.711037Z","iopub.status.idle":"2024-10-24T11:47:01.273911Z","shell.execute_reply.started":"2024-10-24T11:46:59.711005Z","shell.execute_reply":"2024-10-24T11:47:01.272978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del tds,vds\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:47:01.275182Z","iopub.execute_input":"2024-10-24T11:47:01.275518Z","iopub.status.idle":"2024-10-24T11:47:01.565012Z","shell.execute_reply.started":"2024-10-24T11:47:01.275484Z","shell.execute_reply":"2024-10-24T11:47:01.564088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"papermill":{"duration":0.02732,"end_time":"2024-07-21T15:01:31.366492","exception":false,"start_time":"2024-07-21T15:01:31.339172","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-10-24T11:47:01.566171Z","iopub.execute_input":"2024-10-24T11:47:01.566507Z","iopub.status.idle":"2024-10-24T11:47:01.573484Z","shell.execute_reply.started":"2024-10-24T11:47:01.566474Z","shell.execute_reply":"2024-10-24T11:47:01.572569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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#       MinMaxScaling along the class plane to generate a heatmap\n        min_values = x.min(-1)[0].view(-1,1)\n        max_values = x.max(-1)[0].view(-1,1)\n        d = (max_values - min_values)\n        d[d == 0] = 1\n        x = (x - min_values)/d\n        \n        return x.view(-1,self.classes,H,W)","metadata":{"papermill":{"duration":0.03081,"end_time":"2024-07-21T15:01:31.645783","exception":false,"start_time":"2024-07-21T15:01:31.614973","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-10-24T11:47:01.574907Z","iopub.execute_input":"2024-10-24T11:47:01.575305Z","iopub.status.idle":"2024-10-24T11:47:01.585014Z","shell.execute_reply.started":"2024-10-24T11:47:01.575258Z","shell.execute_reply":"2024-10-24T11:47:01.584130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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        \n        return D.mean()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:47:01.586168Z","iopub.execute_input":"2024-10-24T11:47:01.587044Z","iopub.status.idle":"2024-10-24T11:47:01.597575Z","shell.execute_reply.started":"2024-10-24T11:47:01.586999Z","shell.execute_reply":"2024-10-24T11:47:01.596738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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#plt.plot(nt(.25,1,np.arange(EPOCHS),EPOCHS))\n#plt.show()\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":{"papermill":{"duration":0.278967,"end_time":"2024-07-21T15:01:31.997510","exception":false,"start_time":"2024-07-21T15:01:31.718543","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-10-24T11:47:01.598852Z","iopub.execute_input":"2024-10-24T11:47:01.599202Z","iopub.status.idle":"2024-10-24T11:47:01.610042Z","shell.execute_reply.started":"2024-10-24T11:47:01.599160Z","shell.execute_reply":"2024-10-24T11:47:01.609148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for f in FOLDS:\n    seed_everything(SEED)\n    model = myUNet(5)\n    \n    tdf = F[F['fold'] != f]\n    vdf = F[F['fold'] == f]\n\n    tds = Sagittal_T1_sagittal_level_Dataset(tdf)\n    vds = Sagittal_T1_sagittal_level_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(EPOCHS)\n    torch.save(model,'Sagittal_T1_sagittal_level_segmentation_'+str(f))\n    del tdl,vdl,dls,model,learn\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T11:47:01.613727Z","iopub.execute_input":"2024-10-24T11:47:01.614183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f = 1\nmodel = torch.load('Sagittal_T1_sagittal_level_segmentation_'+str(f))\nvdf = F[F['fold'] == f]\nvds = Sagittal_T1_sagittal_level_Dataset(vdf,VALID=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for k in range(5):\n    i = np.random.randint(len(vds))\n    print(i)\n    image,centers = vds.__getitem__(np.random.randint(len(vds)))\n    centers = centers[centers.isnan().sum(1) == 0]\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.sum(0)\n    fig, axes = plt.subplots(1, 2, figsize=(10,10))\n    axes[0].imshow(image.cpu()[0] + .5*(model(image.unsqueeze(0))[0].detach().cpu() > TH).sum(0))\n    axes[1].imshow(image.cpu()[0] + .5*(mask.cpu() > TH))\n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del model,vdf,vds\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}