{"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 cv2\nimport pydicom\nimport numpy as np\nimport os\nimport glob\nfrom tqdm import tqdm\nimport gc\nimport pickle\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 = 777\nFOLDS = [1,2,3,4,5]\nPATCH_SIZE = 512\nBS = 16\nLR = 5e-4\nEPOCHS = 10\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":"train = pd.read_csv('train_split.csv')\ntrain.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_coor = pd.read_csv('C:/Users/Angel/Kaggle/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':[]}\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":"coordinates = np.zeros((len(S),4))\ncoordinates[:] = np.nan\nfor i in range(len(S)):\n    row = S.iloc[i]\n    for side in centers[row['study_id']][row['series_id']][row['instance_number']]:\n            if len(centers[row['study_id']][row['series_id']][row['instance_number']][side]) > 0:\n                center = 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                ] = center","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":"S = S[[\n    'study_id',\n    'series_id',\n    'instance_number'   \n]]\nS[[\n    'x_L',\n    'y_L',\n    'x_R',\n    'y_R' \n]] = coordinates\nS = S.drop_duplicates().reset_index(drop=True)\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[['x_L','y_L','x_R','y_R']].isna().sum(1) > 0].reset_index(drop=True).tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NBH = 2\nS_values = S.values\ntotal = 0\ntotal_imputed = 0\nfor i in range(len(S_values)):\n    mask = np.isnan(S_values[i,-4:])\n    v = [S_values[i,-4:]]\n    if mask.sum() > 0 :\n        total += 1\n        try:\n            if (S_values[i-1,:3] - S_values[i,:3]).sum()**2 < NBH: v.append(S_values[i-1,-4:])\n        except:\n            None\n        try:\n            if (S_values[i+1,:3] - S_values[i,:3]).sum()**2 < NBH: v.append(S_values[i-1,-4:])\n        except:\n            None\n        v = np.nanmean(np.stack(v),0)\n        S_values[i,-4:][mask] = v[mask]\n        if np.isnan(S_values[i,-4:]).sum() == 0: total_imputed += 1\nprint(total_imputed*100/total)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S_values[:,-4:][np.isnan(S_values[:,-4:])] = 0","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S[['x_L','y_L','x_R','y_R']] = S_values[:,-4:]","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('C:/Users/Angel/Kaggle/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":"coor = [\n    'x_L',\n    'y_L',\n    'x_R',\n    'y_R'    \n]","metadata":{"execution":{"iopub.execute_input":"2024-07-21T15:01:31.168524Z","iopub.status.busy":"2024-07-21T15:01:31.167713Z","iopub.status.idle":"2024-07-21T15:01:31.172286Z","shell.execute_reply":"2024-07-21T15:01:31.171418Z"},"papermill":{"duration":0.026097,"end_time":"2024-07-21T15:01:31.174262","exception":false,"start_time":"2024-07-21T15:01:31.148165","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def augment_image_and_centers(image,centers,alpha):\n    # Randomly rotate the image.\n    angle = torch.as_tensor(random.uniform(-180, 180)*alpha)\n    image = torchvision.transforms.functional.rotate(image,angle.item(),interpolation=torchvision.transforms.InterpolationMode.BILINEAR)\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    centers = ((centers.cpu() - PATCH_SIZE//2) @ rot) + PATCH_SIZE//2\n\n    return image,centers","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":"torch_resize = torchvision.transforms.Resize((PATCH_SIZE,PATCH_SIZE),antialias=True)","metadata":{"execution":{"iopub.execute_input":"2024-07-21T15:01:31.260490Z","iopub.status.busy":"2024-07-21T15:01:31.260154Z","iopub.status.idle":"2024-07-21T15:01:31.264806Z","shell.execute_reply":"2024-07-21T15:01:31.263874Z"},"papermill":{"duration":0.026156,"end_time":"2024-07-21T15:01:31.266747","exception":false,"start_time":"2024-07-21T15:01:31.240591","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Axial_T1_UNet_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()\n        \n        sample = 'C:/Users/Angel/kaggle/train/'\n        sample = sample+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_SIZE/W\n        centers[:,1] = centers[:,1]*PATCH_SIZE/H\n\n        if not self.VALID: image,centers = augment_image_and_centers(image,centers,self.alpha)\n\n        return image.to(device),centers.to(device)","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":"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":"idx_map = torch.stack([torch.arange(PATCH_SIZE)]*PATCH_SIZE).float().to(device)\nidx_map = torch.stack([idx_map,idx_map.T]).view(1,1,2,PATCH_SIZE,PATCH_SIZE)","metadata":{"execution":{"iopub.execute_input":"2024-07-21T15:01:31.406014Z","iopub.status.busy":"2024-07-21T15:01:31.405259Z","iopub.status.idle":"2024-07-21T15:01:31.593921Z","shell.execute_reply":"2024-07-21T15:01:31.592947Z"},"papermill":{"duration":0.211011,"end_time":"2024-07-21T15:01:31.596366","exception":false,"start_time":"2024-07-21T15:01:31.385355","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class myUNet(nn.Module):\n    def __init__(self):\n        super(myUNet, self).__init__()\n\n        self.UNet = smp.Unet(\n            encoder_name=\"resnet18\",\n            classes=2,\n            in_channels=1\n        ).to(device)\n\n    def forward(self,X):\n        x = self.UNet(X)\n#       MinMaxScaling along the class plane to generate a heatmap\n        min_values = x.view(-1,2,PATCH_SIZE*PATCH_SIZE).min(-1)[0].view(-1,2,1,1)\n        max_values = x.view(-1,2,PATCH_SIZE*PATCH_SIZE).max(-1)[0].view(-1,2,1,1)\n        x = (x - min_values)/(max_values - min_values)\n        \n        return x","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":"# Let's start by segmentation only\nclass myLoss(nn.Module):\n    def __init__(\n            self,\n            alpha=.5\n        ):\n        super().__init__()\n        self.alpha = alpha\n\n    def clone(self):\n        return myLoss(self.alpha)\n\n    def forward(\n            self,\n            y,# Predictions\n            t # Targets\n        ):\n        confident = ((t  < PATCH_SIZE - 64)*(t > 64)).min(-1)[0].bool()\n        mask_true = t[confident]\n        mask_pred = y[confident]\n#       The heatmap Loss as the distance between the predicted Normal and the ideal one\n#       Let's define the ideal heatmaps as the Normal distributions\n#       centered on the diagnostic centers with s2 = PATCH_SIZE/8\n        s2 = s2 = torch.as_tensor([PATCH_SIZE/8])\n#       Then the corresponding alphas and normalization constants would be\n        A = -1/(2*s2).to(device)\n        K = 1/torch.sqrt(2*math.pi*s2).to(device)\n#       Predicted heatmaps rescaling\n        mask_pred = mask_pred*K\n#       Ideal heatmaps\n        mask = mask_true.view(-1,2,1,1)\n        mask = (idx_map[0] - mask)\n        mask = (mask*mask).sum(1)\n        mask = torch.exp(A*mask)*K\n\n        mask = mask.view(-1,PATCH_SIZE*PATCH_SIZE)\n        mask_pred = mask_pred.view(-1,PATCH_SIZE*PATCH_SIZE)\n#       Distance\n        D = 1 -((mask*mask_pred).sum(-1))**2/((mask*mask).sum(-1)*(mask_pred*mask_pred).sum(-1))\n        \n        return D.nanmean()","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\nplt.plot(nt(.25,1,np.arange(EPOCHS),EPOCHS))\nplt.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    \n    seed_everything(SEED)\n    model = myUNet()\n    \n    tdf = S[S['fold'] != f]\n    vdf = S[S['fold'] == f]\n\n    tds = Axial_T1_UNet_Dataset(tdf)\n    vds = Axial_T1_UNet_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_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_segmentation_'+str(f))\nvdf = S[S['fold'] == f]\nvds = Axial_T1_UNet_Dataset(vdf,VALID=True)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://stackoverflow.com/questions/22777049/how-can-i-draw-a-circle-in-a-data-array-map-in-python\nwidth, height = 11, 11\nA, B = 5, 5\nr = 5\nEPSILON = 2.2","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"i = np.random.randint(len(vds))\nimg,centers = vds.__getitem__(i)\nOUT = model(img.unsqueeze(0)).cpu().detach()\ncenters = centers.cpu().long()\nprint(i)\nimg = img[0].cpu()\nfor k in range(2):\n    img += OUT[0,k].cpu()\n    Y,X = centers.cpu().long()[k]\n    for y in range(height):\n        for x in range(width):\n            # see if we're close to (x-a)**2 + (y-b)**2 == r**2\n            if abs((x-A)**2 + (y-B)**2 - r**2) < EPSILON**2:\n                img[x+X-5,y+Y-5] += 1\nplt.imshow(img)\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"i = np.random.randint(len(vds))\nimg,centers = vds.__getitem__(i)\nOUT = model(img.unsqueeze(0)).cpu().detach()\ncenters = centers.cpu().long()\nprint(i)\nfor k in range(2):\n    image = img[0].cpu() + OUT[0,k].cpu()\n    Y,X = centers.cpu().long()[k]\n    for y in range(height):\n        for x in range(width):\n            # see if we're close to (x-a)**2 + (y-b)**2 == r**2\n            if abs((x-A)**2 + (y-B)**2 - r**2) < EPSILON**2:\n                image[x+X-5,y+Y-5] += 1\n    plt.imshow(image)\n    plt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"i = np.random.randint(len(vds))\nimg,centers = vds.__getitem__(i)\nOUT = model(img.unsqueeze(0)).cpu().detach()\ncenters = centers.cpu().long()\nprint(i)\nfig, axes1 = plt.subplots(1, 2, figsize=(10,10))\nfig, axes2 = plt.subplots(1, 2, figsize=(10,10))\nfor k in range(2):\n    image = img[0].cpu() + OUT[0,k].cpu()\n    c = (OUT[0,k].unsqueeze(0)*idx_map[0,0].cpu()).sum(-1).sum(-1)\n    d = OUT[0,k].sum()\n    c = c/d\n    Y,X = centers.cpu().long()[k]\n    YY,XX = c.long()\n    for y in range(height):\n        for x in range(width):\n            # see if we're close to (x-a)**2 + (y-b)**2 == r**2\n            if abs((x-A)**2 + (y-B)**2 - r**2) < EPSILON**2:\n                image[x+X-5,y+Y-5] = 0\n    axes1[k].imshow(image)\n    axes2[k].imshow(image[XX-64:XX+64,YY-64:YY+64])\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"P = 32\nTH = .5\ni = np.random.randint(len(vds))\nimg,centers = vds.__getitem__(i)\nE = model(img.unsqueeze(0))\nwith torch.no_grad():\n    for rot in [1,2,3]:\n        E += torch.rot90(model(torch.rot90(img.unsqueeze(0),rot,dims=[-2, -1])),-rot,dims=[-2, -1])\nOUT = E.cpu().detach()/4 > TH\ncenters = centers.cpu().long()\nprint(i)\n# fig, axes1 = plt.subplots(1, 2, figsize=(10,10))\nfig, axes2 = plt.subplots(1, 2, figsize=(10,10))\nfor k in range(2):\n    image = img[0].cpu() + OUT[0,k].cpu()\n    c = (OUT[0,k].unsqueeze(0)*idx_map[0,0].cpu()).sum(-1).sum(-1)\n    d = OUT[0,k].sum()\n    c = c/d\n    Y,X = centers.cpu().long()[k]\n    YY,XX = c.long()\n    for y in range(height):\n        for x in range(width):\n            # see if we're close to (x-a)**2 + (y-b)**2 == r**2\n            if abs((x-A)**2 + (y-B)**2 - r**2) < EPSILON**2:\n                image[x+X-5,y+Y-5] = 0\n#   axes1[k].imshow(image)\n    axes2[k].imshow(image[XX-P:XX+P,YY-P:YY+P])\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del model,vdf,vds\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FOLDS = [1,2,3,4,5]\nBS = 64\nOUT = torch.zeros((BS,2,PATCH_SIZE,PATCH_SIZE)).to(device)\ndf = train[['study_id','fold']].merge(df_meta_f[df_meta_f.series_description == 'Axial T2'][['study_id','series_id']],left_on='study_id',right_on='study_id')\ncoord = {}\nfor f in FOLDS:\n    coord[f] = {}\nfor f in FOLDS:\n    for i in range(len(df)):\n        row = df.iloc[i]\n        coord[f][row['study_id']] = {}\nfor f in FOLDS:\n    for i in range(len(df)):\n        row = df.iloc[i]\n        coord[f][row['study_id']][row['series_id']] = {}\nmodels = {}\nfor f in FOLDS:\n    model = torch.load('axial_T2_segmentation_'+str(f))\n    model.eval()\n    models[f] = model\nfor f in FOLDS:\n    for i in tqdm(range(len(df))):\n        row = df.iloc[i]\n        \n        sample = 'C:/Users/Angel/kaggle/train/'\n        sample = sample+str(int(row['study_id']))+'/'+str(int(row['series_id']))\n\n        images = [x.replace('\\\\','/') for x in glob.glob(sample+'/*.dcm')]\n        images.sort(key=lambda x:int(x.split('/')[-1].replace('.dcm','')))\n        images = [torch.as_tensor(pydicom.dcmread(img).pixel_array.astype('float32')) for img in images]\n\n        shapes = [img.shape for img in images]\n        H,W = np.array(shapes).max(0)\n\n        image = torch.concat([torch.nn.functional.pad(\n            images[k].unsqueeze(0),(\n                (W - shapes[k][-1])//2,\n                (W - shapes[k][-1]) - (W - shapes[k][-1])//2,\n                (H - shapes[k][-2])//2,\n                (H - shapes[k][-2]) - (H - shapes[k][-2])//2\n            ),\n        mode='reflect') for k in range(len(images))]).float()\n\n        if H > W:\n            d = W\n            h = (H - d)//2\n            image = image[:,h:h+d]\n            H = W\n        elif H < W:\n            d = H\n            w = (W - d)//2\n            image = image[:,:,w:w+d]\n            W = H\n\n        image = torch_resize(image/image.max()).float().unsqueeze(1).to(device)\n\n        D = len(image)\n        c = torch.zeros((D,2,2)).to(device)\n        with torch.no_grad():\n            for k in range(D//BS+1):\n                img = image[k*BS:k*BS+BS]\n                N = len(img)\n                OUT[:N] = 0 \n                for rot in [0,1,2,3]:\n                    OUT[:N] += torch.rot90(models[f](torch.rot90(img,rot,dims=[-2, -1])),-rot,dims=[-2, -1])\n\n                OUT[:N] = OUT[:N]/4 > TH\n                c[k*BS:k*BS+N] = (OUT[:N].unsqueeze(2)*idx_map[0]).view(N,2,2,PATCH_SIZE*PATCH_SIZE).sum(-1).float()\n                d = OUT[:N].view(N,2,PATCH_SIZE*PATCH_SIZE).sum(-1)\n                m = d > 0\n                c[k*BS:k*BS+N][m] = (c[k*BS:k*BS+N][m]/d[m].unsqueeze(-1))\n                c[k*BS:k*BS+N][~m] = torch.nan\n        coord[f][row['study_id']][row['series_id']]=c\n        del c\n        gc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open('axial_centers.pkl', 'wb') as f:\n    pickle.dump(coord, f)","metadata":{},"execution_count":null,"outputs":[]}]}