{"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\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]\nENCODER_NAME=\"resnet18\"\nPATCH_H = 128\nPATCH_W = 256\nDRIFT = 64\nPATCH_SIZE = 512\nANGLE = 30\nS2 = 16\nBS = 16\nLR = 5e-4\nEPOCHS = 2\nTH = .5","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S2 = torch.as_tensor(64)\nA = -1/(2*S2).to(device)\n#K = 1/torch.sqrt(2*math.pi*S2).to(device)","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":{"iopub.execute_input":"2024-07-21T15:01:24.675045Z","iopub.status.busy":"2024-07-21T15:01:24.674395Z","iopub.status.idle":"2024-07-21T15:01:24.692944Z","shell.execute_reply":"2024-07-21T15:01:24.692058Z"},"papermill":{"duration":0.035373,"end_time":"2024-07-21T15:01:24.695008","exception":false,"start_time":"2024-07-21T15:01:24.659635","status":"completed"},"tags":[]},"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":{"iopub.execute_input":"2024-07-21T15:01:24.724580Z","iopub.status.busy":"2024-07-21T15:01:24.724234Z","iopub.status.idle":"2024-07-21T15:01:24.850435Z","shell.execute_reply":"2024-07-21T15:01:24.849512Z"},"papermill":{"duration":0.143427,"end_time":"2024-07-21T15:01:24.852417","exception":false,"start_time":"2024-07-21T15:01:24.708990","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"left_df_coor = df_coor[df_coor.condition == 'Left Neural Foraminal Narrowing'][['series_id','level','instance_number','x','y']].groupby(['series_id','level']).mean().sort_values(['series_id','level'])\nleft_df_coor = left_df_coor.rename(columns={'instance_number':'left_instance_number','x':'left_x','y':'left_y'})\nleft_df_coor.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"right_df_coor = df_coor[df_coor.condition == 'Right Neural Foraminal Narrowing'][['series_id','level','instance_number','x','y']].groupby(['series_id','level']).mean().sort_values(['series_id','level'])\nright_df_coor = right_df_coor.rename(columns={'instance_number':'right_instance_number','x':'right_x','y':'right_y'})\nright_df_coor.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"flipped_foraminal = left_df_coor.merge(right_df_coor,left_on=left_df_coor.index.to_numpy(),right_on=right_df_coor.index.to_numpy())\nflipped_foraminal.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"flipped_foraminal['flipped'] = flipped_foraminal.left_instance_number > flipped_foraminal.right_instance_number\nflipped_foraminal.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"flipped_foraminal['level'] = flipped_foraminal['key_0'].apply(lambda v: v[1])\nflipped_foraminal['series_id'] = flipped_foraminal['key_0'].apply(lambda v: v[0])\nflipped_foraminal = flipped_foraminal.drop(columns=['key_0'])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"flipped_foraminal.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"instance_numbers = {}\nfor i in tqdm(range(len(flipped_foraminal))):\n    row = flipped_foraminal.iloc[i]\n    instance_numbers[row['series_id']] = {\n        'L':{\n            'L1/L2':torch.nan,\n            'L2/L3':torch.nan,\n            'L3/L4':torch.nan,\n            'L4/L5':torch.nan,\n            'L5/S1':torch.nan\n        },\n        'R':{\n            'L1/L2':torch.nan,\n            'L2/L3':torch.nan,\n            'L3/L4':torch.nan,\n            'L4/L5':torch.nan,\n            'L5/S1':torch.nan\n        },\n        'L_x':{\n            'L1/L2':torch.nan,\n            'L2/L3':torch.nan,\n            'L3/L4':torch.nan,\n            'L4/L5':torch.nan,\n            'L5/S1':torch.nan\n        },\n        'R_x':{\n            'L1/L2':torch.nan,\n            'L2/L3':torch.nan,\n            'L3/L4':torch.nan,\n            'L4/L5':torch.nan,\n            'L5/S1':torch.nan\n        },\n        'L_y':{\n            'L1/L2':torch.nan,\n            'L2/L3':torch.nan,\n            'L3/L4':torch.nan,\n            'L4/L5':torch.nan,\n            'L5/S1':torch.nan\n        },\n        'R_y':{\n            'L1/L2':torch.nan,\n            'L2/L3':torch.nan,\n            'L3/L4':torch.nan,\n            'L4/L5':torch.nan,\n            'L5/S1':torch.nan\n        }\n    }\n\nfor i in tqdm(range(len(flipped_foraminal))):\n    row = flipped_foraminal.iloc[i]\n    instance_numbers[row['series_id']]['L'][row['level']] = row['left_instance_number']\n    instance_numbers[row['series_id']]['R'][row['level']] = row['right_instance_number']\n    instance_numbers[row['series_id']]['L_x'][row['level']] = row['left_x']\n    instance_numbers[row['series_id']]['L_y'][row['level']] = row['left_y']\n    instance_numbers[row['series_id']]['R_x'][row['level']] = row['right_x']\n    instance_numbers[row['series_id']]['R_y'][row['level']] = row['right_y']","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"flipped_foraminal_counts = flipped_foraminal[['series_id','flipped']].groupby('series_id').count().reset_index()\nprint(flipped_foraminal_counts.tail())\nflipped_foraminal = flipped_foraminal[['series_id','flipped']].groupby('series_id').sum().reset_index()\nprint(flipped_foraminal.tail())","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"flipped_foraminal['flipped'] = flipped_foraminal['flipped']/flipped_foraminal_counts['flipped']\nflipped_foraminal.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"flipped_foraminal.groupby('flipped').count()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"flipped_foraminal = flipped_foraminal[flipped_foraminal['flipped'].isin([0.,1.])].reset_index(drop=True)\nflipped_foraminal['flipped'] = flipped_foraminal['flipped'] > 0\nflipped_foraminal.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"v = torch.zeros((len(flipped_foraminal),30))\nfor i in tqdm(range(len(flipped_foraminal))):\n    row = flipped_foraminal.iloc[i]\n    v[i,:5] = torch.as_tensor(list(instance_numbers[row['series_id']]['L'].values())).float()\n    v[i,5:10] = torch.as_tensor(list(instance_numbers[row['series_id']]['R'].values())).float()\n    v[i,10:15] = torch.as_tensor(list(instance_numbers[row['series_id']]['L_x'].values())).float()\n    v[i,15:20] = torch.as_tensor(list(instance_numbers[row['series_id']]['L_y'].values())).float()\n    v[i,20:25] = torch.as_tensor(list(instance_numbers[row['series_id']]['R_x'].values())).float()\n    v[i,25:] = torch.as_tensor(list(instance_numbers[row['series_id']]['R_y'].values())).float()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"flipped_foraminal[[\n    'L_L1L2','L_L2L3','L_L3L4','L_L4L5','L_L5S1',\n    'R_L1L2','R_L2L3','R_L3L4','R_L4L5','R_L5S1',\n    'L_x_L1L2','L_x_L2L3','L_x_L3L4','L_x_L4L5','L_x_L5S1',\n    'L_y_L1L2','L_y_L2L3','L_y_L3L4','L_y_L4L5','L_y_L5S1',\n    'R_x_L1L2','R_x_L2L3','R_x_L3L4','R_x_L4L5','R_x_L5S1',\n    'R_y_L1L2','R_y_L2L3','R_y_L3L4','R_y_L4L5','R_y_L5S1'\n]] = v\nflipped_foraminal.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('C:/Users/Angel/Kaggle/train_split.csv')\ntrain = df_meta_f.merge(\n    flipped_foraminal,\n    left_on='series_id',\n    right_on='series_id'\n).merge(\n    train[['study_id','fold']],\n    left_on='study_id',\n    right_on='study_id'\n)\ntrain['flip'] = False\nftrain = train.copy()\nftrain[['flipped','flip']] = ~ftrain[['flipped','flip']]\ntrain = pd.concat([train,ftrain]).reset_index(drop=True)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.groupby('fold').count()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.groupby('series_description').count()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train[train.isna().sum(1) > 0]","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_SIZE),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_count":null,"outputs":[]},{"cell_type":"code","source":"class Sagittal_T1_axial_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        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        D = len(images)\n        \n        if row.flip:\n            images.sort(reverse=True, key=lambda x: int(x.split('/')[-1].replace('.dcm', '')))\n            L = D - row[[\n                'L_L1L2',\n                'L_L2L3',\n                'L_L3L4',\n                'L_L4L5',\n                'L_L5S1',   \n            ]].values.astype(np.float32)\n            R = D - row[[\n                'R_L1L2',\n                'R_L2L3',\n                'R_L3L4',\n                'R_L4L5',\n                'R_L5S1',   \n            ]].values.astype(np.float32)\n            R_X = row[[\n                'L_x_L1L2',\n                'L_x_L2L3',\n                'L_x_L3L4',\n                'L_x_L4L5',\n                'L_x_L5S1',   \n            ]].values.astype(np.float32)\n            L_X = row[[\n                'R_x_L1L2',\n                'R_x_L2L3',\n                'R_x_L3L4',\n                'R_x_L4L5',\n                'R_x_L5S1',   \n            ]].values.astype(np.float32)\n        else:\n            images.sort(reverse=False, key=lambda x: int(x.split('/')[-1].replace('.dcm', '')))\n            L = row[[\n                'L_L1L2',\n                'L_L2L3',\n                'L_L3L4',\n                'L_L4L5',\n                'L_L5S1',   \n            ]].values.astype(np.float32) - 1\n            R = row[[\n                'R_L1L2',\n                'R_L2L3',\n                'R_L3L4',\n                'R_L4L5',\n                'R_L5S1',   \n            ]].values.astype(np.float32) - 1\n            L_X = row[[\n                'L_x_L1L2',\n                'L_x_L2L3',\n                'L_x_L3L4',\n                'L_x_L4L5',\n                'L_x_L5S1',   \n            ]].values.astype(np.float32)\n            R_X = row[[\n                'R_x_L1L2',\n                'R_x_L2L3',\n                'R_x_L3L4',\n                'R_x_L4L5',\n                'R_x_L5S1',   \n            ]].values.astype(np.float32)\n\n        Z = ((row[[\n            'L_y_L1L2',\n            'L_y_L2L3',\n            'L_y_L3L4',\n            'L_y_L4L5',\n            'L_y_L5S1',   \n        ]].values + row[[\n            'R_y_L1L2',\n            'R_y_L2L3',\n            'R_y_L3L4',\n            'R_y_L4L5',\n            'R_y_L5S1',   \n        ]].values)/2).astype(np.float32)\n        \n        UNK = (np.isnan(L) + np.isnan(R)+ np.isnan(L_X) + np.isnan(R_X) + np.isnan(Z)) > 0\n\n        image = torch.as_tensor(np.stack([pydicom.dcmread(x).pixel_array for x in images]).astype(float))\n        image = (image/image.max()).to(device)\n        img = torch.zeros(5,PATCH_H,PATCH_W).to(device)\n        mask = torch.zeros(5,PATCH_H,PATCH_W).to(device)\n        label = torch.zeros(5,2,2)\n        for i in range(5):\n            if UNK[i]:\n                k = np.arange(5)[~UNK][np.random.randint((~UNK).sum())]\n            else:\n                k = i\n            z = Z[k].astype(int) + 1 - np.random.randint(3)\n            L_x = (L_X[k]*(PATCH_SIZE)/(len(image[0,0])))\n            R_x = (R_X[k]*(PATCH_SIZE)/(len(image[0,0])))\n            x = (L_x + R_x)/2\n            l = (L[k]*(PATCH_H)/(len(image)))\n            r = (R[k]*(PATCH_H)/(len(image)))\n\n            if not self.VALID:\n                c = int((self.alpha*(DRIFT - np.random.randint(2*DRIFT)) + x))\n            else:\n                c = int(x)\n            s = c - PATCH_W//2\n            if s < 0: s = 0\n            if s > PATCH_SIZE - PATCH_W: s = PATCH_SIZE - PATCH_W\n\n            img[i] = torch_resize(image[:,z].unsqueeze(0))[0][:,s:s+PATCH_W]\n            label[i,0,0] = L_x - s\n            label[i,1,0] = R_x - s\n            label[i,0,1] = l\n            label[i,1,1] = r\n            if not self.VALID:\n                label[i] += 1 - np.random.randint(3,size=(2,2))\n                rot_img,label[i] = augment_image_and_centers(img[i].unsqueeze(0),label[i],center=(label[i].mean(0).tolist()))\n                img[i] = rot_img[0]\n\n#           Ideal heatmaps\n            m = idx_map - label[i].view(2,2,1,1).to(device)\n            m = (m*m).sum(1)\n            m = torch.exp(A*m)\n            mask[i] = m.sum(0)\n        \n        return img,mask","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tds = Sagittal_T1_axial_Dataset(train)\ntds.alpha=0.5","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image,mask = tds.__getitem__(np.random.randint(len(tds)))\nfor i in range(5):\n    plt.imshow(image[i].cpu() + .5*(mask[i].cpu() > TH))\n    plt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vds = Sagittal_T1_axial_Dataset(train,VALID=True)\nvds.alpha=0.5","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image,mask = vds.__getitem__(np.random.randint(len(vds)))\nfor i in range(5):\n    plt.imshow(image[i].cpu() + .5*(mask[i].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_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=ENCODER_NAME,\n            classes=1,\n            in_channels=1\n        ).to(device)\n\n    def forward(self,X):\n        x = self.UNet(X.view(-1,1,PATCH_H,PATCH_W))\n#       MinMaxScaling along the class plane to generate a heatmap\n        min_values = x.view(-1,PATCH_H*PATCH_W).min(-1)[0].view(-1,1,1,1)\n        max_values = x.view(-1,PATCH_H*PATCH_W).max(-1)[0].view(-1,1,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,PATCH_H,PATCH_W)","metadata":{},"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            heatmap,# Predictions\n            mask # Targets\n        ):\n#       Distance\n        mask = mask.view(-1,PATCH_H*PATCH_W)\n        heatmap = heatmap.view(-1,PATCH_H*PATCH_W)\n        D = 1 - ((mask*heatmap).sum(-1))**2/((mask*mask).sum(-1)*(heatmap*heatmap).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_count":null,"outputs":[]},{"cell_type":"code","source":"for f in FOLDS:\n    seed_everything(1337)\n    model = myUNet()\n    \n    tdf = train[train['fold'] != f]\n    vdf = train[train['fold'] == f]\n\n    tds = Sagittal_T1_axial_Dataset(tdf)\n    vds = Sagittal_T1_axial_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_axial_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_T1_axial_segmentation_'+str(f))\n-\nvds = Sagittal_T1_axial_Dataset(vdf,VALID=True)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for kk in range(10):\n    i = np.random.randint(len(vds))\n    print(i)\n    image,mask = vds.__getitem__(i)\n    for k in range(5):\n        fig, axes = plt.subplots(1, 2, figsize=(10,10))\n        axes[0].imshow((image.cpu() + (model(image).detach().cpu() > TH))[k])\n        axes[1].imshow((image.cpu() + .5*(mask.cpu() > TH))[k])\n        plt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del model,vdf,vds\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]}]}