{"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\"\nPATCH_H = 512\nPATCH_W = 512\nANGLE = 30\nS2 = 64\nBS = 16\nLR = 1e-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'] == 'Spinal Canal Stenosis'\n].sort_values([\n    'study_id',\n    'series_id',\n    'level'\n]).reset_index(drop=True)\nS.tail()","metadata":{},"execution_count":null,"outputs":[]},{"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":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.boxplot(S['x_mean_fraction'])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S[S['x_mean_fraction'] < .8]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S = S[S['x_mean_fraction'] > .8]","metadata":{},"execution_count":null,"outputs":[]},{"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":{"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()\nS.tail()","metadata":{},"execution_count":null,"outputs":[]},{"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":{"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_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_count":null,"outputs":[]},{"cell_type":"code","source":"S[coor] = v\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":"for (study_id,series_id),df in tqdm(S.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    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":{},"execution_count":null,"outputs":[]},{"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":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S_mean = S.groupby(['study_id','series_id']).mean()\nS_mean.tail()","metadata":{},"execution_count":null,"outputs":[]},{"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":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S_values = S[coor].values\nS_values.shape","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask = S[coor].isna()\nmask.shape","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S_values[mask] = S_mean[mask]\nS[coor] = S_values\nS.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S[S.isna().sum(1) > 0].reset_index(drop=True).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 = S[S.series_description == 'Sagittal T2/STIR'].reset_index(drop=True)\nS.groupby('series_description').count()","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":"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 Sagittal_T2_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":{"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 = Sagittal_T2_sagittal_level_Dataset(S)","metadata":{},"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_count":null,"outputs":[]},{"cell_type":"code","source":"vds = Sagittal_T2_sagittal_level_Dataset(S,VALID=True)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for k in range(5):\n    image,centers = vds.__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":"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_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(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_sagittal_level_Dataset(tdf)\n    vds = Sagittal_T2_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_T2_sagittal_level_segmentation_'+str(f))\n    del tdl,vdl,dls,model,learn\n    gc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f = 1\nmodel = torch.load('Sagittal_T2_sagittal_level_segmentation_'+str(f))\nvdf = S[S['fold'] == f]\nvds = Sagittal_T2_sagittal_level_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":[]}]}