{"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":9538089,"sourceType":"datasetVersion","datasetId":5726807},{"sourceId":9539917,"sourceType":"datasetVersion","datasetId":5726703}],"dockerImageVersionId":30787,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install segmentation_models_pytorch","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-11-01T14:10:51.105336Z","iopub.execute_input":"2024-11-01T14:10:51.106114Z","iopub.status.idle":"2024-11-01T14:11:03.051020Z","shell.execute_reply.started":"2024-11-01T14:10:51.106066Z","shell.execute_reply":"2024-11-01T14:11:03.049985Z"}},"outputs":[],"execution_count":null},{"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 pickle\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":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-01T14:11:03.053548Z","iopub.execute_input":"2024-11-01T14:11:03.054360Z","iopub.status.idle":"2024-11-01T14:11:09.236606Z","shell.execute_reply.started":"2024-11-01T14:11:03.054308Z","shell.execute_reply":"2024-11-01T14:11:09.235571Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"Axial_T2_axial_segmentation_paths = [\n        '/kaggle/input/axial-t2/Axial_T2_axial_side_segmentation_1',\n        '/kaggle/input/axial-t2/Axial_T2_axial_side_segmentation_2',\n        '/kaggle/input/axial-t2/Axial_T2_axial_side_segmentation_3',\n        '/kaggle/input/axial-t2/Axial_T2_axial_side_segmentation_4',\n        '/kaggle/input/axial-t2/Axial_T2_axial_side_segmentation_5'\n]\nTRAIN_PATH = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-01T14:11:09.237991Z","iopub.execute_input":"2024-11-01T14:11:09.238314Z","iopub.status.idle":"2024-11-01T14:11:09.242904Z","shell.execute_reply.started":"2024-11-01T14:11:09.238279Z","shell.execute_reply":"2024-11-01T14:11:09.241974Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myUNet(nn.Module):\n    def __init__(\n        self,\n        classes\n        ):\n        super(myUNet, self).__init__()\n\n        self.classes = classes\n        self.UNet = smp.Unet(\n            encoder_name=ENCODER_NAME,\n            classes=classes,\n            in_channels=1\n        ).to(device)\n\n    def forward(self,X):\n        H,W = X.shape[-2:]\n        x = self.UNet(X.view(-1,1,H,W)).view(-1,H*W)\n#       MinMaxScaling along the class plane to generate a heatmap\n        min_values = x.min(-1)[0].view(-1,1)\n        max_values = x.max(-1)[0].view(-1,1)\n        d = (max_values - min_values)\n        d[d == 0] = 1\n        x = (x - min_values)/d\n        \n        return x.view(-1,self.classes,H,W)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-01T14:11:09.244221Z","iopub.execute_input":"2024-11-01T14:11:09.244765Z","iopub.status.idle":"2024-11-01T14:11:09.260127Z","shell.execute_reply.started":"2024-11-01T14:11:09.244722Z","shell.execute_reply":"2024-11-01T14:11:09.259262Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PATCH_H = 512\nPATCH_W = 512\nTH = .5\nBS = 64","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-01T14:11:09.262587Z","iopub.execute_input":"2024-11-01T14:11:09.262873Z","iopub.status.idle":"2024-11-01T14:11:09.275749Z","shell.execute_reply.started":"2024-11-01T14:11:09.262842Z","shell.execute_reply":"2024-11-01T14:11:09.274857Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/sagittal-t1/train_split.csv')\ntrain.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-01T14:11:09.276863Z","iopub.execute_input":"2024-11-01T14:11:09.277240Z","iopub.status.idle":"2024-11-01T14:11:09.326809Z","shell.execute_reply.started":"2024-11-01T14:11:09.277206Z","shell.execute_reply":"2024-11-01T14:11:09.325777Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"series = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv')\nseries = series[series.series_description == 'Axial T2']\nseries.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-01T14:11:09.327861Z","iopub.execute_input":"2024-11-01T14:11:09.328212Z","iopub.status.idle":"2024-11-01T14:11:09.346888Z","shell.execute_reply.started":"2024-11-01T14:11:09.328168Z","shell.execute_reply":"2024-11-01T14:11:09.345930Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = train.merge(series,left_on='study_id',right_on='study_id')\ntrain.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-01T14:11:09.348043Z","iopub.execute_input":"2024-11-01T14:11:09.348355Z","iopub.status.idle":"2024-11-01T14:11:09.378932Z","shell.execute_reply.started":"2024-11-01T14:11:09.348321Z","shell.execute_reply":"2024-11-01T14:11:09.377925Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"models = [\n    torch.load(path,map_location=torch.device(device)) for path in Axial_T2_axial_segmentation_paths\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-01T14:11:09.380138Z","iopub.execute_input":"2024-11-01T14:11:09.380494Z","iopub.status.idle":"2024-11-01T14:11:09.893740Z","shell.execute_reply.started":"2024-11-01T14:11:09.380448Z","shell.execute_reply":"2024-11-01T14:11:09.892967Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch_resize = torchvision.transforms.Resize((PATCH_H,PATCH_W),antialias=True)\n\ncenters = {k:{} for k in [1,2,3,4,5]}\nfor f in centers:\n    for k in tqdm(range(len(train))):\n        row = train.iloc[k]\n        centers[f][row.study_id] = {}\n    for k in tqdm(range(len(train))):\n        row = train.iloc[k]\n        centers[f][row.study_id][row.series_id] = {}\nwith warnings.catch_warnings():\n    warnings.simplefilter(\"ignore\", category=RuntimeWarning)\n    for k in tqdm(range(len(train))):\n        row = train.iloc[k]\n        sample = TRAIN_PATH + str(row['study_id']) + '/' + str(row['series_id'])\n\n        images = [x for x in glob.glob(sample+'/*.dcm')]\n        images.sort(key=lambda v:int(v.split('/')[-1].replace('.dcm','')))\n\n        dicom = [pydicom.dcmread(dicom_file) for dicom_file in images]\n        images = [torch.as_tensor(dcm.pixel_array.astype(float)) for dcm in dicom]\n\n        HW = np.array([img.shape for img in images])\n\n        H,W = HW.max(0)\n\n        images = torch.concat([torch.nn.functional.pad(\n            images[k].unsqueeze(0),(\n                (W - HW[k][-1])//2,\n                (W - HW[k][-1]) - (W - HW[k][-1])//2,\n                (H - HW[k][-2])//2,\n                (H - HW[k][-2]) - (H - HW[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            images = images[:,h:h+d]\n            H = W\n        elif H < W:\n            d = H\n            w = (W - d)//2\n            images = images[:,:,w:w+d]\n            W = H\n\n        V = torch_resize(images/images.max()).float().to(device)\n\n        D = V.shape[0]\n    \n        MASK = torch.zeros(D,2,PATCH_H,PATCH_W).float().to(device)\n        for f in [1,2,3,4,5]:\n            model = models[f-1]\n            MASK[:] = 0\n            with torch.no_grad():\n                for k in range(D//BS + 1):\n                    START = k*BS\n                    mask = 0\n                    v = V[START:START+BS]\n                    for rot in [0,1,2,3]:\n                        rot_v = torch.rot90(v, rot, dims=[-2, -1])\n                        mask += torch.rot90(model(rot_v), k=-rot, dims=[-2, -1])\n                        mask += torch.rot90(model(rot_v.flip(-1)).flip(-1).flip(1), k=-rot, dims=[-2, -1])\n                \n                    MASK[START:START+BS] = mask\n            MASK = MASK/(2*4)\n            \n            mask = MASK.cpu()\n        \n            y,x = [],[]\n            for m in mask:\n                s,yy,xx = np.where(m > TH)\n                for i in range(2):\n                    y.append(yy[s==i].mean())\n                    x.append(xx[s==i].mean())\n            centers[f][row.study_id][row.series_id] = torch.tensor([x,y]).T.reshape(-1,2,2).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-01T14:11:09.894922Z","iopub.execute_input":"2024-11-01T14:11:09.895271Z","iopub.status.idle":"2024-11-01T14:12:08.935275Z","shell.execute_reply.started":"2024-11-01T14:11:09.895234Z","shell.execute_reply":"2024-11-01T14:12:08.933924Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open('axial_centers.pkl', 'wb') as handle:\n    pickle.dump(centers, handle)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-01T14:12:08.936176Z","iopub.status.idle":"2024-11-01T14:12:08.936584Z","shell.execute_reply.started":"2024-11-01T14:12:08.936382Z","shell.execute_reply":"2024-11-01T14:12:08.936402Z"}},"outputs":[],"execution_count":null}]}