{"metadata":{"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"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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\nimport sklearn\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":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2024-07-22T15:36:47.391937Z","iopub.status.busy":"2024-07-22T15:36:47.389385Z","iopub.status.idle":"2024-07-22T15:37:02.967144Z","shell.execute_reply":"2024-07-22T15:37:02.965578Z","shell.execute_reply.started":"2024-07-22T15:36:47.391871Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"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\npatch_size = 64\nBS = 16\nEPOCHS = 2","metadata":{"execution":{"iopub.execute_input":"2024-07-22T15:37:25.201958Z","iopub.status.busy":"2024-07-22T15:37:25.201502Z","iopub.status.idle":"2024-07-22T15:37:25.208848Z","shell.execute_reply":"2024-07-22T15:37:25.207312Z","shell.execute_reply.started":"2024-07-22T15:37:25.201926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spinal = [\n    'spinal_canal_stenosis_l1_l2',\n    'spinal_canal_stenosis_l2_l3',\n    'spinal_canal_stenosis_l3_l4',\n    'spinal_canal_stenosis_l4_l5',\n    'spinal_canal_stenosis_l5_s1'\n]\ncoor = [\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":"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\n    \ndef augment_image(image):\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    )\n    return image\n\ndef my_collate_fn(data):\n    collation = [torch.cat(s) for s in zip(*data)]\n    return collation\n# https://www.kaggle.com/code/abhinavsuri/anatomy-image-visualization-overview-rsna-raids\n# Function to display images\ndef display_images(images, title, max_images_per_row=4):\n    # Calculate the number of rows needed\n    num_images = len(images)\n    num_rows = (num_images + max_images_per_row - 1) // max_images_per_row  # Ceiling division\n\n    # Create a subplot grid\n    fig, axes = plt.subplots(num_rows, max_images_per_row, figsize=(5, 1.5 * num_rows))\n    \n    # Flatten axes array for easier looping if there are multiple rows\n    if num_rows > 1:\n        axes = axes.flatten()\n    else:\n        axes = [axes]  # Make it iterable for consistency\n\n    # Plot each image\n    for idx, image in enumerate(images):\n        ax = axes[idx]\n        ax.imshow(image, cmap='gray')  # Assuming grayscale for simplicity, change cmap as needed\n        ax.axis('off')  # Hide axes\n\n    # Turn off unused subplots\n    for idx in range(num_images, len(axes)):\n        axes[idx].axis('off')\n    fig.suptitle(title, fontsize=16)\n\n    plt.tight_layout()","metadata":{"execution":{"iopub.execute_input":"2024-07-22T15:37:02.982642Z","iopub.status.busy":"2024-07-22T15:37:02.982119Z","iopub.status.idle":"2024-07-22T15:37:03.364342Z","shell.execute_reply":"2024-07-22T15:37:03.363214Z","shell.execute_reply.started":"2024-07-22T15:37:02.982516Z"},"trusted":true},"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":"S = S[S['x_mean_fraction'] > .8]\nS.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S['instance_number'] = S['instance_number'] - 1\nS.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coordinates = {}\nfor study_id,df in S.groupby('study_id'):\n    coordinates[study_id] = {}\nfor (study_id,series_id),df in tqdm(S.groupby(['study_id','series_id'])):\n    coordinates[study_id][series_id] = {\n                'L1/L2':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'instance_number':torch.nan\n                },\n                    'L2/L3':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'instance_number':torch.nan\n                },\n                'L3/L4':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'instance_number':torch.nan\n                },\n                'L4/L5':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'instance_number':torch.nan\n                },\n                'L5/S1':{\n                    'x':torch.nan,\n                    'y':torch.nan,\n                    'instance_number':torch.nan\n                }\n    }\n    \n    for i in range(len(df)):\n        row = df.iloc[i]\n        coordinates[row['study_id']][row['series_id']][row['level']]['x'] = row['x']\n        coordinates[row['study_id']][row['series_id']][row['level']]['y'] = row['y']\n        coordinates[row['study_id']][row['series_id']][row['level']]['instance_number'] = row['instance_number']","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S =  S[[\n    'study_id',\n    'series_id'\n]].groupby([\n    'study_id',\n    'series_id'\n]).count().reset_index()\nS.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"v = np.zeros((len(S),15))\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']]:\n        v[i,k:k+3] = list(coordinates[row['study_id']][row['series_id']][level].values())\n        k += 3","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S[[\n    'x_L1L2',\n    'y_L1L2',\n    'i_L1L2',\n    'x_L2L3',\n    'y_L2L3',\n    'i_L2L3',\n    'x_L3L4',\n    'y_L3L4',\n    'i_L3L4',\n    'x_L4L5',\n    'y_L4L5',\n    'i_L4L5',\n    'x_L5S1',\n    'y_L5S1',\n    'i_L5S1'\n]] = v\nS.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S = S.merge(train,left_on='study_id',right_on='study_id')\nS.tail()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask = S[[\n    'x_L1L2',\n    'x_L2L3',\n    'x_L3L4',\n    'x_L4L5',\n    'x_L5S1'\n]].isna().values\nmask += S[[\n    'y_L1L2',\n    'y_L2L3',\n    'y_L3L4',\n    'y_L4L5',\n    'y_L5S1'\n]].isna().values\nmask += S[[\n    'i_L1L2',\n    'i_L2L3',\n    'i_L3L4',\n    'i_L4L5',\n    'i_L5S1'\n]].isna().values\nmask = mask > 0\nmask[-5:]\nv = S[spinal].values\nv[mask] = 'UNK'\nS[spinal] = v","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"S.groupby('fold').count()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Sagittal_T2_spine_discriminator_Dataset(Dataset):\n    def __init__(self, df, VALID=False, P=patch_size):\n        self.data = df\n        self.VALID = VALID\n        self.P = P\n        self.resize = torchvision.transforms.Resize((PATCH_H,PATCH_W),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        \n        sample = TRAIN_PATH + str(row['study_id']) + '/'+str(row['series_id'])\n\n        images = [x.replace('\\\\','/') for x in glob.glob(sample+'/*.dcm')]\n        images.sort(reverse=False, key=lambda x: int(x.split('/')[-1].replace('.dcm', '')))\n\n        instance_numbers = row[[\n            'i_L1L2',\n            'i_L2L3',\n            'i_L3L4',\n            'i_L4L5',\n            'i_L5S1'\n        ]].values\n\n        image = torch.stack([\n            torch.as_tensor(pydicom.dcmread(x).pixel_array.astype(np.float32)) for x in images\n        ]).float().to(device)\n        image = image/image.max()\n        D,H,W = image.shape\n\n        c = torch.as_tensor([x for x in row[coor]]).view(5,2).float()\n        missing = c.isnan().sum(1) > 0\n        c[missing] = torch.as_tensor([H/2,W/2])\n\n        if H > W:\n            d = W\n            h = (H - d)//2\n            image = image[:,h:h+d]\n            c[:,1] -= h\n            H = W\n        elif H < W:\n            d = H\n            w = (W - d)//2\n            image = image[:,:,w:w+d]\n            c[:,0] -= w\n            W = H\n\n        image = self.resize(image)\n        image = nn.functional.pad(image,[self.P]*4,'reflect')\n        c[:,1] = c[:,1]*PATCH_H/H + self.P\n        c[:,0] = c[:,0]*PATCH_W/W + self.P\n        c = c.long()\n\n        crops = torch.stack([\n            image[\n                :,\n                xy[1]-self.P:xy[1]+self.P,\n                xy[0]-self.P:xy[0]+self.P\n            ] for xy in c\n        ])\n\n        image = torch.zeros(5,D,2*self.P,2*self.P).to(device)\n        label = torch.zeros(5,D).long().to(device) - 100\n        for i in range(5):\n            if ~missing[i]:\n                instance_number = instance_numbers[i].astype(int)\n                pickeable = torch.ones(D).bool()\n                pickeable[instance_number] = False\n                image[i,0] = crops[i,instance_number]\n                k = 1\n                if instance_number > 0:\n                    pickeable[instance_number-1] = False\n                    image[i,1] = crops[i,instance_number-1]\n                    k += 1\n                if instance_number < D - 1:\n                    pickeable[instance_number+1] = False\n                    image[i,2] = crops[i,instance_number+1]\n                    k += 1\n                if instance_number > 1: pickeable[instance_number-2] = False\n                if instance_number < D - 2: pickeable[instance_number+2] = False\n                \n                pickeable = torch.arange(D)[pickeable]\n                picked = pickeable\n                image[i,k:len(picked)+k] = crops[i,picked]\n\n                label[i,:k] = 1\n                label[i,k:len(picked)+k] = 0\n\n                if not self.VALID:\n                    image[i] = augment_image(image[i].reshape(-1,2*self.P,2*self.P)).reshape(D,2*self.P,2*self.P)\n            \n        image = image[:,:,self.P//2:self.P//2+self.P,self.P//2:self.P//2+self.P]\n\n        mask = label != -100\n        image = image[mask]\n        label = label[mask]\n\n        return image,label","metadata":{"execution":{"iopub.execute_input":"2024-07-22T15:37:03.377957Z","iopub.status.busy":"2024-07-22T15:37:03.376752Z","iopub.status.idle":"2024-07-22T15:37:03.413904Z","shell.execute_reply":"2024-07-22T15:37:03.412746Z","shell.execute_reply.started":"2024-07-22T15:37:03.377915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = Sagittal_T2_spine_discriminator_Dataset(S)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample,label = ds.__getitem__(np.random.randint(len(ds)))\nprint(label)\ndisplay_images(sample[label == 1].cpu(),'Positives',max_images_per_row=5)\ndisplay_images(sample[label == 0].cpu(),'Negatives',max_images_per_row=5)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Sagittal_T2_spine_Discriminator(nn.Module):\n    def __init__(self, dim=512):\n        super().__init__()\n        CNN = torchvision.models.resnet18(weights='DEFAULT')\n        W = nn.Parameter(CNN.conv1.weight.sum(1, keepdim=True))\n        CNN.conv1 = nn.Conv2d(1, patch_size, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n        CNN.conv1.weight = W\n        CNN.fc = nn.Identity()\n        self.emb = CNN.to(device)\n        self.proj_out = nn.Linear(dim,2).to(device)\n    \n    def forward(self, x):        \n        x = self.emb(x.view(-1,1,patch_size,patch_size))\n        x = self.proj_out(x.view(-1,512))\n        return x","metadata":{"execution":{"iopub.execute_input":"2024-07-22T15:37:03.416625Z","iopub.status.busy":"2024-07-22T15:37:03.415818Z","iopub.status.idle":"2024-07-22T15:37:03.443753Z","shell.execute_reply":"2024-07-22T15:37:03.442621Z","shell.execute_reply.started":"2024-07-22T15:37:03.416584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"negatives = 116965\npositives = 29207\ntotal = negatives + positives\n\ndef myLoss(preds,target):\n    target = target.view(-1)\n    return nn.CrossEntropyLoss(weight=torch.as_tensor([total/(negatives*2),total/(positives*2)]).to(device))(preds,target)","metadata":{"execution":{"iopub.execute_input":"2024-07-22T15:37:03.446270Z","iopub.status.busy":"2024-07-22T15:37:03.445500Z","iopub.status.idle":"2024-07-22T15:37:03.459352Z","shell.execute_reply":"2024-07-22T15:37:03.457940Z","shell.execute_reply.started":"2024-07-22T15:37:03.446231Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for f in FOLDS:\n    seed_everything(SEED)\n    model = Sagittal_T2_spine_Discriminator()\n    \n    df = S\n    tdf = df[df.fold != f]\n    vdf = df[df.fold == f]\n    tds = Sagittal_T2_spine_discriminator_Dataset(tdf)\n    vds = Sagittal_T2_spine_discriminator_Dataset(vdf,VALID=True)\n    tdl = torch.utils.data.DataLoader(\n        tds,\n        batch_size=BS,\n        shuffle=True,\n        drop_last=True,\n        collate_fn=my_collate_fn\n    )\n    vdl = torch.utils.data.DataLoader(\n        vds,\n        batch_size=BS,\n        shuffle=False,\n        collate_fn=my_collate_fn\n    )\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=1e-3, wd=0.05, pct_start=0.02)\n    torch.save(model,'Sagittal_T2_spine_discriminator_'+str(f))\n    del model,df,tdf,vdf,tds,vds,tdl,vdl,dls,learn\n    gc.collect()","metadata":{"execution":{"iopub.execute_input":"2024-07-22T15:39:58.974867Z","iopub.status.busy":"2024-07-22T15:39:58.974392Z","iopub.status.idle":"2024-07-22T15:40:10.837835Z","shell.execute_reply":"2024-07-22T15:40:10.835528Z","shell.execute_reply.started":"2024-07-22T15:39:58.974828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_pred = []\ny_true = []\nfor f in FOLDS:\n    model = torch.load('Sagittal_T2_spine_discriminator_'+str(f)) \n    df = S\n    vdf = df[df.fold == f]\n    vds = Sagittal_T2_spine_discriminator_Dataset(vdf,VALID=True)\n    vdl = torch.utils.data.DataLoader(\n        vds,\n        batch_size=BS,\n        shuffle=False,\n        collate_fn=my_collate_fn\n    )\n    with torch.no_grad():\n        for images,target in tqdm(vdl):\n            target = target.view(-1).tolist()\n            preds = model(images).argmax(-1).tolist()\n            y_pred = y_pred + preds\n            y_true = y_true + target\n\n    del model,df,vdf,vds,vdl\n    gc.collect()\n\nsklearn.metrics.confusion_matrix(y_true, y_pred)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"negatives = sum([y == 0 for y in y_true])\npositives = sum([y == 1 for y in y_true])\ntotal = negatives + positives\nprint(negatives,positives,total)","metadata":{},"execution_count":null,"outputs":[]}]}