{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":100869,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":84606,"modelId":108843}],"dockerImageVersionId":30761,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### description and info\n=> We took inspiration from the [deep spine paper](http://proceedings.mlr.press/v85/lu18a/lu18a.pdf) and here have created the unet model used for segmenting spines. \n\n=> The model here generates centroids and masks for every spine from T12-L1. According to the paper these centroids can be further used to prepare croppings of each disc. We have included the pretrained model weights in with this paper.\n\n=> There are a bunch of fine tuned post processing techniques that we use to improve the unet predictions, the end results can be seen in the bottom ssection of this notebook. \n\n=> We have visually inspected the output for first 100 patients and listed down the good and bad results. One can go over the inference section and inspect the results over patients throught dataset","metadata":{}},{"cell_type":"markdown","source":"### imports","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport cv2\nimport pydicom\nimport glob\nfrom tqdm import tqdm\nimport warnings\nfrom PIL import Image, ImageDraw\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nfrom torchvision import models\nfrom torch.nn.functional import relu\nfrom torch.utils.data import DataLoader, Dataset\nimport os\nimport copy\nimport time\nfrom collections import defaultdict\nimport skimage.morphology as morph\nfrom skimage.filters import threshold_otsu\nimport scipy.ndimage as ndi\nfrom skimage.segmentation import watershed\nfrom skimage.feature import peak_local_max, corner_peaks\nfrom skimage import measure\nfrom operator import itemgetter\nfrom scipy.interpolate import CubicSpline","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-10-03T18:19:25.172544Z","iopub.execute_input":"2024-10-03T18:19:25.173670Z","iopub.status.idle":"2024-10-03T18:19:32.591026Z","shell.execute_reply.started":"2024-10-03T18:19:25.173604Z","shell.execute_reply":"2024-10-03T18:19:32.589602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### constants","metadata":{}},{"cell_type":"code","source":"base_path = '/kaggle/input/manually-segmented-lumbar-spine-t12-s1/Lumbar Segmentation/'\nutility = ['train/','valid/','test/']\nimg_path = 'images/'\nlabels_path = 'labels/'\nimage_dim = 512\ncategories = ['null','L1', 'L2', 'L3', 'L4', 'L5', 'S1', 'T12']","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:32.593282Z","iopub.execute_input":"2024-10-03T18:19:32.593909Z","iopub.status.idle":"2024-10-03T18:19:32.600732Z","shell.execute_reply.started":"2024-10-03T18:19:32.593862Z","shell.execute_reply":"2024-10-03T18:19:32.599248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cat_colors = plt.cm.tab20(np.linspace(0,1,len(categories)))\nselect_cat_indices = [idx for idx in range(0,len(categories))]\nselect_cat_rgb_values =  np.array(cat_colors)[select_cat_indices]","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:32.602156Z","iopub.execute_input":"2024-10-03T18:19:32.602646Z","iopub.status.idle":"2024-10-03T18:19:32.617165Z","shell.execute_reply.started":"2024-10-03T18:19:32.602592Z","shell.execute_reply":"2024-10-03T18:19:32.615943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv')\nprint(\"Total Cases: \", len(train))","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:32.619803Z","iopub.execute_input":"2024-10-03T18:19:32.620366Z","iopub.status.idle":"2024-10-03T18:19:32.667304Z","shell.execute_reply.started":"2024-10-03T18:19:32.620310Z","shell.execute_reply":"2024-10-03T18:19:32.665480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### utility functions","metadata":{}},{"cell_type":"code","source":"def visualize(**images):\n    \"\"\"\n    Plot images in one row\n    \"\"\"\n    n_images = len(images)\n    plt.figure(figsize=(20,8))\n    for idx, (name, image) in enumerate(images.items()):\n        plt.subplot(1, n_images, idx + 1)\n        plt.xticks([]); \n        plt.yticks([])\n        # get title from the parameter names\n        plt.title(name.replace('_',' ').title(), fontsize=20)\n        plt.imshow(image)\n    plt.show()\n\ndef reverse_one_hot(img,scale_dim=512):\n    # adding base_arr to represent 0 as a category too\n#     base_arr = torch.as_tensor(np.zeros((1,scale_dim,scale_dim))).float()\n#     new_img = torch.cat((base_arr,img),dim=0)\n    return np.argmax(img,axis=0)\n\ndef colour_code_segmentation(image, label_values):\n    colour_codes = np.array(label_values)\n    img = image.astype(int)\n    x = colour_codes[img]\n    return x\n\ndef generate_mask(path,scale_dim=512,categories=categories):\n    # function to generate masks from polygon coordinates\n    # read polygon info from file\n    p_dict = get_polygon_dict(path)\n    # Create one mask layer for each category except null\n    mask_arr = np.zeros((len(categories),scale_dim,scale_dim))\n    for key in p_dict.keys():\n        img = Image.new('L', (scale_dim, scale_dim), 0)\n        if len(p_dict[key]):\n            ImageDraw.Draw(img).polygon(p_dict[key], outline=1, fill=1)\n        mask_arr[int(key)] = np.array(img)\n\n    return mask_arr\n\ndef get_polygon_dict(filePath,categories=categories):\n    # => returns list of coordinates for polygon corresponding to each label\n    # => the dict init is to ensure a empty list for categories which are \n    # not present in the labeling file\n\n    polygon_dict = dict((str(cat),[]) for cat in range(len(categories)))\n\n    f = open(filePath,\"r\")\n    # every line in file represents a class and the polygon corresponding to it\n    for x in f:\n        el_list = x.replace('\\n','').split(' ')\n        polygon_dict[str(int(el_list[0])+1)] = list(map(scale_n_typecast,el_list[1:]))\n    f.close()\n\n    return polygon_dict\n\ndef scale_n_typecast(x: str,scale_dim=512):\n    # 1 => converts string to float \n    # 2 => then scales it to the dimension of image\n    # 3 => rounds it to nearest integral value\n\n    return round(float(x)*scale_dim)","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:35.314995Z","iopub.execute_input":"2024-10-03T18:19:35.315448Z","iopub.status.idle":"2024-10-03T18:19:35.337001Z","shell.execute_reply.started":"2024-10-03T18:19:35.315406Z","shell.execute_reply":"2024-10-03T18:19:35.334589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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.status.busy":"2024-10-03T18:19:35.663913Z","iopub.execute_input":"2024-10-03T18:19:35.665204Z","iopub.status.idle":"2024-10-03T18:19:35.675058Z","shell.execute_reply.started":"2024-10-03T18:19:35.665149Z","shell.execute_reply":"2024-10-03T18:19:35.673523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process_n_visualise(img, gt_mask, pred):\n    visualize(\n        image = img,\n        ground_truth = colour_code_segmentation(np.argmax(gt_mask,axis=0),select_cat_rgb_values),\n        predicted_mask = colour_code_segmentation(np.argmax(pred,axis=0),select_cat_rgb_values)\n    )\n\ndef mean_blob_area(blob_arr,blob_count):\n    mean_size = 0\n    for i in range(1,blob_count+1):\n        mean_size += np.sum(blob_arr == i)\n    return mean_size/blob_count\n\ndef remove_small_components(blob_arr, threshold, blob_count):\n    labels = [0]\n    core_mask = np.zeros((blob_count+1,*blob_arr.shape))\n    for i in range(1,blob_count+1):\n        region = blob_arr == i\n        if np.sum(region) >= threshold:\n            core_mask[i] = region\n            labels.append(i)\n    comp_core_mask = np.zeros((len(labels),*blob_arr.shape))\n    ind = 1;\n    for i in range(1,len(labels)):\n        comp_core_mask[ind] = core_mask[labels[i]]\n        ind += 1\n    return comp_core_mask\n\ndef seperate_connected_vertebrae(mask, min_peak_gap, min_cluster_intensity):\n    # mask is assumed to be binary with only 1 comp in it\n    distance_map = ndi.distance_transform_edt(mask)\n    local_peaks = corner_peaks(distance_map, footprint=np.ones((3,3)), labels=mask, min_distance=min_peak_gap, threshold_rel=min_cluster_intensity)\n    m_ = np.zeros(distance_map.shape,dtype=bool)\n    m_[tuple(local_peaks.T)] = True\n    markers,_ = ndi.label(m_)\n    labels = watershed(-distance_map, markers, mask=mask)\n    return labels, len(local_peaks)\n\ndef add_labels_to_layer(core_layered_mask, labels, label_count):\n    one_hot_array = np.zeros((label_count,*labels.shape),dtype=int)\n    for i in range(1,label_count+1):\n        one_hot_array[i-1] = (labels == i).astype(int)\n    core_layered_mask = np.concatenate((core_layered_mask,one_hot_array),axis=0)\n    return core_layered_mask\n\ndef sort_n_label(values,spine_order):\n    total_labels = len(spine_order)\n    label_order = {}; centroids = np.zeros((total_labels,2))\n    sort_centroids = values[values[:,1].argsort()[::-1]]\n    for idx in range(0,len(sort_centroids)):\n        if idx < total_labels:\n            if int(sort_centroids[idx,0]) != 0:\n                label_order[spine_order[idx]] = int(sort_centroids[idx,0])\n            else :\n                label_order[spine_order[idx]] = -1\n            centroids[idx] = sort_centroids[idx,1:]\n        else: \n            break\n    return label_order, centroids\n\ndef calculate_centroids(layers,spine_order):\n    centroids = np.zeros((max(layers.shape[0],len(spine_order)),3))\n    for idx in range(1,layers.shape[0]):\n        centroids[idx-1] = np.concatenate((np.ones((1))*idx,measure.centroid(layers[idx,:,:])),axis=0)\n    return centroids","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:35.988166Z","iopub.execute_input":"2024-10-03T18:19:35.988684Z","iopub.status.idle":"2024-10-03T18:19:36.018127Z","shell.execute_reply.started":"2024-10-03T18:19:35.988640Z","shell.execute_reply":"2024-10-03T18:19:36.016696Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### wrapper class","metadata":{}},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n    def __init__(self,in_channels,out_channels):\n        super(DoubleConv, self).__init__()\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels,out_channels, kernel_size=3,padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels,out_channels, kernel_size=3,padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n        )\n    def forward(self,x):\n        return self.double_conv(x)\n    \nclass DownSampleBlock(nn.Module):\n    def __init__(self,in_channels, out_channels):\n        super(DownSampleBlock, self).__init__()\n        self.double_conv = DoubleConv(in_channels, out_channels)\n        self.down_sample = nn.MaxPool2d(kernel_size=2,stride=2)\n    \n    def forward(self, x):\n        skip_out = self.double_conv(x)\n        down_out = self.down_sample(skip_out)\n        return (down_out, skip_out)\n\nclass UpSampleBlock(nn.Module):\n    def __init__(self,in_channels,out_channels):\n        super(UpSampleBlock, self).__init__()\n        self.up_sample = nn.ConvTranspose2d(in_channels-out_channels,in_channels-out_channels,kernel_size=2,stride=2)\n        self.double_conv = DoubleConv(in_channels, out_channels)\n        \n    def forward(self, down_input, skip_input):\n        x = self.up_sample(down_input)\n        x = torch.cat([x,skip_input],dim=1)\n        return self.double_conv(x)\n    \nclass UNet(nn.Module):\n    def __init__(self,out_classes=8):\n        super(UNet, self).__init__()\n        \n        # Encoder\n        # input_dim 512x512x3\n        self.down_block1 = DownSampleBlock(3,64)\n        self.down_block2 = DownSampleBlock(64,128)\n        self.down_block3 = DownSampleBlock(128,256)\n        self.down_block4 = DownSampleBlock(256,512)\n        \n        self.double_conv = DoubleConv(512,1024)\n        #Decoder\n        self.up_block4 = UpSampleBlock(512+1024,512)\n        self.up_block3 = UpSampleBlock(256+512,256)\n        self.up_block2 = UpSampleBlock(256+128,128)\n        self.up_block1 = UpSampleBlock(128+64,64)\n        self.conv_last = nn.Conv2d(64,out_classes,kernel_size=1)\n    \n    def forward(self,x):\n        x, skip1_out = self.down_block1(x)\n        x, skip2_out = self.down_block2(x)\n        x, skip3_out = self.down_block3(x)\n        x, skip4_out = self.down_block4(x)\n        x = self.double_conv(x)\n        x = self.up_block4(x,skip4_out)\n        x = self.up_block3(x,skip3_out)\n        x = self.up_block2(x,skip2_out)\n        x = self.up_block1(x,skip1_out)\n        return self.conv_last(x)\n\n# model = UNet()","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:36.823597Z","iopub.execute_input":"2024-10-03T18:19:36.824100Z","iopub.status.idle":"2024-10-03T18:19:36.846079Z","shell.execute_reply.started":"2024-10-03T18:19:36.824055Z","shell.execute_reply":"2024-10-03T18:19:36.844103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Unet_Pred():\n    def __init__(self):\n        self.device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n        self.model = UNet().to(self.device)\n        if os.path.exists('/kaggle/input/unet-lumbar_segmentation/pytorch/default/1/best_model.pth'):\n            self.model = torch.load('/kaggle/input/unet-lumbar_segmentation/pytorch/default/1/best_model.pth', map_location=self.device)\n            print(\"model from previous session loaded\")\n        else:   \n            print(\"add model.pth to inputs and update the correct path here\")\n        self.post_proc = Post_Processing()\n        \n    def forward(self, img):\n        # takes a tensor input of 3,512,512\n        # the image should be normalized (0-1)\n        # returns a 8,512,512 np array of masks  and a 7,2 np array of centroids\n        self.model.eval()\n        pred = self.model(img)\n        return self.post_proc.forward(pred)","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:37.165686Z","iopub.execute_input":"2024-10-03T18:19:37.166137Z","iopub.status.idle":"2024-10-03T18:19:37.175522Z","shell.execute_reply.started":"2024-10-03T18:19:37.166097Z","shell.execute_reply":"2024-10-03T18:19:37.174250Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Post_Processing():\n    def __init__(\n        self,\n        proc_params = {\n            'small_comp_removal_scale' : 0.4,\n            'comp_sepration_scale' : 1.8,\n            'min_gap_scale' : 0.5,\n            'min_intensity_scale' : 0.3,\n            'erosion_scale' : 2\n        },\n        spine_order = ['S1','L5','L4','L3','L2','L1','T12'],\n        display_order = categories,\n        detail_vis = False,\n        visualise = False\n    ):\n        self.removal_scale = proc_params['small_comp_removal_scale']\n        self.sepration_scale = proc_params['comp_sepration_scale']\n        self.min_gap_scale = proc_params['min_gap_scale']\n        self.min_intensity_scale = proc_params['min_intensity_scale']\n        self.erosion_scale = proc_params['erosion_scale']\n        self.spine_order = spine_order\n        self.display_order = display_order\n        self.vis = visualise\n        self.detail_vis = detail_vis\n        \n    \n    def forward(self,data):\n        # change this when running on test data \n        # as there will be no gt_mask\n#         img, mask, pred = self.tensor_to_np(data)\n        pred = self.tensor_to_np(data)\n        core_mask = self.remove_small_components(pred)\n        core_layered_mask = self.seperate_combined_components(core_mask)\n        centroids, label_layer_map = self.anatomical_labeling(core_layered_mask)\n        \n        permute = []\n        for k in self.display_order:\n            if label_layer_map[k] != -1:\n                permute.append(label_layer_map[k])\n                \n        permuted_array = core_layered_mask[\n            permute,:,:\n        ]\n        \n        core_layered_mask = np.zeros((8,512,512))\n        core_layered_mask[:permuted_array.shape[0]] = permuted_array\n            \n        if(self.vis):\n            visualize(\n                img = img,\n                pred_mask = colour_code_segmentation(np.argmax(pred,axis=0),select_cat_rgb_values),\n                processed_mask = colour_code_segmentation(np.argmax(core_layered_mask, axis=0), select_cat_rgb_values),\n                manually_labeled = colour_code_segmentation(np.argmax(mask,axis=0),select_cat_rgb_values),\n            )\n        if(self.detail_vis):\n            visualize(\n                pred_mask = colour_code_segmentation(np.argmax(pred,axis=0),select_cat_rgb_values),\n                bin_mask = np.argmax(np.stack((pred[0,:,:],np.max(pred[1:,:,:],axis=0)), axis=0),axis=0),\n                small_comp = np.argmax(core_mask,axis=0),\n                seperation = np.argmax(core_layered_mask,axis=0),\n                processed_mask = colour_code_segmentation(np.argmax(core_layered_mask, axis=0), select_cat_rgb_values),\n                manually_labeled = colour_code_segmentation(np.argmax(mask,axis=0),select_cat_rgb_values),\n            )\n        return core_layered_mask, centroids\n    \n    def tensor_to_np(self, data):\n        # change this when running on test data \n        # as there will be no gt_mask\n        pred = data\n        \n#         [inputs, mask, pred] = data\n#         trial_img = inputs.detach().squeeze().cpu().permute(1,2,0).numpy()\n#         trial_mask = mask.detach().squeeze().cpu().numpy()\n        trial_pred = pred.detach().squeeze().cpu().numpy()\n#         return trial_img, trial_mask, trial_pred\n        return trial_pred\n\n    def remove_small_components(self, img):\n        bin_mask = np.argmax(np.stack((img[0,:,:],np.max(img[1:,:,:],axis=0)), axis=0),axis=0)\n        bin_mask = ndi.binary_fill_holes(bin_mask).astype(int)\n        bin_mask = morph.erosion(bin_mask,footprint=morph.disk(self.erosion_scale))\n        comp, label = ndi.label(bin_mask)\n        threshold_area = mean_blob_area(comp,label)*self.removal_scale\n        return remove_small_components(comp,threshold_area,label)\n\n    def seperate_combined_components(self, core_mask):\n        core_bin_mask = np.argmax(core_mask,axis=0)\n        threshold_area = mean_blob_area(core_bin_mask, core_mask.shape[0])*self.sepration_scale\n        core_layered_mask = np.zeros((1,*core_mask.shape[1:]))\n        for i in range(1,core_mask.shape[0]):\n            region = core_bin_mask == i\n            area = np.sum(region)\n            if area >= threshold_area:\n                labels, label_count = seperate_connected_vertebrae(region,int(self.min_gap_scale*np.sqrt(threshold_area)),self.min_intensity_scale)\n                core_layered_mask = add_labels_to_layer(core_layered_mask, labels, label_count)\n            else :\n                core_layered_mask = np.concatenate((core_layered_mask,np.expand_dims(region,axis=0)),axis=0)\n        return core_layered_mask\n    \n    def anatomical_labeling(self, core_layered_mask):\n        label_layer_map, centroids = sort_n_label(\n            calculate_centroids(core_layered_mask,self.spine_order),\n            self.spine_order\n        )\n        label_layer_map['null']=0\n        if (centroids[-1] == np.zeros((1,2))).all():\n            idx = 0\n            for i in range(centroids.shape[0]-1,-1,-1):\n                if (centroids[i] == np.zeros((1,2))).all():\n                    idx -= 1\n                else: \n                    break\n            if (-1*idx) <= 5:\n                centroids = self.spline_interpolation(centroids[:idx][::-1],-1*idx,centroids)\n        return centroids, label_layer_map\n    \n    def spline_interpolation(self, centroids, count, res):\n        n_pred = centroids.shape[0]\n        x_coords = np.array([c[:][1] for c in centroids])\n        y_coords = np.array([c[:][0] for c in centroids])\n        positions = np.arange(1, n_pred+1)  # Positions 1 (S1) to 6 (L1)\n\n        # Fit cubic splines to the known centroids\n        spline_x = CubicSpline(positions, x_coords, bc_type='natural')\n        spline_y = CubicSpline(positions, y_coords, bc_type='natural')\n        \n        # Estimate missing positions\n        print(f'interpolation count : {count}')\n        for i in range(0,count):\n            res[n_pred+i][1] = spline_x(0)\n            res[n_pred+i][0] = spline_y(0)\n            if res[n_pred+i][0] <0 or res[n_pred+i][1]<0:\n                res[n_pred+i:] = res[n_pred+i]\n                break\n\n        return res","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:37.587885Z","iopub.execute_input":"2024-10-03T18:19:37.588344Z","iopub.status.idle":"2024-10-03T18:19:37.628180Z","shell.execute_reply.started":"2024-10-03T18:19:37.588301Z","shell.execute_reply":"2024-10-03T18:19:37.626221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### data loading","metadata":{}},{"cell_type":"code","source":"part_1 = os.listdir('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images')\npart_1 = list(filter(lambda x: x.find('.DS') == -1, part_1))\ndf_meta_f = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv')","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:38.636093Z","iopub.execute_input":"2024-10-03T18:19:38.636552Z","iopub.status.idle":"2024-10-03T18:19:38.789280Z","shell.execute_reply.started":"2024-10-03T18:19:38.636511Z","shell.execute_reply":"2024-10-03T18:19:38.787815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"p1 = [(x, f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{x}\") for x in part_1]\nmeta_obj = { p[0]: { 'folder_path': p[1], \n                    'SeriesInstanceUIDs': [] \n                   } \n            for p in p1 }\n\nfor m in meta_obj:\n    meta_obj[m]['SeriesInstanceUIDs'] = list(\n        filter(lambda x: x.find('.DS') == -1, \n               os.listdir(meta_obj[m]['folder_path'])\n              )\n    )\n\n# grabs the correspoding series descriptions\nfor k in tqdm(meta_obj):\n    for s in meta_obj[k]['SeriesInstanceUIDs']:\n        if 'SeriesDescriptions' not in meta_obj[k]:\n            meta_obj[k]['SeriesDescriptions'] = []\n        try:\n            meta_obj[k]['SeriesDescriptions'].append(\n                df_meta_f[(df_meta_f['study_id'] == int(k)) & \n                (df_meta_f['series_id'] == int(s))]['series_description'].iloc[0])\n        except:\n            print(\"Failed on\", s, k)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:39.094035Z","iopub.execute_input":"2024-10-03T18:19:39.094686Z","iopub.status.idle":"2024-10-03T18:19:49.005538Z","shell.execute_reply.started":"2024-10-03T18:19:39.094635Z","shell.execute_reply":"2024-10-03T18:19:49.004107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta_obj[list(meta_obj.keys())[1]]","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:49.007637Z","iopub.execute_input":"2024-10-03T18:19:49.008168Z","iopub.status.idle":"2024-10-03T18:19:49.018833Z","shell.execute_reply.started":"2024-10-03T18:19:49.008125Z","shell.execute_reply":"2024-10-03T18:19:49.017311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# section to calculate mean masks, one will not need this during inference\n# mean_mask = np.zeros((8,512,512))\n# mean_centroids = np.zeros((7,2))\n# count = 0\n# sample_keys = []","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:34:49.362956Z","iopub.execute_input":"2024-10-03T18:34:49.363498Z","iopub.status.idle":"2024-10-03T18:34:49.370149Z","shell.execute_reply.started":"2024-10-03T18:34:49.363452Z","shell.execute_reply":"2024-10-03T18:34:49.368767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Inference and Visulisation of output","metadata":{}},{"cell_type":"code","source":"pre_process = Unet_Pred()","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:49.037440Z","iopub.execute_input":"2024-10-03T18:19:49.038004Z","iopub.status.idle":"2024-10-03T18:19:50.923460Z","shell.execute_reply.started":"2024-10-03T18:19:49.037929Z","shell.execute_reply":"2024-10-03T18:19:50.922052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"key = 73\npatient = train.iloc[key]\nptobj = meta_obj[str(patient['study_id'])]\nim_list_dcm = {}\n\nidx = ptobj['SeriesDescriptions'].index('Sagittal T2/STIR')\ni = ptobj['SeriesInstanceUIDs'][idx]\n\nim_list_dcm[i] = {'images': [], 'description': ptobj['SeriesDescriptions'][idx]}\nimages = glob.glob(f\"{ptobj['folder_path']}/{ptobj['SeriesInstanceUIDs'][idx]}/*.dcm\")\nimage_sorted = sorted(images, key=lambda x: int(x.split('/')[-1].replace('.dcm', '')))\n\ntarget_dcm = pydicom.dcmread(image_sorted[(len(image_sorted)-1)//2])\n\nresized_img = cv2.resize(target_dcm.pixel_array,dsize=(512,512), interpolation=cv2.INTER_CUBIC)\nresized_img = cv2.normalize(resized_img, None, alpha=0, beta=1, norm_type=cv2.NORM_MINMAX, dtype=cv2.CV_32F)\ntarget_img = cv2.merge((resized_img,resized_img,resized_img))\ntarget_img = torch.from_numpy(target_img).permute(2,0,1).float()\nprint(target_img.shape)\n\nmask, centroids = pre_process.forward(target_img.unsqueeze(dim=0))\n\nvisualize(\n    input_img = target_dcm.pixel_array,\n    mean_mask = colour_code_segmentation(reverse_one_hot(mask),select_cat_rgb_values)\n)","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:24:42.570030Z","iopub.execute_input":"2024-10-03T18:24:42.570492Z","iopub.status.idle":"2024-10-03T18:24:49.761359Z","shell.execute_reply.started":"2024-10-03T18:24:42.570449Z","shell.execute_reply":"2024-10-03T18:24:49.759955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(target_img.permute(1,2,0))\nfor c in centroids:\n    plt.plot(c[1],c[0],'ro')\nplt.plot(centroids[:,1],centroids[:,0],'y')\nplt.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:24:21.345962Z","iopub.execute_input":"2024-10-03T18:24:21.346481Z","iopub.status.idle":"2024-10-03T18:24:21.586350Z","shell.execute_reply.started":"2024-10-03T18:24:21.346440Z","shell.execute_reply":"2024-10-03T18:24:21.584912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# mean_centroids += centroids\n# mean_mask += mask\n# count += 1\n# sample_keys.append(key)\n# key+=1","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:51.589082Z","iopub.status.idle":"2024-10-03T18:19:51.589602Z","shell.execute_reply.started":"2024-10-03T18:19:51.589359Z","shell.execute_reply":"2024-10-03T18:19:51.589384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(f'{key}\\n{sample_keys}')","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:51.591367Z","iopub.status.idle":"2024-10-03T18:19:51.591908Z","shell.execute_reply.started":"2024-10-03T18:19:51.591638Z","shell.execute_reply":"2024-10-03T18:19:51.591661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"106\n[0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 29, 30, 31, 32, 33, 34, 35, 36, 37, 38, 39, 40, 41, 42, 43, 44, 45, 46, 47, 48, 49, 50, 51, 52, 53, 54, 55, 56, 58, 59, 60, 61, 62, 64, 65, 66, 67, 68, 69, 70, 71, 72, 73, 74, 75, 76, 77, 78, 79, 81, 82, 83, 84, 85, 86, 87, 88, 89, 90, 91, 92, 93, 94, 96, 97, 98, 99, 100, 101, 102, 103, 104, 105]","metadata":{}},{"cell_type":"code","source":"# bad_keys.append(key)\n# key+=1","metadata":{"execution":{"iopub.status.busy":"2024-10-03T18:19:51.594124Z","iopub.status.idle":"2024-10-03T18:19:51.594620Z","shell.execute_reply.started":"2024-10-03T18:19:51.594391Z","shell.execute_reply":"2024-10-03T18:19:51.594415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"[11, 14, 28, 57, 63, 80, 93, 95]","metadata":{}},{"cell_type":"code","source":"# bin_mean_mask = np.zeros((8,512,512))\n# for i in range(0,mean_mask.shape[0]):\n#     mean = mean_mask[i].sum()/(mean_mask[i]!=0).sum().astype(float)\n#     bin_mean_mask[i] = np.where(mean_mask[i] > mean, 1, 0)","metadata":{"execution":{"iopub.status.busy":"2024-08-31T12:00:06.879146Z","iopub.execute_input":"2024-08-31T12:00:06.880082Z","iopub.status.idle":"2024-08-31T12:00:06.905303Z","shell.execute_reply.started":"2024-08-31T12:00:06.880035Z","shell.execute_reply":"2024-08-31T12:00:06.904059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# visualize(\n#     input_img = target_dcm.pixel_array,\n#     mean_mask = colour_code_segmentation(reverse_one_hot(bin_mean_mask),select_cat_rgb_values)\n# )","metadata":{"execution":{"iopub.status.busy":"2024-08-31T12:00:25.629354Z","iopub.execute_input":"2024-08-31T12:00:25.629918Z","iopub.status.idle":"2024-08-31T12:00:26.262127Z","shell.execute_reply.started":"2024-08-31T12:00:25.629874Z","shell.execute_reply":"2024-08-31T12:00:26.260627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### mean mask\nthis is the mean of mask generated for each vertebrate across 7 spines (T12-s1) \n\n![Mean Mask](https://imgur.com/a/l5KovAo)","metadata":{}}]}