{"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":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9311247,"sourceType":"datasetVersion","datasetId":5639130},{"sourceId":100869,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":84606,"modelId":108843}],"dockerImageVersionId":30761,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### imports","metadata":{}},{"cell_type":"code","source":"import os\nos.makedirs('/root/.config/kaggle/')\n!echo '{\"username\":\"coopermini\",\"key\":\"91d48a9932f3d031df2b6c7129a8041e\"}' > /root/.config/kaggle/kaggle.json\n!chmod 600 ~/.config/kaggle/kaggle.json","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:14.603285Z","iopub.execute_input":"2024-09-04T22:23:14.603839Z","iopub.status.idle":"2024-09-04T22:23:16.939573Z","shell.execute_reply.started":"2024-09-04T22:23:14.603774Z","shell.execute_reply":"2024-09-04T22:23:16.937898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport math\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 kaggle\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 json\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-09-04T22:23:16.942794Z","iopub.execute_input":"2024-09-04T22:23:16.943411Z","iopub.status.idle":"2024-09-04T22:23:24.497219Z","shell.execute_reply.started":"2024-09-04T22:23:16.943348Z","shell.execute_reply":"2024-09-04T22:23:24.495810Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### constants","metadata":{}},{"cell_type":"code","source":"categories = ['null','L1', 'L2', 'L3', 'L4', 'L5', 'S1', 'T12']\ncat_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-09-04T22:23:24.499095Z","iopub.execute_input":"2024-09-04T22:23:24.499676Z","iopub.status.idle":"2024-09-04T22:23:24.509337Z","shell.execute_reply.started":"2024-09-04T22:23:24.499635Z","shell.execute_reply":"2024-09-04T22:23:24.508054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### meta_obj","metadata":{}},{"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-09-04T22:23:24.512141Z","iopub.execute_input":"2024-09-04T22:23:24.512589Z","iopub.status.idle":"2024-09-04T22:23:24.566047Z","shell.execute_reply.started":"2024-09-04T22:23:24.512546Z","shell.execute_reply":"2024-09-04T22:23:24.564109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-09-04T22:23:24.567882Z","iopub.execute_input":"2024-09-04T22:23:24.568368Z","iopub.status.idle":"2024-09-04T22:23:24.684425Z","shell.execute_reply.started":"2024-09-04T22:23:24.568321Z","shell.execute_reply":"2024-09-04T22:23:24.683063Z"},"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-09-04T22:23:24.686226Z","iopub.execute_input":"2024-09-04T22:23:24.686664Z","iopub.status.idle":"2024-09-04T22:23:33.700112Z","shell.execute_reply.started":"2024-09-04T22:23:24.686622Z","shell.execute_reply":"2024-09-04T22:23:33.698642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# meta_obj[list(meta_obj.keys())[1]]","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:33.702707Z","iopub.execute_input":"2024-09-04T22:23:33.703135Z","iopub.status.idle":"2024-09-04T22:23:33.708340Z","shell.execute_reply.started":"2024-09-04T22:23:33.703093Z","shell.execute_reply":"2024-09-04T22:23:33.706900Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### utilities","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","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:33.710184Z","iopub.execute_input":"2024-09-04T22:23:33.710712Z","iopub.status.idle":"2024-09-04T22:23:33.723461Z","shell.execute_reply.started":"2024-09-04T22:23:33.710654Z","shell.execute_reply":"2024-09-04T22:23:33.721936Z"},"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-09-04T22:23:33.725426Z","iopub.execute_input":"2024-09-04T22:23:33.725846Z","iopub.status.idle":"2024-09-04T22:23:33.751577Z","shell.execute_reply.started":"2024-09-04T22:23:33.725787Z","shell.execute_reply":"2024-09-04T22:23:33.750174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### wrapper","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-09-04T22:23:33.756704Z","iopub.execute_input":"2024-09-04T22:23:33.757161Z","iopub.status.idle":"2024-09-04T22:23:33.964071Z","shell.execute_reply.started":"2024-09-04T22:23:33.757119Z","shell.execute_reply":"2024-09-04T22:23:33.962460Z"},"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-09-04T22:23:33.966354Z","iopub.execute_input":"2024-09-04T22:23:33.966865Z","iopub.status.idle":"2024-09-04T22:23:34.005872Z","shell.execute_reply.started":"2024-09-04T22:23:33.966821Z","shell.execute_reply":"2024-09-04T22:23:34.003932Z"},"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.to(self.device))\n        return self.post_proc.forward(pred)","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:34.007727Z","iopub.execute_input":"2024-09-04T22:23:34.008177Z","iopub.status.idle":"2024-09-04T22:23:34.023888Z","shell.execute_reply.started":"2024-09-04T22:23:34.008136Z","shell.execute_reply":"2024-09-04T22:23:34.022395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def find_index(el_list, el):\n    for i in range(0,len(el_list)):\n        if el_list[i] == el:\n            return i\n    return -1","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:34.025518Z","iopub.execute_input":"2024-09-04T22:23:34.025901Z","iopub.status.idle":"2024-09-04T22:23:34.040176Z","shell.execute_reply.started":"2024-09-04T22:23:34.025846Z","shell.execute_reply":"2024-09-04T22:23:34.038759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Crop_n_Label():\n    def __init__(self, centroids, ptobj, study_id, base_path):\n        # centroids: np.zeros((7,2))\n        self.study_id = study_id;self.base_path = base_path\n        self.centroids = centroids;self.ptobj = ptobj\n        self.centroids_3d = np.zeros((len(centroids),3))\n        self.axial_lables = {\n            '1-2': [],\n            '2-3': [],\n            '3-4': [],\n            '4-5': [],\n            '5-6': []\n        }\n        self.type_label = {'s2':'Sagittal T2/STIR','s1':'Sagittal T1','a':'Axial T2'}\n        \n    def forward_sag_crop(self, img_type):\n        idx = find_index(self.ptobj['SeriesDescriptions'],self.type_label[img_type])\n        if idx == -1:\n            print(f\"skipping sag crop for {img_type}\")\n            return\n        i = self.ptobj['SeriesInstanceUIDs'][idx]\n        images = glob.glob(f\"{self.ptobj['folder_path']}/{self.ptobj['SeriesInstanceUIDs'][idx]}/*.dcm\")\n        self.create_dir_struct(f'{self.base_path}/{self.study_id}/{img_type}',['1-2','2-3','3-4','4-5','5-6'])\n\n        for j in sorted(images, key=lambda x: int(x.split('/')[-1].replace('.dcm', ''))):\n            dcm = pydicom.dcmread(j)\n            self.crop_sag_img(\n                dcm.pixel_array,save_path=f'{self.base_path}/{self.study_id}/{img_type}',\n                img_id=dcm.InstanceNumber\n            )\n\n        \n    def forward_axial_labeling(self, target_dcm, img_type='a'):\n        for c in range(0,len(self.centroids)):\n            self.centroids_3d[c] = self.get_3d_coords_sagg(self.centroids[c], target_dcm)\n#         print(\"centroids in 3d : \",self.centroids_3d)\n        axial_series_list = [self.ptobj['SeriesInstanceUIDs'][idx] for idx, x in enumerate(self.ptobj['SeriesDescriptions']) if x == self.type_label[img_type]]\n        self.create_dir_struct(f'{self.base_path}/{self.study_id}/{img_type}',['1-2','2-3','3-4','4-5','5-6'])\n#         label_dict = {'1-2':0,'2-3':0,'3-4':0,'4-5':0,'5-6':0,}\n        for axial_series in axial_series_list:\n            axial_images = os.listdir(f'{self.ptobj[\"folder_path\"]}/{axial_series}')\n#             print(\"count of input axial images : \",len(axial_images))\n            for img in axial_images:\n                axial_path = f'{self.ptobj[\"folder_path\"]}/{axial_series}/{img}'\n                self.classify_axial_slice(\n                    axial_path = axial_path,\n                    img_id = img,\n                    save_path = f'{self.base_path}/{self.study_id}/{img_type}'\n                )\n#                 label_dict[label] += 1\n                \n    def get_3d_coords_sagg(self, r, target_dcm):\n    \n        j, i = r\n        S = target_dcm.ImagePositionPatient\n        imgOrientation = [float(num) for num in target_dcm.ImageOrientationPatient]\n        X, Y = imgOrientation[:3], imgOrientation[3:]\n        del_i, del_j = target_dcm.PixelSpacing\n        \n        A = np.array([[X[0]*del_i, Y[0]*del_j, 0, S[0]], [X[1]*del_i, Y[1]*del_j, 0, S[1]], [X[2]*del_i, Y[2]*del_j, 0, S[2]], [0, 0, 0, 1]])\n        b = np.array([i, j, 0, 1])\n        r_3d = np.dot(A, b)[:3]\n\n        return r_3d\n\n    def classify_axial_slice(self, axial_path, img_id, save_path):\n        n_centroids = len(self.centroids_3d)-2\n\n        axial_img = pydicom.dcmread(axial_path)\n        imgOrientation = [float(num) for num in axial_img.ImageOrientationPatient]\n        X, Y = imgOrientation[:3], imgOrientation[3:]\n        n = np.cross(X, Y)\n        if n[2] < 0:\n            n *= -1\n        res = n_centroids\n        for i in range(n_centroids,-1,-1):\n            if n @ (self.centroids_3d[i] - axial_img.ImagePositionPatient) <= 0:\n                res = n_centroids - i\n                if res==0:\n                    res += 1\n                break\n        np.save(f'{save_path}/{res}-{res+1}/{img_id.replace(\".dcm\",\"\")}.npy',axial_img.pixel_array)\n        return f'{res}-{res+1}'\n    \n    def rectangle_from_points(self, p1, p2,scale_dim=2):\n        height = np.linalg.norm(p2 - p1)\n        width = scale_dim * height\n\n        midpoint = (p1 + p2) / 2\n\n        delta = p2 - p1\n        perpendicular_vector = np.array([-delta[1], delta[0]])\n\n        perpendicular_vector = (perpendicular_vector / np.linalg.norm(perpendicular_vector)) * (width / 2)\n\n        center = midpoint + perpendicular_vector\n        angle = np.degrees(np.arctan2(delta[1], delta[0]))\n    \n        return (tuple(center), (height, width), angle)\n\n    def crop_rectangle(self, img,rect):\n        # the order of the box points: bottom left, top left, top right,\n        # bottom right\n        box = cv2.boxPoints(rect)\n        box = np.int0(box)\n\n        # get width and height of the detected rectangle\n        width = int(rect[1][0])\n        height = int(rect[1][1])\n\n        src_pts = box.astype(\"float32\")\n        # coordinate of the points in box points after the rectangle has been\n        # straightened\n        dst_pts = np.array([[0, height-1],\n                            [0, 0],\n                            [width-1, 0],\n                            [width-1, height-1]], dtype=\"float32\")\n\n        # the perspective transformation matrix\n        M = cv2.getPerspectiveTransform(src_pts, dst_pts)\n\n        # directly warp the rotated rectangle to get the straightened rectangle\n        return np.rot90(cv2.warpPerspective(img, M, (width, height)))\n    \n    def crop_sag_img(self, img, img_id, save_path):\n        # do this for all images in sag t1 and sag t2\n        resized_img = cv2.resize(img,dsize=(512,512), interpolation=cv2.INTER_CUBIC)\n        resized_img = cv2.normalize(resized_img, None, alpha=0, beta=1, norm_type=cv2.NORM_MINMAX, dtype=cv2.CV_32F)\n\n        for i in range(0,len(self.centroids)-2):\n            crop_img = self.crop_rectangle(\n                resized_img, \n                self.rectangle_from_points(self.centroids[i][::-1],self.centroids[i+1][::-1])\n            )\n            np.save(f'{save_path}/{6-i-1}-{6-i}/{img_id}.npy', crop_img)\n            \n    def create_dir_struct(self, path, labels):\n        os.makedirs(path,exist_ok=True)\n        for label in labels:\n            os.makedirs(f'{path}/{label}',exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:34.042620Z","iopub.execute_input":"2024-09-04T22:23:34.043130Z","iopub.status.idle":"2024-09-04T22:23:34.084261Z","shell.execute_reply.started":"2024-09-04T22:23:34.043077Z","shell.execute_reply":"2024-09-04T22:23:34.083054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Restructure_DB():\n    def __init__(self, meta_obj, train, output_path):\n        self.unet = Unet_Pred()\n        self.meta_obj = meta_obj\n        self.train = train\n        self.base_path = output_path\n    \n    def forward(self,start,end):\n        for i in range(start,end):\n            print(i)\n            patient = self.train.iloc[i]\n            ptobj = self.meta_obj[str(patient['study_id'])]\n            # get center sagital image\n            center_t2 = self.get_center_t2(ptobj)\n            if center_t2 == None:\n                print('skipping patient as there are no s2 slices')\n                continue\n            # predict centroid for img\n            _, centroids = self.generate_centroids(center_t2)\n#             self.show_centroids(center_t2.pixel_array, centroids)\n            # iterate throught all t2 slices and crop\n            crop_n_label = Crop_n_Label(centroids, ptobj, patient['study_id'], self.base_path)\n            crop_n_label.forward_sag_crop('s2')\n            crop_n_label.forward_sag_crop('s1')\n            crop_n_label.forward_axial_labeling(center_t2)\n            if i==start or i==(end-1):\n                self.upload_n_refresh(i,start)\n            \n    def get_center_t2(self,ptobj):\n        idx = find_index(ptobj['SeriesDescriptions'],'Sagittal T2/STIR')\n        if idx == -1:\n            return None\n        i = ptobj['SeriesInstanceUIDs'][idx]; im_list_dcm = {}\n\n        im_list_dcm[i] = {'images': [], 'description': ptobj['SeriesDescriptions'][idx]}\n        images = glob.glob(f\"{ptobj['folder_path']}/{ptobj['SeriesInstanceUIDs'][idx]}/*.dcm\")\n        image_sorted = sorted(images, key=lambda x: int(x.split('/')[-1].replace('.dcm', '')))\n\n        return pydicom.dcmread(image_sorted[(len(image_sorted)-1)//2])\n    \n    def generate_centroids(self, target_dcm):\n        resized_img = cv2.resize(target_dcm.pixel_array,dsize=(512,512), interpolation=cv2.INTER_CUBIC)\n        resized_img = cv2.normalize(resized_img, None, alpha=0, beta=1, norm_type=cv2.NORM_MINMAX, dtype=cv2.CV_32F)\n        target_img = cv2.merge((resized_img,resized_img,resized_img))\n        target_img = torch.from_numpy(target_img).permute(2,0,1).float()\n        _, centroids = self.unet.forward(target_img.unsqueeze(dim=0))\n        return _, self.scaled_centroids(centroids, target_dcm.pixel_array.shape,resized_img.shape)\n    \n    def show_centroids(self,img, centroids):\n        plt.imshow(img,cmap='gray')\n        for c in centroids:\n            plt.plot(c[1],c[0],'ro')\n        plt.plot(centroids[:,1],centroids[:,0],'y')\n        plt.show()\n        \n    def scaled_centroids(self, centroids, old_shape, new_shape):\n        h_x, h_y = old_shape\n        scale_x, scale_y = h_x/new_shape[0], h_y/new_shape[1]\n        centroids[:,0] *= scale_x\n        centroids[:,1] *= scale_y\n        return centroids\n    \n    def upload_n_refresh(self,i,start):\n        if i==start:\n            !kaggle datasets create -p /kaggle/working/ds_up --dir-mode zip\n        else:\n            !kaggle datasets version -p /kaggle/working/ds_up --dir-mode zip -m f'chunk no. {i/700}'\n        !rm -r /kaggle/working/ds_up/train_db","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:34.085822Z","iopub.execute_input":"2024-09-04T22:23:34.086288Z","iopub.status.idle":"2024-09-04T22:23:34.144319Z","shell.execute_reply.started":"2024-09-04T22:23:34.086246Z","shell.execute_reply":"2024-09-04T22:23:34.142932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### restructuring db","metadata":{}},{"cell_type":"code","source":"os.makedirs('/kaggle/working/ds_up',exist_ok=True)\ndataset_metadata = {\n    \"title\": \"Processed RSNA Lumbar Stenosis Dataset S3\",\n    \"id\": \"coopermini/processed-rsna-lumbar-stenosis-dataset-s3\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}\nwith open('ds_up/dataset-metadata.json', 'w') as f:\n    json.dump(dataset_metadata, f)","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:35.339516Z","iopub.execute_input":"2024-09-04T22:23:35.339991Z","iopub.status.idle":"2024-09-04T22:23:35.348496Z","shell.execute_reply.started":"2024-09-04T22:23:35.339928Z","shell.execute_reply":"2024-09-04T22:23:35.347058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"restructure_db = Restructure_DB(meta_obj, train, '/kaggle/working/ds_up/train_db')","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:35.350308Z","iopub.execute_input":"2024-09-04T22:23:35.350830Z","iopub.status.idle":"2024-09-04T22:23:37.302147Z","shell.execute_reply.started":"2024-09-04T22:23:35.350772Z","shell.execute_reply":"2024-09-04T22:23:37.300567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# restructure_db.forward(0,701)\n# restructure_db.forward(700,1401)\nrestructure_db.forward(1400,len(train))","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:37.304399Z","iopub.execute_input":"2024-09-04T22:23:37.305117Z","iopub.status.idle":"2024-09-04T22:23:53.122843Z","shell.execute_reply.started":"2024-09-04T22:23:37.305050Z","shell.execute_reply":"2024-09-04T22:23:53.120896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## preidiction","metadata":{}},{"cell_type":"code","source":"# import os\n# import matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:53.125373Z","iopub.execute_input":"2024-09-04T22:23:53.126596Z","iopub.status.idle":"2024-09-04T22:23:53.132805Z","shell.execute_reply.started":"2024-09-04T22:23:53.126541Z","shell.execute_reply":"2024-09-04T22:23:53.131306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pat_obj = {\n#     's2' : {\n#             '1-2':[],\n#             '2-3':[],\n#             '3-4':[],\n#             '4-5':[],\n#             '5-6':[],\n#         },\n#     's1' : {\n#             '1-2':[],\n#             '2-3':[],\n#             '3-4':[],\n#             '4-5':[],\n#             '5-6':[],\n#         },\n#     'a' : {\n#             '1-2':[],\n#             '2-3':[],\n#             '3-4':[],\n#             '4-5':[],\n#             '5-6':[],\n#         }\n# }","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:53.134671Z","iopub.execute_input":"2024-09-04T22:23:53.135138Z","iopub.status.idle":"2024-09-04T22:23:53.145619Z","shell.execute_reply.started":"2024-09-04T22:23:53.135091Z","shell.execute_reply":"2024-09-04T22:23:53.144337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# meta_obj = {}\n# base_path = '/kaggle/input/predictive-data/train_zip'\n# index = {'1-2':0,'2-3':1,'3-4':2,'4-5':3,'5-6':4}","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:53.147616Z","iopub.execute_input":"2024-09-04T22:23:53.148200Z","iopub.status.idle":"2024-09-04T22:23:53.162662Z","shell.execute_reply.started":"2024-09-04T22:23:53.148136Z","shell.execute_reply":"2024-09-04T22:23:53.161280Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dir_list = os.listdir(base_path)","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:53.164677Z","iopub.execute_input":"2024-09-04T22:23:53.165261Z","iopub.status.idle":"2024-09-04T22:23:53.642762Z","shell.execute_reply.started":"2024-09-04T22:23:53.165201Z","shell.execute_reply":"2024-09-04T22:23:53.640251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# type_dict = [0,0,0,0,0]\n# len_dict = {'1-2':[],'2-3':[],'3-4':[],'4-5':[],'5-6':[],}\n# typ = 'a'\n# cnt_target = 0\n# for pat in dir_list:\n# #         meta_obj[pat] = {\n# #             's2' : {\n# #                     '1-2':[],\n# #                     '2-3':[],\n# #                     '3-4':[],\n# #                     '4-5':[],\n# #                     '5-6':[],\n# #                 },\n# #             's1' : {\n# #                     '1-2':[],\n# #                     '2-3':[],\n# #                     '3-4':[],\n# #                     '4-5':[],\n# #                     '5-6':[],\n# #                 },\n# #             'a' : {\n# #                     '1-2':[],\n# #                     '2-3':[],\n# #                     '3-4':[],\n# #                     '4-5':[],\n# #                     '5-6':[],\n# #                 }\n# #         }\n#     path = f'{base_path}/{pat}/{typ}'\n#     if os.path.exists(path):\n#         for dsc in os.listdir(path):\n# #                 meta_obj[pat][typ][dsc] = os.listdir(f'{path}/{dsc}')\n#             type_dict[index[dsc]] += 1\n#             len_dict[dsc].append(len(os.listdir(f'{path}/{dsc}')))\n#             cnt_target += len_dict[dsc][-1]\n            ","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:53.644369Z","iopub.status.idle":"2024-09-04T22:23:53.645019Z","shell.execute_reply.started":"2024-09-04T22:23:53.644673Z","shell.execute_reply":"2024-09-04T22:23:53.644704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(type_dict,cnt_target)","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:53.647674Z","iopub.status.idle":"2024-09-04T22:23:53.648298Z","shell.execute_reply.started":"2024-09-04T22:23:53.647967Z","shell.execute_reply":"2024-09-04T22:23:53.648018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plt.bar(index.keys(),type_dict)\n# plt.title(typ)\n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:53.650736Z","iopub.status.idle":"2024-09-04T22:23:53.651385Z","shell.execute_reply.started":"2024-09-04T22:23:53.651048Z","shell.execute_reply":"2024-09-04T22:23:53.651079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plt.boxplot(len_dict.values(),labels=list(len_dict.keys()))\n# plt.title(typ)\n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:53.653692Z","iopub.status.idle":"2024-09-04T22:23:53.654236Z","shell.execute_reply.started":"2024-09-04T22:23:53.653992Z","shell.execute_reply":"2024-09-04T22:23:53.654017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# inp_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images'","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:53.656777Z","iopub.status.idle":"2024-09-04T22:23:53.657369Z","shell.execute_reply.started":"2024-09-04T22:23:53.657114Z","shell.execute_reply":"2024-09-04T22:23:53.657145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# db_index = ['Sagittal T1', 'Axial T2', 'Sagittal T2/STIR']\n# cnt = 0\n# ppd = []\n# for pat in dir_list:\n#     ind = meta_obj[pat]['SeriesDescriptions'].index(db_index[1])\n#     ppd.append(len(os.listdir(f'{inp_path}/{pat}/{meta_obj[pat][\"SeriesInstanceUIDs\"][ind]}')))\n#     cnt +=  ppd[-1]\n# print(cnt)","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:53.659696Z","iopub.status.idle":"2024-09-04T22:23:53.660246Z","shell.execute_reply.started":"2024-09-04T22:23:53.659961Z","shell.execute_reply":"2024-09-04T22:23:53.660008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# bp = plt.boxplot(ppd)\n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:53.662524Z","iopub.status.idle":"2024-09-04T22:23:53.663249Z","shell.execute_reply.started":"2024-09-04T22:23:53.662899Z","shell.execute_reply":"2024-09-04T22:23:53.662934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# bp['medians'][0].get_ydata()","metadata":{"execution":{"iopub.status.busy":"2024-09-04T22:23:53.665383Z","iopub.status.idle":"2024-09-04T22:23:53.666017Z","shell.execute_reply.started":"2024-09-04T22:23:53.665681Z","shell.execute_reply":"2024-09-04T22:23:53.665715Z"},"trusted":true},"execution_count":null,"outputs":[]}]}