{"metadata":{"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":30746,"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]\nPATCH_SIZE = 512\npatch_size = 128\nLmax = 36\nBS = 16\nEPOCHS = 5","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":"df = df_coor[\n    df_coor.condition.isin(\n        [\n            'Left Subarticular Stenosis',\n            'Right Subarticular Stenosis'\n        ]\n    )\n][['study_id','series_id','instance_number','level']].drop_duplicates()\ndf_min = df.groupby(['study_id','series_id','level'],sort=True).min().reset_index().rename(columns={'instance_number':'instance_number_min'})\ndf_max = df.groupby(['study_id','series_id','level'],sort=True).max().reset_index().rename(columns={'instance_number':'instance_number_max'})\ndf_min['instance_number_max'] = df_max['instance_number_max']\ndf = df_min","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"centers = {}\nfor i in range(len(df)):\n    row = df.iloc[i]\n    centers[row['study_id']]={}\nfor i in range(len(df)):\n    row = df.iloc[i]\n    centers[row['study_id']][row['series_id']]={}\nfor i in range(len(df)):\n    row = df.iloc[i]\n    centers[row['study_id']][row['series_id']][row['level']] = [row['instance_number_min'],row['instance_number_max']]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = df.groupby(['study_id','series_id']).count().reset_index()[['study_id','series_id']]\ndf.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# We'll use 0 as missing value\nv = np.zeros((len(df),10)).astype(int)\nfor i in range(len(df)):\n    row = df.iloc[i]\n    for level in centers[row['study_id']][row['series_id']]:\n        v_min,v_max = centers[row['study_id']][row['series_id']][level]\n        v[i,{'L1/L2':0,'L2/L3':2,'L3/L4':4,'L4/L5':6,'L5/S1':8}[level]] = v_min\n        v[i,{'L1/L2':1,'L2/L3':3,'L3/L4':5,'L4/L5':7,'L5/S1':9}[level]] = v_max","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[[\n    'L1L2_min','L1L2_max',\n    'L2L3_min','L2L3_max',\n    'L3L4_min','L3L4_max',\n    'L4L5_min','L4L5_max',\n    'L5S1_min','L5S1_max'\n]] = v\ndf.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[(df[[\n    'L1L2_min','L1L2_max',\n    'L2L3_min','L2L3_max',\n    'L3L4_min','L3L4_max',\n    'L4L5_min','L4L5_max',\n    'L5S1_min','L5S1_max'\n]] == 0).sum(1)>0].reset_index(drop=True).tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['flip'] = False\nfdf = df.copy()\nfdf['flip'] = True\ndf = pd.concat([df,fdf]).reset_index(drop=True)\ndf.tail()","metadata":{},"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":"level_train = df.merge(train[['study_id','fold']],left_on='study_id',right_on='study_id')\nlevel_train.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open('axial_centers.pkl', 'rb') as f:\n    coord = pickle.load(f)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"I'll try to train a vertebrate level locator for Axial T2.","metadata":{}},{"cell_type":"code","source":"class Axial_ViT_Dataset(Dataset):\n    def __init__(self, df, f, VALID=False, INFERENCE=False, alpha=0):\n        self.data = df\n        self.f = f\n        self.VALID = VALID\n        self.INFERENCE = INFERENCE\n        self.resize = torchvision.transforms.Resize((PATCH_SIZE,PATCH_SIZE),antialias=True)\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        slices = list(np.arange(len(images)))\n#       array = np.array([pydicom.dcmread(img).ImageOrientationPatient for img in images]).mean(0)\n#       kk = np.argmin(array[:3] + array[3:])\n#       slices.sort(key=lambda k:pydicom.dcmread(images[k]).ImagePositionPatient[kk])\n        slices.sort(key=lambda k:int(images[k].split('/')[-1].replace('.dcm','')))\n        instance_numbers = torch.as_tensor([int(x.split('/')[-1].replace('.dcm','')) for x in images])[slices]\n        images = np.array(images)[slices]\n        D = len(images)\n\n        c = coord[self.f][row['study_id']][row['series_id']].clone()\n        c[c < 64] = torch.nan\n        c[c > 512 - 64] = torch.nan\n        if D > Lmax:\n            slices = np.rint(torch.arange(Lmax)*D/Lmax).long()\n            images = images[slices]\n            c = c[slices]\n        elif D < Lmax:\n            N = Lmax//D + 1\n            slices = torch.repeat_interleave(torch.arange(D), N) \n            slices = slices[np.rint(torch.arange(Lmax)*(D*N)/Lmax).long()]\n            images = images[slices]\n            c = c[slices]\n\n        images = [torch.as_tensor(pydicom.dcmread(img).pixel_array.astype('float32')) for img in images]\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 = self.resize(image/image.max()).float().unsqueeze(1).to(device)\n\n        c = torch.nanmean(c, dim=1)\n        mask = torch.isnan(c)\n        c_mean = torch.nanmean(c,0)\n        c[mask[:,0],0] = c_mean[0]\n        c[mask[:,1],1] = c_mean[1]\n        c = c.long()\n        image = torch.stack([\n            image[\n                i,\n                :,\n                c[i,1]-patch_size//2:c[i,1]+patch_size-patch_size//2,\n                c[i,0]-patch_size//2:c[i,0]+patch_size-patch_size//2\n            ] for i in range(len(images))\n        ])\n\n        if self.INFERENCE:\n            return image.to(device)\n        else:\n            indices = np.arange(len(images))\n            label_min = -torch.ones(5).int()\n            label_max = -torch.ones(5).int()\n            label = -torch.ones(5).float()\n            values_min = torch.as_tensor(row[['L1L2_min','L2L3_min','L3L4_min','L4L5_min','L5S1_min']].values.astype(int)).view(-1)\n            values_max = torch.as_tensor(row[['L1L2_max','L2L3_max','L3L4_max','L4L5_max','L5S1_max']].values.astype(int)).view(-1)\n            mask = values_min != 0\n            label_min[mask] = torch.as_tensor(np.array([indices[l == instance_numbers] for l in values_min[mask]])).view(-1)\n            label_max[mask] = torch.as_tensor(np.array([indices[l == instance_numbers] for l in values_max[mask]])).view(-1)\n            label[mask] = ((label_min + label_max)/2)[mask]\n            label = label*Lmax/D\n            if row['flip']:\n                image = image.flip(0)\n                c = c.flip(0)\n                label[label != -1] = Lmax - 1 - label[label != -1]\n            return image.to(device),label.to(device)","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=\"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) # Bug, I've been MinMaxScaling with the wrong values\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_count":null,"outputs":[]},{"cell_type":"code","source":"class SinusoidalPosEmb(nn.Module):\n    def __init__(self, dim=16, M=10000):\n        super().__init__()\n        self.dim = dim\n        self.M = M\n\n    def forward(self, x):\n        device = x.device\n        half_dim = self.dim // 2\n        emb = math.log(self.M) / half_dim\n        emb = torch.exp(torch.arange(half_dim, device=device) * (-emb))\n        emb = x[...,None] * emb[None,...]\n        emb = torch.cat((emb.sin(), emb.cos()), dim=-1)\n        return emb\n\nclass Axial_ViT(nn.Module):\n    def __init__(self, ENCODER, dim=512, depth=24, head_size=64, **kwargs):\n        super().__init__()\n        self.ENCODER = ENCODER\n        self.AvgPool = nn.AdaptiveAvgPool2d(output_size=1).to(device)\n        self.slices_enc = nn.Parameter(SinusoidalPosEmb(dim)(torch.arange(Lmax, device=device).unsqueeze(0)))\n        self.slices_transformer = nn.TransformerEncoder(\n                nn.TransformerEncoderLayer(d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim,\n                dropout=0.1, activation=nn.GELU(), batch_first=True, norm_first=True, device=device), depth)\n        self.proj_out = nn.Linear(dim,5).to(device)\n    \n    def forward(self, x):\n        x = self.ENCODER(x.view(-1,1,patch_size,patch_size))[-1]\n        x = self.AvgPool(x)\n        x = x.view(-1,Lmax,512)\n        x = x + self.slices_enc\n        x = self.slices_transformer(x)#,src_key_padding_mask=slices_mask)\n        x = self.proj_out(x.view(-1,512)).view(-1,Lmax,5).permute(0,2,1)\n#       MinMaxScaling along the class plane to generate a heatmap\n        min_values = x.min(-1)[0].view(-1,5,1)\n        max_values = x.max(-1)[0].view(-1,5,1)\n        x = (x - min_values)/(max_values - min_values)\n        return x","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":"idx_map = torch.arange(Lmax).to(device)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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        available = t >= 0\n        mask_true = t[available]\n        mask_pred = y[available]\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/16])\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,1) - idx_map.view(1,-1)\n        mask = mask*mask\n        mask = torch.exp(A*mask)*K\n\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_count":null,"outputs":[]},{"cell_type":"code","source":"df = level_train\nfor f in FOLDS:\n    seed_everything(SEED)\n    seg_model = torch.load('Axial_T2_segmentation_'+str(f))\n    model = Axial_ViT(seg_model.UNet.encoder)\n    \n    tdf = df[df.fold != f]\n    vdf = df[df.fold == f]\n    tds = Axial_ViT_Dataset(tdf,f)\n    vds = Axial_ViT_Dataset(vdf,f,VALID=True)\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    dls = DataLoaders(tdl,vdl)\n\n    n_iter = len(tds)//BS\n\n    learn = Learner(\n        dls,\n        model,\n        loss_func=myLoss(),\n        cbs=[\n            ShowGraphCallback(),\n            GradientClip(3.0)\n        ]\n    )\n    learn.fit_one_cycle(EPOCHS, lr_max=5e-4, wd=0.05, pct_start=0.02)\n    torch.save(model,'axial_T2_levels_'+str(f))\n    del model,seg_model,tdf,vdf,tds,vds,tdl,vdl,dls,learn\n    gc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = torch.load('axial_T2_levels_1')\ndf = level_train\nvdf = df[df.fold == 1]\nvds = Axial_ViT_Dataset(vdf,1,VALID=True)\nvdl = torch.utils.data.DataLoader(vds, batch_size=BS, shuffle=False)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"levels = []\ny_true = []\nwith torch.no_grad():\n    for X,l in tqdm(vdl):\n        levels = levels + torch.argmax((model(X)+model(X.flip(1)).flip(-1)),-1).cpu().tolist()\n        y_true = y_true + l.cpu().tolist()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_true = np.array(y_true).flatten()\nlevels = np.array(levels).flatten()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(y_true,levels,'.')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for X,l in tqdm(vdl):\n    with torch.no_grad():\n            levels = (model(X)+model(X.flip(1)).flip(-1)).cpu()/2\n            x0 = torch.argmax(levels,-1)\n            break","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from scipy.optimize import curve_fit\n\ndef func(x, x0, b):\n    r = x - x0\n    return np.exp(-b*r*r)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x = np.linspace(0, 35, 1000)\nindices = np.arange(Lmax)\ncolors = ['r','lime','blue','yellow','cyan']\nref_preds = []\nfor i in range(len(levels)):\n    for k in range(len(levels[i])):\n        plt.plot(levels[i][k],'.',color=colors[k])\n        popt, pcov = curve_fit(func, indices, levels[i][k],[x0[i][k],1])\n        ref_preds = ref_preds + list(func(np.arange(Lmax), *popt))\n        plt.plot(x, func(x, *popt),color=colors[k])\n    for k in range(5): plt.axvline(x=l[i,k].cpu().detach(),c=colors[k])\n    plt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ref_preds = np.array(ref_preds).reshape(-1,5,Lmax)\nassigned_levels = np.argmax(ref_preds,-2)\nassigned_levels[:10]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(assigned_levels)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del X,l\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = 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').sort_values('fold').reset_index(drop=True)\ntrain.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for f in [1,2,3,4,5]:\n    model = torch.load('axial_T2_levels_'+str(f))\n    model.eval()\n    dl = torch.utils.data.DataLoader(Axial_ViT_Dataset(train,f,VALID=True,INFERENCE=True), batch_size=BS, shuffle=False)\n\n    popt = []\n    with torch.no_grad():\n        for X in tqdm(dl):\n            y_pred = (model(X)+model(X.flip(1)).flip(-1)).cpu()/2\n            x0 = torch.argmax(y_pred,-1)\n\n            for i in range(len(y_pred)):\n                for k in range(5):\n                    popt = popt + list(curve_fit(func, indices, y_pred[i][k],[x0[i][k],1],maxfev = 9999)[0])\n\n    \n    popt = torch.as_tensor(popt).view(-1,5,2).permute(0,2,1).reshape(-1,2*5)\n    mask = (popt[:,:5] < 0) + (popt[:,:5] > Lmax - 1)\n    popt[:,:5][mask] = torch.nan\n    popt[:,5:][mask] = torch.nan\n    train[[\n        'L1L2_x0','L2L3_x0','L3L4_x0','L4L5_x0','L5S1_x0',\n        'L1L2_b','L2L3_b','L3L4_b','L4L5_b','L5S1_b'\n    ]] = popt\n\n    train[[\n        'study_id','fold','series_id',\n        'L1L2_x0','L2L3_x0','L3L4_x0','L4L5_x0','L5S1_x0',\n        'L1L2_b','L2L3_b','L3L4_b','L4L5_b','L5S1_b'\n    ]].to_csv('axial_T2_levels_'+str(f)+'.csv',index=False)\n    del model,dl,popt\n    gc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"v = train[['L1L2_b','L2L3_b','L3L4_b','L4L5_b','L5S1_b']].values\nplt.boxplot(v[~np.isnan(v)])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.percentile(v[~np.isnan(v)],99)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train[(train[['L1L2_b','L2L3_b','L3L4_b','L4L5_b','L5S1_b']]>np.percentile(v[~np.isnan(v)],99)).sum(1)>0]","metadata":{},"execution_count":null,"outputs":[]}]}