{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","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":7229399,"sourceType":"datasetVersion","datasetId":4185566},{"sourceId":7229919,"sourceType":"datasetVersion","datasetId":4185960},{"sourceId":7229957,"sourceType":"datasetVersion","datasetId":4185994},{"sourceId":7324315,"sourceType":"datasetVersion","datasetId":4250949},{"sourceId":9054265,"sourceType":"datasetVersion","datasetId":5459494},{"sourceId":9284588,"sourceType":"datasetVersion","datasetId":5597030},{"sourceId":9538089,"sourceType":"datasetVersion","datasetId":5726807},{"sourceId":9539833,"sourceType":"datasetVersion","datasetId":5729440},{"sourceId":9539917,"sourceType":"datasetVersion","datasetId":5726703},{"sourceId":9562501,"sourceType":"datasetVersion","datasetId":5827535},{"sourceId":9332336,"sourceType":"datasetVersion","datasetId":5639221}],"dockerImageVersionId":30747,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nfrom tqdm import tqdm\nimport glob\nimport gc\nimport pydicom\nimport math\nimport warnings\nimport pickle\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset\n\n!pip install '/kaggle/input/hengck23-ver-1-demo-workflow-2-stage-approach/natsort-8.4.0-py3-none-any.whl'\n\nimport sys, os\nsys.path.append('/kaggle/input/hengck23-ver-1-demo-workflow-2-stage-approach')\n\nfrom _dir_setting_ import *\n\nimport matplotlib\nimport matplotlib.pyplot as plt\n\nfrom helper import *\nfrom data import *\n\nimport sys\nsys.path.append(\"/kaggle/input/segmentation-models-pytorch/segmentation_models.pytorch-master/\")\nsys.path.append(\"/kaggle/input/pretrainedmodels-0-7-4/pretrainedmodels-0.7.4\")\nsys.path.append(\"/kaggle/input/efficientnet-pytorch-0-7-1/efficientnet_pytorch-0.7.1\")\nimport efficientnet_pytorch\nimport pretrainedmodels\nimport segmentation_models_pytorch as smp\nimport torchvision\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-10-07T20:14:07.317603Z","iopub.execute_input":"2024-10-07T20:14:07.317934Z","iopub.status.idle":"2024-10-07T20:14:45.052037Z","shell.execute_reply.started":"2024-10-07T20:14:07.317903Z","shell.execute_reply":"2024-10-07T20:14:45.050955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = False\nN = 100 # 0 for all","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:14:45.057296Z","iopub.execute_input":"2024-10-07T20:14:45.057610Z","iopub.status.idle":"2024-10-07T20:14:45.061996Z","shell.execute_reply.started":"2024-10-07T20:14:45.057584Z","shell.execute_reply":"2024-10-07T20:14:45.061033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Swiss Army Knife model\nSagittal_T1_sagittal_segmentation_paths = [\n        '/kaggle/input/sagittal-t1/Sagittal_T1_sagittal_level_segmentation_1',\n        '/kaggle/input/sagittal-t1/Sagittal_T1_sagittal_level_segmentation_2',\n        '/kaggle/input/sagittal-t1/Sagittal_T1_sagittal_level_segmentation_3',\n        '/kaggle/input/sagittal-t1/Sagittal_T1_sagittal_level_segmentation_4',\n        '/kaggle/input/sagittal-t1/Sagittal_T1_sagittal_level_segmentation_5'\n]\nAxial_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]\nSagittal_T2_sagittal_segmentation_paths = [\n        '/kaggle/input/sagittal-t2/Sagittal_T2_sagittal_level_segmentation_1',\n        '/kaggle/input/sagittal-t2/Sagittal_T2_sagittal_level_segmentation_2',\n        '/kaggle/input/sagittal-t2/Sagittal_T2_sagittal_level_segmentation_3',\n        '/kaggle/input/sagittal-t2/Sagittal_T2_sagittal_level_segmentation_4',\n        '/kaggle/input/sagittal-t2/Sagittal_T2_sagittal_level_segmentation_5'\n]\nSagittal_T1_foraminal_paths = [\n    '/kaggle/input/sagittal-t1/Sagittal_T1_pretrained_foraminal_ViT_1',\n    '/kaggle/input/sagittal-t1/Sagittal_T1_pretrained_foraminal_ViT_2',\n    '/kaggle/input/sagittal-t1/Sagittal_T1_pretrained_foraminal_ViT_3',\n    '/kaggle/input/sagittal-t1/Sagittal_T1_pretrained_foraminal_ViT_4',\n    '/kaggle/input/sagittal-t1/Sagittal_T1_pretrained_foraminal_ViT_5'\n]\nSagittal_T2_spinal_paths = [\n    '/kaggle/input/sagittal-t2/Sagittal_T2_Spinal_ViT_1',\n    '/kaggle/input/sagittal-t2/Sagittal_T2_Spinal_ViT_2',\n    '/kaggle/input/sagittal-t2/Sagittal_T2_Spinal_ViT_3',\n    '/kaggle/input/sagittal-t2/Sagittal_T2_Spinal_ViT_4',\n    '/kaggle/input/sagittal-t2/Sagittal_T2_Spinal_ViT_5'\n]\nAxial_T2_subarticular_paths = [\n    '/kaggle/input/subarticular-dicom-v2-vit/subarticular_DICOM_V2_ViT_1',\n    '/kaggle/input/subarticular-dicom-v2-vit/subarticular_DICOM_V2_ViT_2',\n    '/kaggle/input/subarticular-dicom-v2-vit/subarticular_DICOM_V2_ViT_3',\n    '/kaggle/input/subarticular-dicom-v2-vit/subarticular_DICOM_V2_ViT_4',\n    '/kaggle/input/subarticular-dicom-v2-vit/subarticular_DICOM_V2_ViT_5'\n]\nAxial_T2_spinal_paths = [\n    '/kaggle/input/subarticular-dicom-v2-vit/spinal_DICOM_ViT_1.pth',\n    '/kaggle/input/subarticular-dicom-v2-vit/spinal_DICOM_ViT_2.pth',\n    '/kaggle/input/subarticular-dicom-v2-vit/spinal_DICOM_ViT_3.pth',\n    '/kaggle/input/subarticular-dicom-v2-vit/spinal_DICOM_ViT_4.pth',\n    '/kaggle/input/subarticular-dicom-v2-vit/spinal_DICOM_ViT_5.pth'\n]\n# Definitions\nspinal = [\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]\nlforaminal = [\n    'left_neural_foraminal_narrowing_l1_l2',\n    'left_neural_foraminal_narrowing_l2_l3',\n    'left_neural_foraminal_narrowing_l3_l4',\n    'left_neural_foraminal_narrowing_l4_l5',\n    'left_neural_foraminal_narrowing_l5_s1'\n]\nrforaminal = [\n    'right_neural_foraminal_narrowing_l1_l2',\n    'right_neural_foraminal_narrowing_l2_l3',\n    'right_neural_foraminal_narrowing_l3_l4',\n    'right_neural_foraminal_narrowing_l4_l5',\n    'right_neural_foraminal_narrowing_l5_s1'\n]\nlsubarticular = [\n    'left_subarticular_stenosis_l1_l2',\n    'left_subarticular_stenosis_l2_l3',\n    'left_subarticular_stenosis_l3_l4',\n    'left_subarticular_stenosis_l4_l5',\n    'left_subarticular_stenosis_l5_s1'\n]\nrsubarticular = [\n    'right_subarticular_stenosis_l1_l2',\n    'right_subarticular_stenosis_l2_l3',\n    'right_subarticular_stenosis_l3_l4',\n    'right_subarticular_stenosis_l4_l5',\n    'right_subarticular_stenosis_l5_s1'\n]\nforaminal = lforaminal + rforaminal\nsubarticular = lsubarticular + rsubarticular\ndiagnosis = spinal + foraminal + subarticular\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":{"iopub.status.busy":"2024-10-07T20:14:45.063270Z","iopub.execute_input":"2024-10-07T20:14:45.063570Z","iopub.status.idle":"2024-10-07T20:14:45.075696Z","shell.execute_reply.started":"2024-10-07T20:14:45.063543Z","shell.execute_reply":"2024-10-07T20:14:45.074902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    TEST_PATH = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/'\n    test_description = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv')\n    if N > 0: test_description = test_description[:N]\n    test_description.head()\n    TEST = False\nelse:\n    TEST_PATH = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/'\n    test_description = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv')\n    test_description.head()\n    TEST = True","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:14:45.077868Z","iopub.execute_input":"2024-10-07T20:14:45.078152Z","iopub.status.idle":"2024-10-07T20:14:45.099328Z","shell.execute_reply.started":"2024-10-07T20:14:45.078129Z","shell.execute_reply":"2024-10-07T20:14:45.098535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocessing","metadata":{}},{"cell_type":"code","source":"PATCH_H = 512\nPATCH_W = 512\nTH = .5\nBS = 64","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:14:45.100413Z","iopub.execute_input":"2024-10-07T20:14:45.100787Z","iopub.status.idle":"2024-10-07T20:14:45.105925Z","shell.execute_reply.started":"2024-10-07T20:14:45.100757Z","shell.execute_reply":"2024-10-07T20:14:45.105056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-07T20:14:45.107115Z","iopub.execute_input":"2024-10-07T20:14:45.108009Z","iopub.status.idle":"2024-10-07T20:14:45.118324Z","shell.execute_reply.started":"2024-10-07T20:14:45.107976Z","shell.execute_reply":"2024-10-07T20:14:45.117614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cases = list(test_description.groupby('study_id'))\nmodels = {\n    'Sagittal T1':[\n        torch.load(path,map_location=torch.device(device)) for path in Sagittal_T1_sagittal_segmentation_paths\n    ],\n    'Sagittal T2/STIR':[\n        torch.load(path,map_location=torch.device(device)) for path in Sagittal_T2_sagittal_segmentation_paths\n    ],\n    'Axial T2':[\n        torch.load(path,map_location=torch.device(device)) for path in Axial_T2_axial_segmentation_paths\n    ]\n}","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:14:45.119338Z","iopub.execute_input":"2024-10-07T20:14:45.119637Z","iopub.status.idle":"2024-10-07T20:14:46.350615Z","shell.execute_reply.started":"2024-10-07T20:14:45.119614Z","shell.execute_reply":"2024-10-07T20:14:46.349569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch_resize = torchvision.transforms.Resize((PATCH_H,PATCH_W),antialias=True)\n\ndef read_volume(df):\n    sample = TEST_PATH + str(df['study_id']) + '/' + str(df['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#   TODO: Maintain proportions also in segmentation\n    V = torch.concat([\n        torch_resize(images[i].unsqueeze(0)) for i in range(len(images))\n    ]).float().to(device)\n    V = V/V.max()\n    D = V.shape[0]\n    \n    if df.series_description == 'Axial T2':\n        MASK = torch.zeros(D,2,PATCH_H,PATCH_W).float().to(device)\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                    for model in models['Axial T2']:\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*len( models[df['series_description']]))\n    else:\n        MASK = torch.zeros(D,5,PATCH_H,PATCH_W).float().to(device)\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 model in models[df['series_description']]:\n                    mask += model(v)\n                \n                MASK[START:START+BS] = mask\n        MASK =MASK/len(models[df['series_description']])\n    #   Trust middle slices only\n        head = tail = D//5\n        MASK = MASK[head:D-tail]\n            \n    mask = MASK.cpu()\n        \n    if df.series_description == 'Axial T2':\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 = np.array([x,y]).T.reshape(-1,2,2)\n        centers[:,:,0] = centers[:,:,0]*(HW[:,1].reshape(-1,1))/PATCH_W\n        centers[:,:,1] = centers[:,:,1]*(HW[:,0].reshape(-1,1))/PATCH_H\n        all_centers = centers.copy()\n#       As firts approach we'll impute volume mean to missing values\n#       Once levels have been assigned we'll impute level mean instead\n        center_mean = np.tile(np.nanmean(centers,0).reshape(1,2,2),(len(centers),1,1))\n        center_mask = np.isnan(centers)\n        all_centers[center_mask] = center_mean[center_mask]\n    else:\n        mask = mask.sum(0).view(5,-1)\n        mask_max = mask.max(-1)[0].view(-1,1)\n        mask_min = mask.min(-1)[0].view(-1,1)\n        d = mask_max - mask_min\n        d[d == 0] = 1\n        mask = ((mask - mask_min)/d).view(5,PATCH_H,PATCH_W)\n\n        l,yy,xx = np.where(mask > TH)\n        y,x = [],[]\n        for i in range(5):\n            y.append(yy[l==i].mean())\n            x.append(xx[l==i].mean())\n        centers = np.array([x,y]).T\n        \n        centers[:,0] = centers[:,0]*HW[0,1]/PATCH_W\n        centers[:,1] = centers[:,1]*HW[0,0]/PATCH_H\n        all_centers = np.stack([centers]*len(dicom))\n        \n        if study_id == cases[0][0]:\n            plt.imshow(V.sum(0).cpu()/len(V) + .5*(mask > TH).sum(0))\n            plt.show()\n#   https://www.kaggle.com/code/hengck23/2d-to-3d-projection-for-dicom\n    xx,yy,zz = [],[],[]\n    for i in range(len(dicom)):\n        sx,sy,sz = [float(v) for v in dicom[i].ImagePositionPatient]\n        o0, o1, o2, o3, o4, o5 = [float(v) for v in dicom[i].ImageOrientationPatient]\n        delx,dely = dicom[i].PixelSpacing\n\n        for x,y in all_centers[i]:\n            xx.append(o0*delx*x + o3*dely*y + sx)\n            yy.append(o1*delx*x + o4*dely*y + sy)\n            zz.append(o2*delx*x + o5*dely*y + sz)\n\n    return np.stack([xx,yy,zz]).swapaxes(0,1).reshape(len(dicom),-1,3),centers","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:14:46.352541Z","iopub.execute_input":"2024-10-07T20:14:46.352926Z","iopub.status.idle":"2024-10-07T20:14:46.382266Z","shell.execute_reply.started":"2024-10-07T20:14:46.352893Z","shell.execute_reply":"2024-10-07T20:14:46.381295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"xyz = {}\nAxial_T2_xy = {}\nAxial_T2_flipped = {}\nAxial_T2_assignments = {}\nAxial_T2_instance_number = []\nT1_study_ids = []\nT1_series_ids = []\nT1_flipped = []\nT1_instance_number = []\nT1_levels = []\nT2_study_ids = []\nT2_series_ids = []\nT2_instance_number = []\nT2_levels = []\nfor study_id,df in tqdm(cases):\n    Sagittal_T1_df = df[df.series_description == 'Sagittal T1']\n    Sagittal_T2_df = df[df.series_description == 'Sagittal T2/STIR']\n    Axial_T2_df = df[df.series_description == 'Axial T2']\n    \n    xyz['Sagittal_T1'] = {}\n    xyz['Sagittal_T2'] = {}\n    xyz['Axial_T2'] = {}\n    \n    if len(Sagittal_T1_df) > 0:\n        xyz['Sagittal_T1'] = {}\n        for k in range(len(Sagittal_T1_df)):\n            row = Sagittal_T1_df.iloc[k]\n            xyz['Sagittal_T1'][row.series_id] = read_volume(row)\n            \n    if len(Sagittal_T2_df) > 0:\n        xyz['Sagittal_T2'] = {}\n        for k in range(len(Sagittal_T2_df)):\n            row = Sagittal_T2_df.iloc[k]\n            xyz['Sagittal_T2'][row.series_id] = read_volume(row)\n\n    if len(Axial_T2_df) > 0:\n        xyz['Axial_T2'] = {}\n        Axial_T2_xy[study_id] = {}\n        Axial_T2_flipped[study_id] = {}\n        Axial_T2_assignments[study_id] = {}\n        with warnings.catch_warnings():\n            warnings.simplefilter(\"ignore\", category=RuntimeWarning)\n            for k in range(len(Axial_T2_df)):\n                row = Axial_T2_df.iloc[k]\n                xyz['Axial_T2'][row.series_id],Axial_T2_xy[study_id][row.series_id] = read_volume(row)\n                points = xyz['Axial_T2'][row.series_id][:,:,2].mean(1)\n                Axial_T2_flipped[study_id][row.series_id] = points[-1] > points[0]\n                Axial_T2_assignments[study_id][row.series_id] = []\n\n    if study_id == cases[0][0]:\n        for Sagittal_T1_series in xyz['Sagittal_T1']:\n            for l in [0,1,2,3,4]:\n                plt.plot(\n                    xyz['Sagittal_T1'][Sagittal_T1_series][0][:,l,0],\n                    xyz['Sagittal_T1'][Sagittal_T1_series][0][:,l,2],\n                    '.'\n                )\n        for Axial_T2_series in xyz['Axial_T2']:\n            for s in [0,1]:\n                 plt.plot(\n                    xyz['Axial_T2'][Axial_T2_series][:,s,0],\n                    xyz['Axial_T2'][Axial_T2_series][:,s,2],\n                    '.'\n                 )\n        plt.show()\n        \n        for Sagittal_T2_series in xyz['Sagittal_T2']:\n            for l in [0,1,2,3,4]:\n                plt.plot(\n                    xyz['Sagittal_T2'][Sagittal_T2_series][0][:,l,0],\n                    xyz['Sagittal_T2'][Sagittal_T2_series][0][:,l,2],\n                    '.'\n                )\n        for Axial_T2_series in xyz['Axial_T2']:\n            for s in [0,1]:\n                 plt.plot(\n                    xyz['Axial_T2'][Axial_T2_series][:,s,0],\n                    xyz['Axial_T2'][Axial_T2_series][:,s,2],\n                    '.'\n                 )\n        plt.show()\n    \n    for Sagittal_T1_series in xyz['Sagittal_T1']:\n        T1_study_ids.append(study_id)\n        T1_series_ids.append(Sagittal_T1_series)\n        T1_levels.append(xyz['Sagittal_T1'][Sagittal_T1_series][1].reshape(-1).tolist())\n        sagittal_points = xyz['Sagittal_T1'][Sagittal_T1_series][0][:,:,0].mean(1)\n        T1_flipped.append(sagittal_points[-1] < sagittal_points[0])\n        \n        sagittal_assignments = []\n        xxyyzz = torch.as_tensor(xyz['Sagittal_T1'][Sagittal_T1_series][0]).view(1,-1,5,3).to(device)\n        for Axial_T2_series in xyz['Axial_T2']:\n            points = torch.as_tensor(xyz['Axial_T2'][Axial_T2_series]).to(device)\n            points = points.mean(1)\n            d = points.view(-1,1,1,3) - xxyyzz\n            d = d*d\n            d = d.sum(-1)\n            level_assignments = torch.argmin(d.min(1)[0],-1)\n            if study_id == cases[0][0]: print(level_assignments)           \n            for l in [0,1,2,3,4]:\n                level_mask = level_assignments == l\n                if level_mask.sum() > 0:\n                    dd = d[level_mask,:,l]\n                    axial_idx,sagittal_idx = torch.where(dd == dd.min())\n                    sagittal_assignments.append([l,sagittal_idx[0].item()])\n                \n        v = torch.zeros(5)\n        v[:] = torch.nan\n        slices_df = pd.DataFrame(sagittal_assignments).groupby([0]).mean()\n        for i in range(len(slices_df)):\n            row = slices_df.iloc[i]\n            v[\n                {\n                    0:0,\n                    1:1,\n                    2:2,\n                    3:3,\n                    4:4\n                }[row.name]\n            ] = row.values[0]\n    \n        v_mean = v.nanmean()\n        mask = v.isnan()\n        v[mask] = v_mean\n        T1_instance_number.append(v.tolist())\n        \n    for Sagittal_T2_series in xyz['Sagittal_T2']:\n        T2_study_ids.append(study_id)\n        T2_series_ids.append(Sagittal_T2_series)\n        T2_levels.append(xyz['Sagittal_T2'][Sagittal_T2_series][1].reshape(-1).tolist())\n        \n        sagittal_assignments = []\n        xxyyzz = torch.as_tensor(xyz['Sagittal_T2'][Sagittal_T2_series][0]).view(1,-1,5,3).to(device)\n        for Axial_T2_series in xyz['Axial_T2']:\n            points = torch.as_tensor(xyz['Axial_T2'][Axial_T2_series]).to(device)\n            points = points.mean(1)\n            axial_indices = torch.arange(len(points)).to(device)\n            d = points.view(-1,1,1,3) - xxyyzz\n            d = d*d\n            d = d.sum(-1)\n            level_assignments = torch.argmin(d.min(1)[0],-1)\n            Axial_T2_assignments[study_id][Axial_T2_series].append(level_assignments.tolist())\n            if study_id == cases[0][0]: print(level_assignments)\n            v = torch.zeros(5)\n            v[:] = torch.nan\n            for l in [0,1,2,3,4]:\n                level_mask = level_assignments == l\n                if level_mask.sum() > 0:\n                    dd = d[level_mask,:,l]\n                    axial_idx,sagittal_idx = torch.where(dd == dd.min())\n                    sagittal_assignments.append([l,sagittal_idx[0].item()])\n                    v[l] = axial_indices[level_mask][axial_idx[0]].item()\n            Axial_T2_instance_number.append([study_id,Axial_T2_series]+v.tolist())\n                \n        v = torch.zeros(5)\n        v[:] = torch.nan\n        slices_df = pd.DataFrame(sagittal_assignments).groupby([0]).mean()\n        for i in range(len(slices_df)):\n            row = slices_df.iloc[i]\n            v[\n                {\n                    0:0,\n                    1:1,\n                    2:2,\n                    3:3,\n                    4:4\n                }[row.name]\n            ] = row.values[0]\n    \n        v_mean = v.nanmean()\n        mask = v.isnan()\n        v[mask] = v_mean\n        T2_instance_number.append(v.tolist())","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:14:46.383709Z","iopub.execute_input":"2024-10-07T20:14:46.384057Z","iopub.status.idle":"2024-10-07T20:15:01.145007Z","shell.execute_reply.started":"2024-10-07T20:14:46.384027Z","shell.execute_reply":"2024-10-07T20:15:01.144052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Axial_T2_instane_number_df = pd.DataFrame(Axial_T2_instance_number).groupby([0,1]).mean().reset_index().rename(columns={\n    0:'study_id',\n    1:'series_id',\n    2:'L1L2',\n    3:'L2L3',\n    4:'L3L4',\n    5:'L4L5',\n    6:'L5S1'\n})\nAxial_T2_instane_number_df.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:01.146086Z","iopub.execute_input":"2024-10-07T20:15:01.146344Z","iopub.status.idle":"2024-10-07T20:15:01.165438Z","shell.execute_reply.started":"2024-10-07T20:15:01.146321Z","shell.execute_reply":"2024-10-07T20:15:01.164555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"T1_cases = pd.DataFrame(\n    {\n        'study_id':T1_study_ids,\n        'series_id':T1_series_ids,\n        'flipped':T1_flipped\n    }\n)\nT1_cases[[\n    'L1L2','L2L3','L3L4','L4L5','L5S1'\n]] = T1_instance_number\nT1_cases[[\n    'x_L1L2','y_L1L2',\n    'x_L2L3','y_L2L3',\n    'x_L3L4','y_L3L4',\n    'x_L4L5','y_L4L5',\n    'x_L5S1','y_L5S1'\n]] = T1_levels\nT1_cases.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:01.166684Z","iopub.execute_input":"2024-10-07T20:15:01.166958Z","iopub.status.idle":"2024-10-07T20:15:01.191083Z","shell.execute_reply.started":"2024-10-07T20:15:01.166932Z","shell.execute_reply":"2024-10-07T20:15:01.190031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"T2_cases = pd.DataFrame(\n    {\n        'study_id':T2_study_ids,\n        'series_id':T2_series_ids\n    }\n)\nT2_cases[[\n    'L1L2','L2L3','L3L4','L4L5','L5S1'\n]] = T2_instance_number\nT2_cases[[\n    'x_L1L2','y_L1L2',\n    'x_L2L3','y_L2L3',\n    'x_L3L4','y_L3L4',\n    'x_L4L5','y_L4L5',\n    'x_L5S1','y_L5S1'\n]] = T2_levels\nT2_cases.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:01.192135Z","iopub.execute_input":"2024-10-07T20:15:01.192385Z","iopub.status.idle":"2024-10-07T20:15:01.219234Z","shell.execute_reply.started":"2024-10-07T20:15:01.192363Z","shell.execute_reply":"2024-10-07T20:15:01.218405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    T1_cases.to_csv('T1_cases.csv',index=False)\n    T2_cases.to_csv('T2_cases.csv',index=False)\n    Axial_T2_instane_number_df.to_csv('Axial_T2_instane_number.csv',index=False)\n    \n    with open('Axial_T2_xy.pkl', 'wb') as handle:\n        pickle.dump(Axial_T2_xy, handle, protocol=pickle.HIGHEST_PROTOCOL)\n    with open('Axial_T2_flipped.pkl', 'wb') as handle:\n        pickle.dump(Axial_T2_flipped, handle, protocol=pickle.HIGHEST_PROTOCOL)\n    with open('Axial_T2_assignments.pkl', 'wb') as handle:\n        pickle.dump(Axial_T2_assignments, handle, protocol=pickle.HIGHEST_PROTOCOL)","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:01.224280Z","iopub.execute_input":"2024-10-07T20:15:01.224538Z","iopub.status.idle":"2024-10-07T20:15:01.231458Z","shell.execute_reply.started":"2024-10-07T20:15:01.224495Z","shell.execute_reply":"2024-10-07T20:15:01.230615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del models\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:01.232586Z","iopub.execute_input":"2024-10-07T20:15:01.232860Z","iopub.status.idle":"2024-10-07T20:15:01.412841Z","shell.execute_reply.started":"2024-10-07T20:15:01.232838Z","shell.execute_reply":"2024-10-07T20:15:01.411967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Sagittal T1","metadata":{}},{"cell_type":"code","source":"Lmax = 10\npatch_size = 64\nBS = 32\n\nclass Sagittal_T1_foraminal_Dataset(Dataset):\n    def __init__(self, df, P=patch_size):\n        self.data = df\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 = TEST_PATH + str(row['study_id']) + '/'+str(row['series_id'])\n\n        images = [x for x in glob.glob(sample+'/*.dcm')]\n        images.sort(reverse=False, key=lambda x: int(x.split('/')[-1].replace('.dcm', '')))\n\n        M = row[[\n            'L1L2',\n            'L2L3',\n            'L3L4',\n            'L4L5',\n            '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] = 0\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(\n            image,\n            [\n                self.P//2, self.P - self.P//2,\n                self.P//2, self.P - self.P//2\n            ],'reflect')\n        c[:,1] = c[:,1]*PATCH_H/H + self.P//2\n        c[:,0] = c[:,0]*PATCH_W/W + self.P//2\n        c = c.long()\n\n        crops = torch.stack([\n            image[\n                :,\n                xy[1]-self.P//2:xy[1]+self.P-self.P//2,\n                xy[0]-self.P//2:xy[0]+self.P-self.P//2\n            ] for xy in c\n        ])\n\n        image = torch.zeros(2,5,Lmax,self.P,self.P).to(device)\n        slices_mask = torch.ones(2,5,Lmax).bool().to(device)\n        for i in range(5):\n            if ~missing[i]:\n                if row.flipped:\n                    left_start = int(M[i]) + 1\n                    left_end = min([D,left_start + Lmax])\n                    left_crop = crops[i,left_start:left_end]\n\n                    right_end = int(M[i])\n                    right_start = max([0,right_end - Lmax])\n                    right_crop = crops[i,right_start:right_end]\n\n                    image[1,i,:len(left_crop)] = left_crop\n                    image[0,i,:len(right_crop)] = right_crop.flip(0)\n\n                    slices_mask[1,i,:len(left_crop)] = False\n                    slices_mask[0,i,:len(right_crop)] = False\n                else:\n                    left_start = int(M[i]) + 1\n                    left_end = min([D,left_start + Lmax])\n                    left_crop = crops[i,left_start:left_end]\n\n                    right_end = int(M[i])\n                    right_start = max([0,right_end - Lmax])\n                    right_crop = crops[i,right_start:right_end]\n\n                    image[0,i,:len(left_crop)] = left_crop\n                    image[1,i,:len(right_crop)] = right_crop.flip(0)\n\n                    slices_mask[0,i,:len(left_crop)] = False\n                    slices_mask[1,i,:len(right_crop)] = False\n\n        return row.study_id,image,missing,slices_mask","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:01.413865Z","iopub.execute_input":"2024-10-07T20:15:01.414131Z","iopub.status.idle":"2024-10-07T20:15:01.437058Z","shell.execute_reply.started":"2024-10-07T20:15:01.414109Z","shell.execute_reply":"2024-10-07T20:15:01.436133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = Sagittal_T1_foraminal_Dataset(T1_cases)","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:01.438179Z","iopub.execute_input":"2024-10-07T20:15:01.438449Z","iopub.status.idle":"2024-10-07T20:15:01.451845Z","shell.execute_reply.started":"2024-10-07T20:15:01.438427Z","shell.execute_reply":"2024-10-07T20:15:01.450905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study_id,sample,m,mask = ds.__getitem__(np.random.randint(len(ds)))\nprint(study_id)\nprint(m)\nfor k in range(5):\n    fig, axes = plt.subplots(2, Lmax, figsize=(10,2))\n    for i in range(2):\n        for j in range(Lmax):\n            axes[i,j].imshow(sample.cpu()[i,k,j])\n    plt.show()\n\nplt.imshow(mask.cpu().view(-1,Lmax))","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:01.452838Z","iopub.execute_input":"2024-10-07T20:15:01.453151Z","iopub.status.idle":"2024-10-07T20:15:11.431419Z","shell.execute_reply.started":"2024-10-07T20:15:01.453122Z","shell.execute_reply":"2024-10-07T20:15:11.430484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Sagittal_T1_Foraminal_ViT(nn.Module):\n    def __init__(\n            self,\n            ENCODER,\n            dim=512,\n            depth=24,\n            head_size=64\n        ):\n        super().__init__()\n        self.ENCODER = ENCODER\n        self.slices_enc = SinusoidalPosEmb(dim)(torch.arange(Lmax, device=device).unsqueeze(0))\n        pos_enc = SinusoidalPosEmb(dim)(torch.arange(5, device=device).unsqueeze(0))\n        self.pos_enc = nn.Parameter(pos_enc)\n        self.slices_transformer = nn.TransformerEncoder(\n                nn.TransformerEncoderLayer(d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim,\n                dropout=0.1, activation=nn.GELU(), batch_first=True, norm_first=True, device=device), depth)\n        self.transformer = nn.TransformerEncoder(\n                nn.TransformerEncoderLayer(d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim,\n                dropout=0.1, activation=nn.GELU(), batch_first=True, norm_first=True, device=device), depth)\n        self.proj_out = nn.Linear(dim,3).to(device)\n    \n    def forward(self, x):\n        x,slices_mask = x\n        slices_mask = slices_mask.view(-1,Lmax)\n        mask = slices_mask.sum(-1) < Lmax\n        \n        x = self.ENCODER(x.view(-1,1,patch_size,patch_size))\n\n        x = x.view(-1,Lmax,512)\n        x = x + self.slices_enc\n        x[mask] = self.slices_transformer(x[mask],src_key_padding_mask=slices_mask[mask])\n\n        x[slices_mask] = 0\n        d = (~slices_mask).sum(1).unsqueeze(-1).tile(1,512)\n        x = x.sum(1)\n        x[d > 0] = x[d > 0]/d[d > 0]\n\n        level_mask = (slices_mask.sum(1) == Lmax).view(-1,10)\n        x = x.view(-1,10,512) + torch.concat([self.pos_enc,self.pos_enc],1)\n        x = self.transformer(x,src_key_padding_mask=level_mask)\n        x = self.proj_out(x.view(-1,512)).view(-1,2,5,3)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:11.432713Z","iopub.execute_input":"2024-10-07T20:15:11.433770Z","iopub.status.idle":"2024-10-07T20:15:11.447110Z","shell.execute_reply.started":"2024-10-07T20:15:11.433742Z","shell.execute_reply":"2024-10-07T20:15:11.446236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if len(T1_cases) > 0:\n    \n    foraminal_models = []\n    for path in Sagittal_T1_foraminal_paths:\n        foraminal_models.append(torch.load(path,map_location=device))\n        \n    dl = torch.utils.data.DataLoader(ds, batch_size=BS, shuffle=False, drop_last=False)\n    \n    study_ids = []\n    foraminal_predictions = []\n    with torch.no_grad():\n        OUT = torch.zeros((BS,2,5,3)).to(device)\n        for study_id,X,mask,slices_mask in tqdm(dl):\n            study_ids = study_ids + study_id.tolist()\n            mask = mask.view(-1,1,5).tile(1,2,1)\n            OUT[:] = 0\n            for model in foraminal_models:\n                    y_pred = model([X,slices_mask]).view(-1,2,5,3)\n                    OUT[:len(X)] += nn.Softmax(dim=-1)(y_pred)\n                    \n            OUT[:len(X)] = OUT[:len(X)]/OUT[:len(X)].sum(-1,keepdim=True)\n            OUT[:len(X)][mask] = torch.nan\n\n            foraminal_predictions = foraminal_predictions + OUT[:len(X)].tolist()\n            \n    del ds,foraminal_models,dl,OUT\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:11.448413Z","iopub.execute_input":"2024-10-07T20:15:11.449134Z","iopub.status.idle":"2024-10-07T20:15:15.428158Z","shell.execute_reply.started":"2024-10-07T20:15:11.449103Z","shell.execute_reply":"2024-10-07T20:15:15.427215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"foraminal_predictions = torch.as_tensor(foraminal_predictions)\nids = []\nfor study_id in study_ids:\n    ids = ids + [str(study_id) + '_' + d for d in foraminal]","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:15.431849Z","iopub.execute_input":"2024-10-07T20:15:15.432129Z","iopub.status.idle":"2024-10-07T20:15:15.437689Z","shell.execute_reply.started":"2024-10-07T20:15:15.432107Z","shell.execute_reply":"2024-10-07T20:15:15.436536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"foraminal_predictions = pd.DataFrame({\n    'row_id':ids,\n    'normal_mild':foraminal_predictions[...,0].flatten(),\n    'moderate':foraminal_predictions[...,1].flatten(),\n    'severe':foraminal_predictions[...,2].flatten()\n})\nforaminal_predictions.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:15.440543Z","iopub.execute_input":"2024-10-07T20:15:15.440895Z","iopub.status.idle":"2024-10-07T20:15:15.454819Z","shell.execute_reply.started":"2024-10-07T20:15:15.440863Z","shell.execute_reply":"2024-10-07T20:15:15.453881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Sagittal T2","metadata":{}},{"cell_type":"code","source":"Lmax = 15\nBS = 1\n\nclass Sagittal_T2_Spinal_Dataset(Dataset):\n    def __init__(self, df, P=patch_size):\n        self.data = df\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#       Sanity Check getitem        \n        row = self.data.iloc[index]\n        try:\n            return self.__raw_getitem__(row)\n        except:\n            image = torch.zeros(2,5,Lmax,self.P,self.P).to(device)\n            slices_mask = torch.ones(2,5,Lmax).bool().to(device)\n            missing = torch.ones(5).bool()\n            return int(row.study_id),image,missing,slices_mask\n\n    def __raw_getitem__(self, row):\n                \n        sample = TEST_PATH + str(int(row['study_id'])) + '/'+str(int(row['series_id']))\n\n        images = [x for x in glob.glob(sample+'/*.dcm')]\n        images.sort(reverse=False, key=lambda x: int(x.split('/')[-1].replace('.dcm', '')))\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        instance_numbers = row[[\n            'L1L2',\n            'L2L3',\n            'L3L4',\n            'L4L5',\n            'L5S1'\n        ]].values\n        instance_missing = np.isnan(instance_numbers)\n        instance_numbers[instance_missing] = D/2\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,0] = W/2\n        c[missing,1] = H/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(\n            image,\n            [\n                self.P//2, self.P - self.P//2,\n                self.P//2, self.P - self.P//2\n            ],'reflect')\n        c[:,1] = c[:,1]*PATCH_H/H + self.P//2\n        c[:,0] = c[:,0]*PATCH_W/W + self.P//2\n        c = c.long()\n\n        crops = torch.stack([\n            image[\n                :,\n                xy[1]-self.P//2:xy[1]+self.P-self.P//2,\n                xy[0]-self.P//2:xy[0]+self.P-self.P//2\n            ] for xy in c\n        ])\n\n        image = torch.zeros(2,5,Lmax,self.P,self.P).to(device)\n        slices_mask = torch.ones(2,5,Lmax).bool().to(device)\n        for i in range(5):\n            if ~missing[i]:\n                instance_number = instance_numbers[i].astype(int)\n                start = max([0,instance_number - Lmax//2])\n                end = min([D,start + Lmax])\n                crop = crops[i,start:end]\n\n                image[0,i,:len(crop)] = crop\n                image[1,i,:len(crop)] = crop.flip(0)\n                slices_mask[:,i,:len(crop)] = False\n\n        return int(row.study_id),image,missing,slices_mask","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:15.456120Z","iopub.execute_input":"2024-10-07T20:15:15.456451Z","iopub.status.idle":"2024-10-07T20:15:15.477628Z","shell.execute_reply.started":"2024-10-07T20:15:15.456422Z","shell.execute_reply":"2024-10-07T20:15:15.476772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = Sagittal_T2_Spinal_Dataset(T2_cases)","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:15.478770Z","iopub.execute_input":"2024-10-07T20:15:15.479081Z","iopub.status.idle":"2024-10-07T20:15:15.491213Z","shell.execute_reply.started":"2024-10-07T20:15:15.479051Z","shell.execute_reply":"2024-10-07T20:15:15.490327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study_id,sample,m,mask = ds.__getitem__(np.random.randint(len(ds)))\nprint(study_id)\nprint(m)\nfor k in range(5):\n    fig, axes = plt.subplots(2, Lmax, figsize=(10,2))\n    for i in range(2):\n        for j in range(Lmax):\n            axes[i,j].imshow(sample.cpu()[i,k,j])\n    plt.show()\n\nplt.imshow(mask.cpu().view(-1,Lmax))","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:15.492216Z","iopub.execute_input":"2024-10-07T20:15:15.492538Z","iopub.status.idle":"2024-10-07T20:15:29.954926Z","shell.execute_reply.started":"2024-10-07T20:15:15.492489Z","shell.execute_reply":"2024-10-07T20:15:29.954012Z"},"trusted":true},"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\n    \nclass Sagittal_T2_Spinal_ViT(nn.Module):\n    def __init__(\n            self,\n            ENCODER,\n            dim=512,\n            depth=24,\n            head_size=64\n        ):\n        super().__init__()\n        self.ENCODER = ENCODER\n        self.slices_enc = SinusoidalPosEmb(dim)(torch.arange(Lmax, device=device).unsqueeze(0))\n        self.slices_enc = nn.Parameter(self.slices_enc)\n        pos_enc = SinusoidalPosEmb(dim)(torch.arange(5, device=device).unsqueeze(0))\n        self.pos_enc = nn.Parameter(pos_enc)\n        self.slices_transformer = nn.TransformerEncoder(\n                nn.TransformerEncoderLayer(d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim,\n                dropout=0.1, activation=nn.GELU(), batch_first=True, norm_first=True, device=device), depth)\n        self.transformer = nn.TransformerEncoder(\n                nn.TransformerEncoderLayer(d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim,\n                dropout=0.1, activation=nn.GELU(), batch_first=True, norm_first=True, device=device), depth)\n        self.proj_out = nn.Linear(dim,3).to(device)\n    \n    def forward(self, x):\n        x,slices_mask = x\n        slices_mask = slices_mask.view(-1,Lmax)\n        mask = slices_mask.sum(-1) < Lmax\n        \n        x = self.ENCODER(x.view(-1,1,patch_size,patch_size))\n\n        x = x.view(-1,Lmax,512)\n        x = x + self.slices_enc\n        x[mask] = self.slices_transformer(x[mask],src_key_padding_mask=slices_mask[mask])\n\n        x[slices_mask] = 0\n        d = (~slices_mask).sum(1).unsqueeze(-1).tile(1,512)\n        x = x.sum(1)\n        x[d > 0] = x[d > 0]/d[d > 0]\n\n        level_mask = (slices_mask.sum(1) == Lmax).view(-1,5)\n        mask = level_mask.sum(-1) < 5\n        x = x.view(-1,5,512) + self.pos_enc\n        x[mask] = self.transformer(x[mask],src_key_padding_mask=level_mask[mask])\n        x = self.proj_out(x.view(-1,512)).view(-1,5,3)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:29.956236Z","iopub.execute_input":"2024-10-07T20:15:29.956594Z","iopub.status.idle":"2024-10-07T20:15:29.975192Z","shell.execute_reply.started":"2024-10-07T20:15:29.956561Z","shell.execute_reply":"2024-10-07T20:15:29.974195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if len(T2_cases) > 0:\n    \n    spinal_models = []\n    for path in Sagittal_T2_spinal_paths:\n        spinal_models.append(torch.load(path,map_location=device))\n        \n    dl = torch.utils.data.DataLoader(ds, batch_size=BS, shuffle=False, drop_last=False)\n    \n    study_ids = []\n    spinal_predictions = []\n    with torch.no_grad():\n        OUT = torch.zeros(BS*2,5,3).to(device)\n        for study_id,X,mask,slices_mask in tqdm(dl):\n            study_ids = study_ids + study_id.tolist()\n            OUT[:] = 0\n            for model in spinal_models:\n                OUT[:len(X)*2] += nn.Softmax(dim=-1)(model([X,slices_mask]))\n                    \n            OUT[:len(X)*2] = OUT[:len(X)*2]/OUT[:len(X)*2].sum(-1,keepdim=True)\n            OUT[:len(X)] = OUT[:len(X)*2].view(-1,2,5,3).mean(1)\n            OUT[:len(X)][mask] = torch.nan\n            \n            spinal_predictions = spinal_predictions + OUT[:len(X)].tolist()\n            \n    del ds,spinal_models,OUT,dl\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:29.976324Z","iopub.execute_input":"2024-10-07T20:15:29.976637Z","iopub.status.idle":"2024-10-07T20:15:34.009598Z","shell.execute_reply.started":"2024-10-07T20:15:29.976613Z","shell.execute_reply":"2024-10-07T20:15:34.008665Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spinal_predictions = torch.as_tensor(spinal_predictions)\nids = []\nfor study_id in study_ids:\n    ids = ids + [str(study_id) + '_' + d for d in spinal]","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:34.011705Z","iopub.execute_input":"2024-10-07T20:15:34.012009Z","iopub.status.idle":"2024-10-07T20:15:34.017228Z","shell.execute_reply.started":"2024-10-07T20:15:34.011983Z","shell.execute_reply":"2024-10-07T20:15:34.016148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spinal_predictions = pd.DataFrame({\n    'row_id':ids,\n    'normal_mild':spinal_predictions[...,0].flatten(),\n    'moderate':spinal_predictions[...,1].flatten(),\n    'severe':spinal_predictions[...,2].flatten()\n})\nspinal_predictions.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:34.018278Z","iopub.execute_input":"2024-10-07T20:15:34.018547Z","iopub.status.idle":"2024-10-07T20:15:34.035379Z","shell.execute_reply.started":"2024-10-07T20:15:34.018519Z","shell.execute_reply":"2024-10-07T20:15:34.034459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Axial T2","metadata":{}},{"cell_type":"code","source":"id_df = test_description\nvalid_id = id_df.study_id.unique()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:34.036849Z","iopub.execute_input":"2024-10-07T20:15:34.037112Z","iopub.status.idle":"2024-10-07T20:15:34.044396Z","shell.execute_reply.started":"2024-10-07T20:15:34.037090Z","shell.execute_reply":"2024-10-07T20:15:34.043552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"assignments = {}\nfor i in tqdm(range(len(T2_cases))):\n    row = T2_cases.iloc[i]\n    study_id = int(row['study_id'])\n    axial_t2_ids    = id_df[(id_df.study_id == int(row['study_id'])) & (id_df.series_description=='Axial T2')].series_id\n    sagittal_t2_id = int(row['series_id'])\n    for axial_t2_id in axial_t2_ids:\n        try:\n                sagittal_t2_point = torch.as_tensor([x for x in row[[\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                ]]]).view(5,2).float()\n\n\n                data = read_study(study_id, axial_t2_id=axial_t2_id, sagittal_t2_id=sagittal_t2_id, test=TEST)\n\n\n                #--- step.1 : detect 2d point in sagittal_t2\n                sagittal_t2 = data.sagittal_t2.volume\n                sagittal_t2_df = data.sagittal_t2.df\n                axial_t2_df = data.axial_t2.df\n\n                D,H,W = sagittal_t2.shape\n                \n                sagittal_t2_point[:,0] = sagittal_t2_point[:,0]*256/W\n                sagittal_t2_point[:,1] = sagittal_t2_point[:,1]*256/H\n                \n                sagittal_t2_z = D//2            \n\n                #for debug and development\n                point_hat, z_hat = sagittal_t2_point_hat = get_true_sagittal_t2_point(study_id, sagittal_t2_df)\n                point_hat = point_hat*[[256/W, 256/H]]\n\n\n                #--- step.2 : perdict slice level of axial_t2\n                world_point = view_to_world(sagittal_t2_point, sagittal_t2_z, sagittal_t2_df, 256)\n                assigned_level, closest_z, dis  = axial_t2_level = point_to_level(world_point, axial_t2_df)\n                \n                key = str(study_id)+'_'+str(sagittal_t2_id)+'_'+str(axial_t2_id)\n                assignments[key] = {1:{},2:{},3:{},4:{},5:{}}\n                for k in range(5):\n                    assignments[key][k + 1]['dis'] = dis[k][assigned_level == k + 1]\n                    assignments[key][k + 1]['instance_numbers'] = axial_t2_df.instance_number[assigned_level == k + 1].values\n\n                if i == 0:\n                    print('assigned_level:', assigned_level)\n                    ###################################################################\n                    #visualisation\n                    # https://matplotlib.org/stable/gallery/mplot3d/mixed_subplots.html\n                    fig = plt.figure(figsize=(23, 6))\n                    ax2 = fig.add_subplot(1, 2, 2, projection='3d')\n\n                    # draw  assigned_level\n                    level_ncolor = np.array(level_color) / 255\n                    coloring = level_ncolor[assigned_level].tolist()\n                    draw_slice(\n                        ax2, axial_t2_df,\n                        is_slice=True,   scolor=coloring, salpha=[0.1],\n                        is_border=True,  bcolor=coloring, balpha=[0.2],\n                        is_origin=False, ocolor=[[0, 0, 0]], oalpha=[0.0],\n                        is_arrow=True\n                    )\n\n                    # draw world_point\n                    ax2.scatter(world_point[:, 0], world_point[:, 1], world_point[:, 2], alpha=1, color='black')\n\n\n                    ### draw closest slice\n                    coloring = level_ncolor[1:].tolist()\n                    draw_slice(\n                        ax2, axial_t2_df.iloc[closest_z],\n                        is_slice=True, scolor=coloring, salpha=[0.1],\n                        is_border=True, bcolor=coloring, balpha=[1],\n                        is_origin=False, ocolor=[[1, 0, 0]], oalpha=[0],\n                        is_arrow=False\n                    )\n\n                    ax2.set_aspect('equal')\n                    ax2.set_title(f'axial slice assignment\\n series_id:{sagittal_t2_id}')\n                    ax2.set_xlabel('x')\n                    ax2.set_ylabel('y')\n                    ax2.set_zlabel('z')\n                    ax2.view_init(elev=0, azim=-10, roll=0)\n                    plt.tight_layout(pad=2)\n                    plt.show()\n        \n        except:\n            print(study_id,sagittal_t2_id,axial_t2_id)","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:34.046644Z","iopub.execute_input":"2024-10-07T20:15:34.047039Z","iopub.status.idle":"2024-10-07T20:15:36.265454Z","shell.execute_reply.started":"2024-10-07T20:15:34.047008Z","shell.execute_reply":"2024-10-07T20:15:36.264567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.DataFrame({'key':assignments.keys()})\ndf['study_id'] = df['key'].apply(lambda v:int(v.split('_')[0]))\ndf['series_id'] = df['key'].apply(lambda v:int(v.split('_')[-1]))\ndf.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:36.266864Z","iopub.execute_input":"2024-10-07T20:15:36.267313Z","iopub.status.idle":"2024-10-07T20:15:36.280577Z","shell.execute_reply.started":"2024-10-07T20:15:36.267278Z","shell.execute_reply":"2024-10-07T20:15:36.279766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class subarticular_Dataset(Dataset):\n    def __init__(self, df, P=patch_size):\n        self.data = df\n        self.P = P\n        self.resize = torchvision.transforms.Resize((PATCH_SIZE,PATCH_SIZE),antialias=True)\n        self.indices = torch.arange(Lmax).float()\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, index):\n\n        row = self.data.iloc[index]\n        \n        sample = TEST_PATH+str(int(row['study_id']))+'/'+str(int(row['series_id']))\n\n        images = [x for x in glob.glob(sample+'/*.dcm')]\n        images.sort(key=lambda k:int(k.split('/')[-1].replace('.dcm','')))\n        instance_numbers = [int(k.split('/')[-1].replace('.dcm','')) for k in images]\n        images = [torch.as_tensor(pydicom.dcmread(img).pixel_array.astype('float32')) for img in images]\n        shapes = [img.shape for img in images]\n        H,W = np.array(shapes).max(0)\n        \n        centers = torch.as_tensor(coord[row['study_id']][row['series_id']]).clone().float()\n        levels = torch.as_tensor(Axial_T2_assignments[row['study_id']][row['series_id']]).float().mean(0).long()\n        \n        for l in range(5):\n            mask = levels == l\n            if mask.sum() > 0:\n                level_mean = centers[mask].nanmean(0).view(1,2,2).tile(mask.sum(),1,1)\n                missing = centers[mask].isnan()\n                centers[mask][missing] = level_mean[missing]\n\n        centers_mean = centers.nanmean(0).view(1,2,2).tile(len(centers),1,1)\n        missing = centers.isnan()\n        centers[missing] = centers_mean[missing]\n        \n        c=centers\n        for k in range(len(c)):\n            h,w = shapes[k]\n            c[k,0] += (W - w)//2\n            c[k,1] += (H - h)//2\n\n        images = torch.concat([torch.nn.functional.pad(\n            images[k].unsqueeze(0),(\n                (W - shapes[k][-1])//2,\n                (W - shapes[k][-1]) - (W - shapes[k][-1])//2,\n                (H - shapes[k][-2])//2,\n                (H - shapes[k][-2]) - (H - shapes[k][-2])//2\n            ),\n        mode='reflect') for k in range(len(images))]).float()\n\n        if H > W:\n            d = W\n            h = (H - d)//2\n            c[:,1] -= h\n            images = images[:,h:h+d]\n            H = W\n        elif H < W:\n            d = H\n            w = (W - d)//2\n            c[:,0] -= w\n            images = images[:,:,w:w+d]\n            W = H\n            \n        c[:,0] = c[:,0]*PATCH_SIZE/W\n        c[:,1] = c[:,1]*PATCH_SIZE/H\n\n        images = self.resize(images/images.max()).float().to(device)\n\n        c[c < 64] = torch.nan\n        c[c > 512 - 64] = torch.nan\n        c_mean = torch.nanmean(c, dim=0)\n        instance_to_k = {instance_numbers[k]:k for k in range(len(c))}\n        \n        img = torch.zeros(5,2,Lmax,self.P,self.P)\n        slices_mask = torch.ones(5,2,Lmax).bool()\n        for k in [1,2,3,4,5]:\n            instance_numbers = assignments[row['key']][k]['instance_numbers']\n            if len(instance_numbers) == 0: continue\n            distances = assignments[row['key']][k]['dis']\n            dis_sign = np.sign(distances)\n            if dis_sign[0] != dis_sign[-1]:\n                c_k = torch.stack([c[instance_to_k[i]] for i in instance_numbers])\n                c_k_mean = torch.nanmean(c_k, dim=0)\n                mask =  torch.isnan(c_k_mean)\n                c_k_mean[mask] = c_mean[mask]\n                for i in instance_numbers:\n                    mask = torch.isnan(c[instance_to_k[i]])\n                    c[instance_to_k[i],mask] = c_k_mean[mask]\n        \n                images_k = torch.stack([\n                    torch.stack([\n                        images[\n                            instance_to_k[i],\n                            c[instance_to_k[i],0,1].long()-self.P//2:c[instance_to_k[i],0,1].long()+self.P-self.P//2,\n                            c[instance_to_k[i],0,0].long()-self.P//2:c[instance_to_k[i],0,0].long()+self.P-self.P//2\n                        ] for i in instance_numbers\n                    ]),\n                    torch.stack([\n                        images[\n                            instance_to_k[i],\n                            c[instance_to_k[i],1,1].long()-self.P//2:c[instance_to_k[i],1,1].long()+self.P-self.P//2,\n                            c[instance_to_k[i],1,0].long()-self.P//2:c[instance_to_k[i],1,0].long()+self.P-self.P//2\n                        ] for i in instance_numbers\n                    ]).flip(-1)\n                ])\n\n                if len(distances) > Lmax:\n#                   abs_dist = [abs(v) for v in distances]\n#                   abs_dist.sort()\n#                   dis_th = abs_dist[Lmax]\n#                   img[k-1] = images_k[:,abs(distances) < dis_th]\n#                   slices_mask[k-1] = False\n                    indices = [i for i in range(len(distances))]\n                    indices.sort(key=lambda i:abs(distances[i]))\n                    indices = indices[:Lmax]\n                    indices.sort()\n                    img[k-1] = images_k[:,indices]\n                    slices_mask[k-1] = False\n\n                else:\n                    d = (Lmax - len(distances))//2\n                    img[k-1,:,d:d+len(distances)] = images_k\n                    slices_mask[k-1,:,d:d+len(distances)] = False\n\n        return torch.as_tensor(int(row['study_id'])).to(device),img.to(device),slices_mask.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:36.281914Z","iopub.execute_input":"2024-10-07T20:15:36.282202Z","iopub.status.idle":"2024-10-07T20:15:36.314824Z","shell.execute_reply.started":"2024-10-07T20:15:36.282179Z","shell.execute_reply":"2024-10-07T20:15:36.313905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PATCH_SIZE = 512\ncoord = Axial_T2_xy\nds = subarticular_Dataset(df)","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:36.315929Z","iopub.execute_input":"2024-10-07T20:15:36.317821Z","iopub.status.idle":"2024-10-07T20:15:36.328092Z","shell.execute_reply.started":"2024-10-07T20:15:36.317795Z","shell.execute_reply":"2024-10-07T20:15:36.327350Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample = ds.__getitem__(np.random.randint(len(ds)))\nfig, axesL = plt.subplots(1, 5, figsize=(10,10))\nfig, axesR = plt.subplots(1, 5, figsize=(10,10))\nfor k in range(5):\n    axesL[k].imshow(sample[1][k,0].sum(0).cpu())\n    axesR[k].imshow(sample[1][k,1].sum(0).cpu())\nplt.show()\nplt.imshow(sample[2].cpu().view(-1,Lmax))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:36.329141Z","iopub.execute_input":"2024-10-07T20:15:36.329392Z","iopub.status.idle":"2024-10-07T20:15:38.139007Z","shell.execute_reply.started":"2024-10-07T20:15:36.329366Z","shell.execute_reply":"2024-10-07T20:15:38.138123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Old version\nclass myUNet(nn.Module):\n    def __init__(self):\n        super(myUNet, self).__init__()\n\n        self.UNet = smp.Unet(\n            encoder_name=\"resnet18\",\n            classes=2,\n            in_channels=1\n        ).to(device)\n\n    def forward(self,X):\n        x = self.UNet(X)\n#       MinMaxScaling along the class plane to generate a heatmap\n        min_values = x.view(-1,2,PATCH_SIZE*PATCH_SIZE).min(-1)[0].view(-1,2,1,1)\n        max_values = x.view(-1,2,PATCH_SIZE*PATCH_SIZE).max(-1)[0].view(-1,2,1,1)\n        x = (x - min_values)/(max_values - min_values)\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:38.140164Z","iopub.execute_input":"2024-10-07T20:15:38.140481Z","iopub.status.idle":"2024-10-07T20:15:38.147349Z","shell.execute_reply.started":"2024-10-07T20:15:38.140455Z","shell.execute_reply":"2024-10-07T20:15:38.146412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SinusoidalPosEmb(nn.Module):\n    def __init__(self, dim=16, M=10000):\n        super().__init__()\n        self.dim = dim\n        self.M = M\n\n    def forward(self, x):\n        device = x.device\n        half_dim = self.dim // 2\n        emb = math.log(self.M) / half_dim\n        emb = torch.exp(torch.arange(half_dim, device=device) * (-emb))\n        emb = x[...,None] * emb[None,...]\n        emb = torch.cat((emb.sin(), emb.cos()), dim=-1)\n        return emb\n\nclass myViT(nn.Module):\n    def __init__(self, ENCODER, dim=512, depth=24, head_size=64, **kwargs):\n        super().__init__()\n        self.ENCODER = ENCODER\n        self.AvgPool = nn.AdaptiveAvgPool2d(output_size=1).to(device)\n        self.slices_pos_enc = nn.Parameter(SinusoidalPosEmb(dim)(torch.arange(Lmax, device=device).unsqueeze(0)))\n        self.side_pos_enc = nn.Parameter(SinusoidalPosEmb(dim)(torch.arange(2, device=device).unsqueeze(0)))\n        self.level_pos_enc = nn.Parameter(SinusoidalPosEmb(dim)(torch.arange(5, device=device).unsqueeze(0)))\n        self.slices_transformer = nn.TransformerEncoder(\n                nn.TransformerEncoderLayer(d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim,\n                dropout=0.1, activation=nn.GELU(), batch_first=True, norm_first=True, device=device), 24)\n        self.side_transformer = nn.TransformerEncoder(\n                nn.TransformerEncoderLayer(d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim,\n                dropout=0.1, activation=nn.GELU(), batch_first=True, norm_first=True, device=device), 12)\n        self.level_transformer = nn.TransformerEncoder(\n                nn.TransformerEncoderLayer(d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim,\n                dropout=0.1, activation=nn.GELU(), batch_first=True, norm_first=True, device=device), 12)\n        self.proj_out = nn.Linear(dim,3).to(device)\n    \n    def forward(self, x):\n        x,slices_mask = x\n        \n        '''for kk in range(BS):\n            fig, axesL = plt.subplots(1, 5, figsize=(10,10))\n            fig, axesR = plt.subplots(1, 5, figsize=(10,10))\n            for k in range(5):\n                axesL[k].imshow(x[kk,k,0].sum(0).cpu())\n                axesR[k].imshow(x[kk,k,1].sum(0).cpu())\n            plt.show()'''\n        \n        x = self.ENCODER(x.view(-1,1,patch_size,patch_size))[-1]\n        x = self.AvgPool(x)\n        slices_mask = slices_mask.view(-1,Lmax)\n        mask = slices_mask.sum(-1) < Lmax\n        x = x.view(-1,Lmax,512) + self.slices_pos_enc\n        x[mask] = self.slices_transformer(x[mask],src_key_padding_mask=slices_mask[mask])\n        x[slices_mask] = 0\n        d = (~slices_mask).sum(1).unsqueeze(-1).tile(1,512)\n        x = x.sum(1)\n        x[d > 0] = x[d > 0]/d[d > 0]\n\n        side_mask = slices_mask.view(-1,2,Lmax).sum(-1) == Lmax\n        mask = side_mask.sum(-1) < 2\n        x = x.view(-1,2,512) + self.side_pos_enc\n        x[mask] = self.side_transformer(x[mask],src_key_padding_mask=side_mask[mask])\n\n        level_mask = side_mask.view(-1,5,2).permute(0,2,1).reshape(-1,5)\n        mask = level_mask.sum(-1) < 5\n        x = x.view(-1,5,2,512).permute(0,2,1,3).reshape(-1,5,512)\n        x = x + self.level_pos_enc\n        x[mask] = self.level_transformer(x[mask],src_key_padding_mask=level_mask[mask])\n\n        x = self.proj_out(x.reshape(-1,512)).view(-1,2,5,3)\n\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:38.148559Z","iopub.execute_input":"2024-10-07T20:15:38.148846Z","iopub.status.idle":"2024-10-07T20:15:38.169981Z","shell.execute_reply.started":"2024-10-07T20:15:38.148824Z","shell.execute_reply":"2024-10-07T20:15:38.169122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if len(df) > 0:\n    \n    Lmax = 5\n    BS = 2\n    \n    subarticular_models = []\n    for path in Axial_T2_subarticular_paths:\n        subarticular_models.append(torch.load(path,map_location=device))\n        \n    dl = torch.utils.data.DataLoader(ds, batch_size=BS, shuffle=False, drop_last=False)\n    \n    OUT = torch.zeros((BS*10,3)).to(device)\n    fOUT = torch.zeros((BS*10,3)).to(device)\n    \n    with torch.no_grad():\n        study_ids = []\n        subarticular_predictions = []\n        for study_id,X,mask in tqdm(dl):\n            UNK = (mask.sum(-1) == Lmax).permute(0,2,1).reshape(-1,10)\n            study_ids = study_ids + study_id.tolist()\n            OUT[:] = 0\n            fOUT[:] = 0\n            for model in subarticular_models:\n                    OUT[:len(X)*10] += nn.Softmax(dim=-1)(model([X,mask]).view(-1,3))\n                    fOUT[:len(X)*10] += nn.Softmax(dim=-1)(model([X.flip(2),mask]).view(-1,3))\n                    \n            OUT[:len(X)*10] = (OUT[:len(X)*10]/OUT[:len(X)*10].sum(-1,keepdim=True))\n            fOUT[:len(X)*10] = (fOUT[:len(X)*10]/fOUT[:len(X)*10].sum(-1,keepdim=True))\n            PREDS = OUT.view(-1,10,3)[:len(X)]\n            fPREDS = fOUT.view(-1,10,3)[:len(X)]\n            PREDS[:,5:] += fPREDS[:,:5]\n            PREDS[:,:5] += fPREDS[:,5:]\n            PREDS /= 2\n            PREDS[UNK] = torch.nan#1./3\n            \n            subarticular_predictions = subarticular_predictions + PREDS.tolist()\n            \n    del ds,subarticular_models,dl,OUT\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:38.170986Z","iopub.execute_input":"2024-10-07T20:15:38.171246Z","iopub.status.idle":"2024-10-07T20:15:41.230252Z","shell.execute_reply.started":"2024-10-07T20:15:38.171224Z","shell.execute_reply":"2024-10-07T20:15:41.229049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"subarticular_predictions = torch.as_tensor(subarticular_predictions)\nids = []\nfor study_id in study_ids:\n    ids = ids + [str(study_id) + '_' + d for d in subarticular]","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:41.231975Z","iopub.execute_input":"2024-10-07T20:15:41.232267Z","iopub.status.idle":"2024-10-07T20:15:41.237510Z","shell.execute_reply.started":"2024-10-07T20:15:41.232242Z","shell.execute_reply":"2024-10-07T20:15:41.236684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"subarticular_predictions = pd.DataFrame({\n    'row_id':ids,\n    'normal_mild':subarticular_predictions[...,0].flatten(),\n    'moderate':subarticular_predictions[...,1].flatten(),\n    'severe':subarticular_predictions[...,2].flatten()\n})\nsubarticular_predictions.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:41.238492Z","iopub.execute_input":"2024-10-07T20:15:41.238804Z","iopub.status.idle":"2024-10-07T20:15:41.257326Z","shell.execute_reply.started":"2024-10-07T20:15:41.238781Z","shell.execute_reply":"2024-10-07T20:15:41.256401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"angle = -60\ntheta = (angle/180.) * np.pi\n\nrotMatrix = torch.as_tensor([\n    [np.cos(theta), -np.sin(theta)],\n    [np.sin(theta),  np.cos(theta)]\n]).float().to(device)","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:41.263443Z","iopub.execute_input":"2024-10-07T20:15:41.263829Z","iopub.status.idle":"2024-10-07T20:15:41.269408Z","shell.execute_reply.started":"2024-10-07T20:15:41.263805Z","shell.execute_reply":"2024-10-07T20:15:41.268564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class spinal_Axial_Dataset(Dataset):\n    def __init__(self, df, P=patch_size):\n        self.data = df\n        self.P = P\n        self.resize = torchvision.transforms.Resize((PATCH_SIZE,PATCH_SIZE),antialias=True)\n        self.indices = torch.arange(Lmax).float()\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, index):\n\n        row = self.data.iloc[index]\n        \n        sample = TEST_PATH+str(int(row['study_id']))+'/'+str(int(row['series_id']))\n\n        images = [x.replace('\\\\','/') for x in glob.glob(sample+'/*.dcm')]\n        images.sort(key=lambda k:int(k.split('/')[-1].replace('.dcm','')))\n        instance_numbers = [int(k.split('/')[-1].replace('.dcm','')) for k in images]\n        images = [torch.as_tensor(pydicom.dcmread(img).pixel_array.astype('float32')) for img in images]\n        shapes = [img.shape for img in images]\n        H,W = np.array(shapes).max(0)\n\n        centers = torch.as_tensor(coord[row['study_id']][row['series_id']]).clone().float().to(device)\n        levels = torch.as_tensor(Axial_T2_assignments[row['study_id']][row['series_id']]).float().mean(0).long()\n        \n        for l in range(5):\n            mask = levels == l\n            if mask.sum() > 0:\n                level_mean = centers[mask].nanmean(0).view(1,2,2).tile(mask.sum(),1,1)\n                missing = centers[mask].isnan()\n                centers[mask][missing] = level_mean[missing]\n\n        centers_mean = centers.nanmean(0).view(1,2,2).tile(len(centers),1,1)\n        missing = centers.isnan()\n        centers[missing] = centers_mean[missing]\n        \n        c=centers\n        for k in range(len(c)):\n            h,w = shapes[k]\n            c[k,0] += (W - w)//2\n            c[k,1] += (H - h)//2\n\n        images = torch.concat([torch.nn.functional.pad(\n            images[k].unsqueeze(0),(\n                (W - shapes[k][-1])//2,\n                (W - shapes[k][-1]) - (W - shapes[k][-1])//2,\n                (H - shapes[k][-2])//2,\n                (H - shapes[k][-2]) - (H - shapes[k][-2])//2\n            ),\n        mode='reflect') for k in range(len(images))]).float()\n\n        if H > W:\n            d = W\n            h = (H - d)//2\n            c[:,1] -= h\n            images = images[:,h:h+d]\n            H = W\n        elif H < W:\n            d = H\n            w = (W - d)//2\n            c[:,0] -= w\n            images = images[:,:,w:w+d]\n            W = H\n            \n        c[:,0] = c[:,0]*PATCH_SIZE/W\n        c[:,1] = c[:,1]*PATCH_SIZE/H\n\n        images = self.resize(images/images.max()).float().to(device)\n\n        c = torch.cat([\n            c,\n            (c[:,1] + torch.matmul(c[:,0] - c[:,1],rotMatrix)).unsqueeze(1)\n        ],1)\n\n        c[c < 64] = torch.nan\n        c[c > 512 - 64] = torch.nan\n\n        c_mean = torch.nanmean(c, dim=0)\n        instance_to_k = {instance_numbers[k]:k for k in range(len(c))}\n        \n        img = torch.zeros(5,Lmax,self.P,self.P)\n        slices_mask = torch.ones(5,Lmax).bool()\n        for k in [1,2,3,4,5]:\n            instance_numbers = assignments[row['key']][k]['instance_numbers']\n            if len(instance_numbers) == 0: continue\n            distances = assignments[row['key']][k]['dis']\n            dis_sign = np.sign(distances)\n            if dis_sign[0] != dis_sign[-1]:\n                c_k = torch.stack([c[instance_to_k[i]] for i in instance_numbers])\n                c_k_mean = torch.nanmean(c_k, dim=0)\n                mask =  torch.isnan(c_k_mean)\n                c_k_mean[mask] = c_mean[mask]\n                c_k_mean = c_k_mean.unsqueeze(0).tile(len(c_k),1,1)\n                mask = torch.isnan(c_k)\n                c_k[mask] = c_k_mean[mask]\n                    \n                c_spine = c_k.mean(1)\n                '''print(c_mean)\n                print(c_k_mean)\n                print(c_k)\n                print(c_spine)'''\n        \n                images_k = torch.stack([\n                        images[\n                            instance_to_k[instance_numbers[i]],\n                            c_spine[i,1].long()-self.P//2:c_spine[i,1].long()+self.P-self.P//2,\n                            c_spine[i,0].long()-self.P//2:c_spine[i,0].long()+self.P-self.P//2\n                        ] for i in range(len(distances))#instance_numbers\n                ])\n\n                if len(distances) > Lmax:\n#                   abs_dist = [abs(v) for v in distances]\n#                   abs_dist.sort()\n#                   dis_th = abs_dist[Lmax]\n#                   img[k-1] = images_k[:,abs(distances) < dis_th]\n#                   slices_mask[k-1] = False\n                    indices = [i for i in range(len(distances))]\n                    indices.sort(key=lambda i:abs(distances[i]))\n                    indices = indices[:Lmax]\n                    indices.sort()\n                    img[k-1] = images_k[indices]\n                    slices_mask[k-1] = False\n\n                else:\n                    d = (Lmax - len(distances))//2\n                    img[k-1,d:d+len(distances)] = images_k\n                    slices_mask[k-1,d:d+len(distances)] = False\n\n        return torch.as_tensor(int(row['study_id'])).to(device),img.to(device),slices_mask.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:41.270925Z","iopub.execute_input":"2024-10-07T20:15:41.271187Z","iopub.status.idle":"2024-10-07T20:15:41.302486Z","shell.execute_reply.started":"2024-10-07T20:15:41.271165Z","shell.execute_reply":"2024-10-07T20:15:41.301691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = spinal_Axial_Dataset(df)","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:41.303801Z","iopub.execute_input":"2024-10-07T20:15:41.304541Z","iopub.status.idle":"2024-10-07T20:15:41.316392Z","shell.execute_reply.started":"2024-10-07T20:15:41.304489Z","shell.execute_reply":"2024-10-07T20:15:41.315439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample = ds.__getitem__(np.random.randint(len(ds)))\nfig, axes = plt.subplots(1, 5, figsize=(10,10))\nfor k in range(5):\n    axes[k].imshow(sample[1][k].sum(0).cpu())\nplt.show()\nplt.imshow(sample[2].cpu().view(-1,Lmax))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:41.317581Z","iopub.execute_input":"2024-10-07T20:15:41.317870Z","iopub.status.idle":"2024-10-07T20:15:42.545960Z","shell.execute_reply.started":"2024-10-07T20:15:41.317841Z","shell.execute_reply":"2024-10-07T20:15:42.545024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AxialSpinalViT(nn.Module):\n    def __init__(\n        self,\n        ENCODER,\n        slices_pos_enc,\n        level_pos_enc,\n        slices_transformer,\n        level_transformer,\n        proj_out,\n        dim=512, depth=24, head_size=64, **kwargs\n    ):\n        super().__init__()\n        self.ENCODER = ENCODER\n        self.AvgPool = nn.AdaptiveAvgPool2d(output_size=1).to(device)\n        self.slices_pos_enc = slices_pos_enc\n        self.level_pos_enc = level_pos_enc\n        self.slices_transformer = slices_transformer\n        self.level_transformer = level_transformer\n        self.proj_out = proj_out\n    \n    def forward(self, x):\n        x,slices_mask = x\n        \n        '''for kk in range(BS):\n            fig, axes = plt.subplots(1, 5, figsize=(10,10))\n            for k in range(5):\n                axes[k].imshow(x[kk,k].sum(0).cpu())\n            plt.show()'''\n        \n        x = self.ENCODER(x.view(-1,1,patch_size,patch_size))[-1]\n        x = self.AvgPool(x)\n        slices_mask = slices_mask.view(-1,Lmax)\n        mask = slices_mask.sum(-1) < Lmax\n        x = x.view(-1,Lmax,512) + self.slices_pos_enc\n        x[mask] = self.slices_transformer(x[mask],src_key_padding_mask=slices_mask[mask])\n        x[slices_mask] = 0\n        d = (~slices_mask).sum(1).unsqueeze(-1).tile(1,512)\n        x = x.sum(1)\n        x[d > 0] = x[d > 0]/d[d > 0]\n\n        level_mask = slices_mask.view(-1,5,Lmax).sum(-1) == Lmax\n        mask = level_mask.sum(-1) < 5\n        x = x.view(-1,5,512) + self.level_pos_enc\n        x[mask] = self.level_transformer(x[mask],src_key_padding_mask=level_mask[mask])\n\n        x = self.proj_out(x.view(-1,512)).view(-1,5,3)\n\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:42.547055Z","iopub.execute_input":"2024-10-07T20:15:42.547307Z","iopub.status.idle":"2024-10-07T20:15:42.558991Z","shell.execute_reply.started":"2024-10-07T20:15:42.547285Z","shell.execute_reply":"2024-10-07T20:15:42.558076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if len(df) > 0:\n    \n    spinal_Axial_models = []\n    for path in Axial_T2_spinal_paths:\n        sub_model = torch.load('/kaggle/input/subarticular-dicom-v2-vit/subarticular_DICOM_V2_ViT_1')\n        model = AxialSpinalViT(\n            sub_model.ENCODER,\n            sub_model.slices_pos_enc,\n            sub_model.level_pos_enc,\n            sub_model.slices_transformer,\n            sub_model.level_transformer,\n            sub_model.proj_out\n        )\n        model.load_state_dict(torch.load(path,map_location=device))#torch.load(PATH), weights_only=True)\n        model.eval()\n        spinal_Axial_models.append(model)\n        \n    dl = torch.utils.data.DataLoader(ds, batch_size=BS, shuffle=False, drop_last=False)\n    \n    OUT = torch.zeros((BS*5,3)).to(device)\n    fOUT = torch.zeros((BS*5,3)).to(device)\n    \n    with torch.no_grad():\n        study_ids = []\n        Axial_T2_spinal_predictions = []\n        for study_id,X,mask in tqdm(dl):\n            UNK = mask.view(-1,5,Lmax).sum(-1) == Lmax\n            study_ids = study_ids + study_id.tolist()\n            OUT[:] = 0\n            fOUT[:] = 0\n            for model in spinal_Axial_models:\n                    OUT[:len(X)*5] += nn.Softmax(dim=-1)(model([X,mask]).view(-1,3))\n                    fOUT[:len(X)*5] += nn.Softmax(dim=-1)(model([X.flip(-1),mask]).view(-1,3))\n                    \n            OUT[:len(X)*5] = (OUT[:len(X)*5]/OUT[:len(X)*5].sum(-1,keepdim=True))\n            fOUT[:len(X)*5] = (fOUT[:len(X)*5]/fOUT[:len(X)*5].sum(-1,keepdim=True))\n            PREDS = OUT.view(-1,5,3)[:len(X)]\n            fPREDS = fOUT.view(-1,5,3)[:len(X)]\n            PREDS += fPREDS\n            PREDS /= 2\n            PREDS[UNK] = torch.nan#1./3\n\n            Axial_T2_spinal_predictions = Axial_T2_spinal_predictions + PREDS.tolist()\n            \n    del ds,spinal_Axial_models,dl,OUT\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:42.560084Z","iopub.execute_input":"2024-10-07T20:15:42.560348Z","iopub.status.idle":"2024-10-07T20:15:47.154345Z","shell.execute_reply.started":"2024-10-07T20:15:42.560326Z","shell.execute_reply":"2024-10-07T20:15:47.153389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Axial_T2_spinal_predictions = torch.as_tensor(Axial_T2_spinal_predictions)\nids = []\nfor study_id in study_ids:\n    ids = ids + [str(study_id) + '_' + d for d in spinal]","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:47.156468Z","iopub.execute_input":"2024-10-07T20:15:47.156786Z","iopub.status.idle":"2024-10-07T20:15:47.161817Z","shell.execute_reply.started":"2024-10-07T20:15:47.156755Z","shell.execute_reply":"2024-10-07T20:15:47.160840Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Axial_T2_spinal_predictions = pd.DataFrame({\n    'row_id':ids,\n    'normal_mild':Axial_T2_spinal_predictions[...,0].flatten(),\n    'moderate':Axial_T2_spinal_predictions[...,1].flatten(),\n    'severe':Axial_T2_spinal_predictions[...,2].flatten()\n})\nAxial_T2_spinal_predictions.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:47.162894Z","iopub.execute_input":"2024-10-07T20:15:47.163185Z","iopub.status.idle":"2024-10-07T20:15:47.180526Z","shell.execute_reply.started":"2024-10-07T20:15:47.163163Z","shell.execute_reply":"2024-10-07T20:15:47.179825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"submission = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/sample_submission.csv')\nsubmission.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:47.181770Z","iopub.execute_input":"2024-10-07T20:15:47.182066Z","iopub.status.idle":"2024-10-07T20:15:47.202688Z","shell.execute_reply.started":"2024-10-07T20:15:47.182043Z","shell.execute_reply":"2024-10-07T20:15:47.201666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = pd.concat([\n    foraminal_predictions,\n    spinal_predictions,\n    subarticular_predictions,\n    Axial_T2_spinal_predictions\n]).groupby('row_id').mean().reset_index()\npredictions.tail()","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:47.203907Z","iopub.execute_input":"2024-10-07T20:15:47.204160Z","iopub.status.idle":"2024-10-07T20:15:47.218171Z","shell.execute_reply.started":"2024-10-07T20:15:47.204138Z","shell.execute_reply":"2024-10-07T20:15:47.217294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.DataFrame(submission['row_id']).merge(predictions,left_on='row_id',right_on='row_id',how='left').fillna(1./3)\nv = submission[['normal_mild','moderate','severe']].values\nv = v/v.sum(1).reshape(-1,1)\nsubmission[['normal_mild','moderate','severe']] = v\nsubmission","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:47.219269Z","iopub.execute_input":"2024-10-07T20:15:47.219609Z","iopub.status.idle":"2024-10-07T20:15:47.241340Z","shell.execute_reply.started":"2024-10-07T20:15:47.219580Z","shell.execute_reply":"2024-10-07T20:15:47.240404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv',index=False)","metadata":{"execution":{"iopub.status.busy":"2024-10-07T20:15:47.242390Z","iopub.execute_input":"2024-10-07T20:15:47.242743Z","iopub.status.idle":"2024-10-07T20:15:47.250794Z","shell.execute_reply.started":"2024-10-07T20:15:47.242720Z","shell.execute_reply":"2024-10-07T20:15:47.249944Z"},"trusted":true},"execution_count":null,"outputs":[]}]}