{"metadata":{"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.7.12"},"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\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":{"execution":{"iopub.execute_input":"2024-07-21T15:00:53.042579Z","iopub.status.busy":"2024-07-21T15:00:53.042231Z","iopub.status.idle":"2024-07-21T15:01:24.461173Z","shell.execute_reply":"2024-07-21T15:01:24.460287Z"},"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_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 1337\nFOLDS = [1,2,3,4,5]\nPATH = 'C:/Users/Angel/kaggle/'# Main path\nTRAIN_PATH = 'C:/Users/Angel/kaggle/train/'# Training images folder\nENCODER_NAME = \"resnet18\"\nANGLE = 180#30\nS2 = 64\nPATCH_H = 512\nPATCH_W = 512\nBS = 16\nLR = 2.5e-4\nEPOCHS = 1\nTH = .5","metadata":{"execution":{"iopub.execute_input":"2024-07-21T15:01:24.492075Z","iopub.status.busy":"2024-07-21T15:01:24.491716Z","iopub.status.idle":"2024-07-21T15:01:24.496542Z","shell.execute_reply":"2024-07-21T15:01:24.495663Z"},"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_count":null,"outputs":[]},{"cell_type":"code","source":"S2 = torch.as_tensor(S2)\nA = -1/(2*S2).to(device)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv(PATH + 'train_split.csv')\ntrain.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_coor = pd.read_csv(PATH + 'train_label_coordinates.csv')\ndf_coor.tail()","metadata":{},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.execute_input":"2024-07-21T15:01:24.882922Z","iopub.status.busy":"2024-07-21T15:01:24.882156Z","iopub.status.idle":"2024-07-21T15:01:24.915648Z","shell.execute_reply":"2024-07-21T15:01:24.914644Z"},"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_count":null,"outputs":[]},{"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":{"execution":{"iopub.execute_input":"2024-07-21T15:01:25.319656Z","iopub.status.busy":"2024-07-21T15:01:25.319335Z","iopub.status.idle":"2024-07-21T15:01:27.438123Z","shell.execute_reply":"2024-07-21T15:01:27.437297Z"},"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_count":null,"outputs":[]},{"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":{},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.execute_input":"2024-07-21T15:01:27.473771Z","iopub.status.busy":"2024-07-21T15:01:27.473424Z","iopub.status.idle":"2024-07-21T15:01:30.212739Z","shell.execute_reply":"2024-07-21T15:01:30.211950Z"},"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_count":null,"outputs":[]},{"cell_type":"code","source":"coor = [\n    'x_L',\n    'y_L',\n    'x_R',\n    'y_R'    \n]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S[coor] = coordinates\nS.tail()","metadata":{"execution":{"iopub.execute_input":"2024-07-21T15:01:30.248435Z","iopub.status.busy":"2024-07-21T15:01:30.248079Z","iopub.status.idle":"2024-07-21T15:01:30.279735Z","shell.execute_reply":"2024-07-21T15:01:30.278823Z"},"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_count":null,"outputs":[]},{"cell_type":"code","source":"S[S[coor].isna().sum(1) > 0].reset_index(drop=True).tail()","metadata":{},"execution_count":null,"outputs":[]},{"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":{},"execution_count":null,"outputs":[]},{"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":{},"execution_count":null,"outputs":[]},{"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":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S[S[coor].isna().sum(1) > 0].reset_index(drop=True).tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S['flip'] = False\nfS = S.copy()\nfS['flip'] = True\nS = pd.concat([S,fS]).reset_index(drop=True)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S = S.merge(train[['study_id','fold']], left_on='study_id', right_on='study_id')\nS.tail()","metadata":{},"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_count":null,"outputs":[]},{"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":{"execution":{"iopub.execute_input":"2024-07-21T15:01:30.640373Z","iopub.status.busy":"2024-07-21T15:01:30.639426Z","iopub.status.idle":"2024-07-21T15:01:30.667189Z","shell.execute_reply":"2024-07-21T15:01:30.666215Z"},"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_count":null,"outputs":[]},{"cell_type":"code","source":"S.groupby('series_description').count()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S.groupby('fold').count()","metadata":{},"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":{"execution":{"iopub.execute_input":"2024-07-21T15:01:31.213308Z","iopub.status.busy":"2024-07-21T15:01:31.212554Z","iopub.status.idle":"2024-07-21T15:01:31.220122Z","shell.execute_reply":"2024-07-21T15:01:31.219249Z"},"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_count":null,"outputs":[]},{"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#       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,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":{"execution":{"iopub.execute_input":"2024-07-21T15:01:31.305523Z","iopub.status.busy":"2024-07-21T15:01:31.304849Z","iopub.status.idle":"2024-07-21T15:01:31.318422Z","shell.execute_reply":"2024-07-21T15:01:31.317488Z"},"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_count":null,"outputs":[]},{"cell_type":"code","source":"tds = Axial_T2_axial_side_Dataset(S)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for k in range(10):\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_count":null,"outputs":[]},{"cell_type":"code","source":"vds = Axial_T2_axial_side_Dataset(S,VALID=True)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for k in range(10):\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_count":null,"outputs":[]},{"cell_type":"code","source":"del tds,vds\ngc.collect()","metadata":{},"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":{"execution":{"iopub.execute_input":"2024-07-21T15:01:31.359638Z","iopub.status.busy":"2024-07-21T15:01:31.359020Z","iopub.status.idle":"2024-07-21T15:01:31.364428Z","shell.execute_reply":"2024-07-21T15:01:31.363571Z"},"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_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":{"execution":{"iopub.execute_input":"2024-07-21T15:01:31.635757Z","iopub.status.busy":"2024-07-21T15:01:31.635401Z","iopub.status.idle":"2024-07-21T15:01:31.643447Z","shell.execute_reply":"2024-07-21T15:01:31.642550Z"},"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_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.execute_input":"2024-07-21T15:01:31.687492Z","iopub.status.busy":"2024-07-21T15:01:31.686705Z","iopub.status.idle":"2024-07-21T15:01:31.695723Z","shell.execute_reply":"2024-07-21T15:01:31.694854Z"},"papermill":{"duration":0.031904,"end_time":"2024-07-21T15:01:31.697861","exception":false,"start_time":"2024-07-21T15:01:31.665957","status":"completed"},"tags":[]},"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":{"execution":{"iopub.execute_input":"2024-07-21T15:01:31.741286Z","iopub.status.busy":"2024-07-21T15:01:31.740964Z","iopub.status.idle":"2024-07-21T15:01:31.995395Z","shell.execute_reply":"2024-07-21T15:01:31.994477Z"},"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_count":null,"outputs":[]},{"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(EPOCHS)\n    torch.save(model,'Axial_T2_axial_side_segmentation_'+str(f))\n    del tdl,vdl,dls,model,learn\n    gc.collect()","metadata":{"execution":{"iopub.execute_input":"2024-07-21T15:01:32.037325Z","iopub.status.busy":"2024-07-21T15:01:32.036519Z","iopub.status.idle":"2024-07-21T19:59:55.569497Z","shell.execute_reply":"2024-07-21T19:59:55.568640Z"},"papermill":{"duration":17903.555334,"end_time":"2024-07-21T19:59:55.571935","exception":false,"start_time":"2024-07-21T15:01:32.016601","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f = 1\nmodel = torch.load('Axial_T2_axial_side_segmentation_'+str(f))\nvdf = S[S['fold'] == f]\nvds = Axial_T2_axial_side_Dataset(vdf,VALID=True)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for k in range(20):\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":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del model,vdf,vds\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]}]}