{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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":214311605,"sourceType":"kernelVersion"},{"sourceId":214313282,"sourceType":"kernelVersion"},{"sourceId":214327434,"sourceType":"kernelVersion"},{"sourceId":214443091,"sourceType":"kernelVersion"},{"sourceId":214458863,"sourceType":"kernelVersion"},{"sourceId":214467890,"sourceType":"kernelVersion"}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport gc\nimport sys\nfrom PIL import Image\nimport cv2\nimport math, random\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn.functional as F\nfrom torch import nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision.models.segmentation import lraspp_mobilenet_v3_large\n\nimport timm\n\nimport albumentations as A\n\nimport pydicom","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:29:24.806268Z","iopub.execute_input":"2024-12-22T12:29:24.806558Z","iopub.status.idle":"2024-12-22T12:30:07.878722Z","shell.execute_reply.started":"2024-12-22T12:29:24.806531Z","shell.execute_reply":"2024-12-22T12:30:07.877486Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"rd = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:07.879849Z","iopub.execute_input":"2024-12-22T12:30:07.880511Z","iopub.status.idle":"2024-12-22T12:30:07.885013Z","shell.execute_reply.started":"2024-12-22T12:30:07.880473Z","shell.execute_reply":"2024-12-22T12:30:07.883872Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = 'cuda:0' if torch.cuda.is_available() else 'cpu'\nN_WORKERS = os.cpu_count()\nUSE_AMP = True\nSEED = 8620\n\nIMG_SIZE = [512, 512]\nIN_CHANS = 1\n\nLEVEL_MODEL_NAME = \"edgenext_small.usi_in1k\"\nSEVERITY_MODEL_NAME = \"edgenext_base.usi_in1k\"\nBATCH_SIZE = 1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:07.88713Z","iopub.execute_input":"2024-12-22T12:30:07.887403Z","iopub.status.idle":"2024-12-22T12:30:07.908191Z","shell.execute_reply.started":"2024-12-22T12:30:07.887379Z","shell.execute_reply":"2024-12-22T12:30:07.907162Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(f'{rd}/test_series_descriptions.csv')\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:07.909506Z","iopub.execute_input":"2024-12-22T12:30:07.909837Z","iopub.status.idle":"2024-12-22T12:30:07.955182Z","shell.execute_reply.started":"2024-12-22T12:30:07.909807Z","shell.execute_reply":"2024-12-22T12:30:07.953983Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def read_dcm(src_path):\n    dicom_data = pydicom.dcmread(src_path)\n    image = dicom_data.pixel_array\n    image = (image - image.min()) / (image.max() - image.min() +1e-6) * 255\n    image = image.astype('uint8')[:,:,None]\n    return image\n\ndef make_masks(row,series_desc='Axial T2'):\n\n    allimgs=glob(f\"{rd}/test_images/{row['study_id']}/{row['series_id']}/*.dcm\")\n    test_images_dir=f\"test_images/{row['study_id']}/{series_desc}\"\n    \n    os.makedirs(test_images_dir,exist_ok=True)  \n    i=len(os.listdir(test_images_dir))\n    for img_path in allimgs:\n        img=read_dcm(img_path)\n\n        cv2.imwrite(f\"{test_images_dir}/{i}.jpg\",img)\n        i+=1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:07.956512Z","iopub.execute_input":"2024-12-22T12:30:07.956871Z","iopub.status.idle":"2024-12-22T12:30:07.963536Z","shell.execute_reply.started":"2024-12-22T12:30:07.95684Z","shell.execute_reply":"2024-12-22T12:30:07.962338Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"axial_desc=df[df['series_description']=='Axial T2']\nsagt1_desc=df[df['series_description']=='Sagittal T1']\nsagt2_desc=df[df['series_description']=='Sagittal T2/STIR']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:07.964734Z","iopub.execute_input":"2024-12-22T12:30:07.965145Z","iopub.status.idle":"2024-12-22T12:30:07.988572Z","shell.execute_reply.started":"2024-12-22T12:30:07.965103Z","shell.execute_reply":"2024-12-22T12:30:07.987387Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i, row in axial_desc.iterrows():\n    make_masks(row=row,series_desc='Axial T2')\n\nprint('Axial T2 done')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:07.98964Z","iopub.execute_input":"2024-12-22T12:30:07.989966Z","iopub.status.idle":"2024-12-22T12:30:08.715673Z","shell.execute_reply.started":"2024-12-22T12:30:07.989936Z","shell.execute_reply":"2024-12-22T12:30:08.71454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i, row in sagt1_desc.iterrows():\n    make_masks(row=row,series_desc='Sagittal T1')\n\nprint('Sagittal T1 done')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:08.718608Z","iopub.execute_input":"2024-12-22T12:30:08.718929Z","iopub.status.idle":"2024-12-22T12:30:09.607536Z","shell.execute_reply.started":"2024-12-22T12:30:08.718901Z","shell.execute_reply":"2024-12-22T12:30:09.606441Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i, row in sagt2_desc.iterrows():\n    make_masks(row=row,series_desc='Sagittal T2_STIR')\n\nprint('Sagittal T2 done')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:09.60921Z","iopub.execute_input":"2024-12-22T12:30:09.609876Z","iopub.status.idle":"2024-12-22T12:30:10.785349Z","shell.execute_reply.started":"2024-12-22T12:30:09.60984Z","shell.execute_reply":"2024-12-22T12:30:10.784188Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"study_ids = list(df['study_id'].unique())\nsample_sub = pd.read_csv(f'{rd}/sample_submission.csv')\nLABELS = list(sample_sub.columns[1:])\nLABELS","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:10.786486Z","iopub.execute_input":"2024-12-22T12:30:10.786895Z","iopub.status.idle":"2024-12-22T12:30:10.804054Z","shell.execute_reply.started":"2024-12-22T12:30:10.786855Z","shell.execute_reply":"2024-12-22T12:30:10.802842Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RSNA24Dataset(Dataset):\n    def __init__(self, folder_path, view='Sagittal T2_STIR', transform=None):\n        self.images = glob(f'{folder_path}/*/{view}/*.jpg')\n        self.transform = transform\n    \n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        img_path = self.images[idx]\n        study_id = int(img_path.split('/')[-3])\n\n        img = Image.open(img_path).convert('L')\n        img = np.array(img).astype(np.uint8)\n            \n        if self.transform is not None:\n            img = self.transform(image=img)['image']\n\n        img = img[None]\n                \n        return img, study_id","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:10.8055Z","iopub.execute_input":"2024-12-22T12:30:10.805921Z","iopub.status.idle":"2024-12-22T12:30:10.813264Z","shell.execute_reply.started":"2024-12-22T12:30:10.805878Z","shell.execute_reply":"2024-12-22T12:30:10.812226Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transforms_test = A.Compose([\n    A.Resize(IMG_SIZE[0], IMG_SIZE[1]),\n    A.Normalize(mean=0.5, std=0.5)\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:10.814204Z","iopub.execute_input":"2024-12-22T12:30:10.814538Z","iopub.status.idle":"2024-12-22T12:30:10.833573Z","shell.execute_reply.started":"2024-12-22T12:30:10.814512Z","shell.execute_reply":"2024-12-22T12:30:10.832388Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"levels=['l1_l2','l2_l3','l3_l4','l4_l5','l5_s1']\nsagt2_conditions=['spinal_canal_stenosis']\nsagt1_conditions=['left_neural_foraminal_narrowing', 'right_neural_foraminal_narrowing']\naxial_conditions=['left_subarticular_stenosis','right_subarticular_stenosis']\n\nsagt1_condition_levels = np.array([cond+'_'+lv for cond in sagt1_conditions for lv in levels])\nsagt2_condition_levels = np.array([cond+'_'+lv for cond in sagt2_conditions for lv in levels])\naxial_condition_levels = np.array([cond+'_'+lv for cond in axial_conditions for lv in levels])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T13:08:03.668445Z","iopub.execute_input":"2024-12-22T13:08:03.668843Z","iopub.status.idle":"2024-12-22T13:08:03.67529Z","shell.execute_reply.started":"2024-12-22T13:08:03.668809Z","shell.execute_reply":"2024-12-22T13:08:03.674183Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"axial_test_ds = RSNA24Dataset('/kaggle/working/test_images/', view='Axial T2', transform=transforms_test)\naxial_test_dl = DataLoader(\n    axial_test_ds, \n    batch_size=1, \n    shuffle=False,\n    num_workers=N_WORKERS,\n    pin_memory=True,\n    drop_last=False\n)\n\nsagt1_test_ds = RSNA24Dataset('/kaggle/working/test_images/', view='Sagittal T1', transform=transforms_test)\nsagt1_test_dl = DataLoader(\n    sagt1_test_ds, \n    batch_size=1, \n    shuffle=False,\n    num_workers=N_WORKERS,\n    pin_memory=True,\n    drop_last=False\n)\n\nsagt2_test_ds = RSNA24Dataset('/kaggle/working/test_images/', view='Sagittal T2_STIR', transform=transforms_test)\nsagt2_test_dl = DataLoader(\n    sagt2_test_ds, \n    batch_size=1, \n    shuffle=False,\n    num_workers=N_WORKERS,\n    pin_memory=True,\n    drop_last=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:10.850825Z","iopub.execute_input":"2024-12-22T12:30:10.851122Z","iopub.status.idle":"2024-12-22T12:30:10.868731Z","shell.execute_reply.started":"2024-12-22T12:30:10.851091Z","shell.execute_reply":"2024-12-22T12:30:10.867817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DetectModel(nn.Module):\n    def __init__(self, model_name, in_c=3, n_classes=5, pretrained=False, features_only=False):\n        super().__init__()\n        self.model = timm.create_model(model_name,\n                                        pretrained=pretrained, \n                                        features_only=features_only,\n                                        in_chans=in_c,\n                                        num_classes=n_classes,\n                                        global_pool='avg'\n                                        )\n    def forward(self, x):\n        y = self.model(x)\n        return y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:42:10.416124Z","iopub.execute_input":"2024-12-22T12:42:10.416521Z","iopub.status.idle":"2024-12-22T12:42:10.422308Z","shell.execute_reply.started":"2024-12-22T12:42:10.416482Z","shell.execute_reply":"2024-12-22T12:42:10.421061Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SegmentModel(nn.Module):\n    def __init__(self, in_c=1, n_classes=2, pretrained=False, features_only=False):\n        super().__init__()\n\n        self.model = lraspp_mobilenet_v3_large(weights=None,weights_backbone=None)\n        self.model.backbone['0'][0]=nn.Conv2d(in_c, 16, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\n        self.model.classifier.high_classifier=nn.Conv2d(128, n_classes, kernel_size=(1, 1), stride=(1, 1))\n        self.model.classifier.low_classifier=nn.Conv2d(40, n_classes, kernel_size=(1, 1), stride=(1, 1))\n    def forward(self, x):\n        y = self.model(x)[\"out\"]\n        return y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:10.882521Z","iopub.execute_input":"2024-12-22T12:30:10.882917Z","iopub.status.idle":"2024-12-22T12:30:10.901625Z","shell.execute_reply.started":"2024-12-22T12:30:10.882863Z","shell.execute_reply":"2024-12-22T12:30:10.900523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RSNA24Model(nn.Module):\n    def __init__(self, model_name, in_c=2, n_classes=30, pretrained=False, features_only=False):\n        super().__init__()\n        self.model = timm.create_model(model_name,\n                                       in_chans=in_c,\n                                       num_classes=n_classes,\n                                       features_only=features_only,\n                                       pretrained=pretrained\n                                       )\n    def forward(self, x):\n        y = self.model(x)\n        return y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:30:10.902672Z","iopub.execute_input":"2024-12-22T12:30:10.903063Z","iopub.status.idle":"2024-12-22T12:30:10.927566Z","shell.execute_reply.started":"2024-12-22T12:30:10.903024Z","shell.execute_reply":"2024-12-22T12:30:10.926248Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"axial = RSNA24Model(SEVERITY_MODEL_NAME, in_c=IN_CHANS, n_classes=len(LABELS)*len(axial_conditions))\nsagt1 = RSNA24Model(SEVERITY_MODEL_NAME, in_c=IN_CHANS, n_classes=len(LABELS)*len(sagt1_condition_levels))\nsagt2 = RSNA24Model(SEVERITY_MODEL_NAME, in_c=IN_CHANS, n_classes=len(LABELS)*len(sagt2_condition_levels))\n\naxial_detect = DetectModel(LEVEL_MODEL_NAME, in_c=IN_CHANS, n_classes=2*len(axial_condition_levels))\nsagt1_detect = DetectModel(LEVEL_MODEL_NAME, in_c=IN_CHANS, n_classes=2*len(sagt1_condition_levels))\nsagt2_detect = DetectModel(LEVEL_MODEL_NAME, in_c=IN_CHANS, n_classes=2*len(sagt2_condition_levels))\n\n#axial_segment=SegmentModel()\n#sagt1_segment=SegmentModel()\n#sagt2_segment=SegmentModel()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:55:40.489373Z","iopub.execute_input":"2024-12-22T12:55:40.489857Z","iopub.status.idle":"2024-12-22T12:55:42.757146Z","shell.execute_reply.started":"2024-12-22T12:55:40.489819Z","shell.execute_reply":"2024-12-22T12:55:42.756283Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"axial.load_state_dict(torch.load('/kaggle/input/fork-of-rsna2024-no-masked-training-axial/rsna24-results/best_wll_model_Axial T2_fold-0.pt',map_location=device,weights_only=True))\nsagt1.load_state_dict(torch.load('/kaggle/input/fork-of-rsna2024-no-masked-training-sagt1/rsna24-results/best_wll_model_Sagittal T1_fold-0.pt',map_location=device,weights_only=True))\nsagt2.load_state_dict(torch.load('/kaggle/input/fork-of-rsna2024-no-masked-training-sagt2/rsna24-results/best_wll_model_Sagittal T2_STIR_fold-0.pt',map_location=device,weights_only=True))\n\naxial_detect.load_state_dict(torch.load('/kaggle/input/rsna-training-label-detection-axial/rsna-results/best_wll_model.pt',map_location=device,weights_only=True))\nsagt1_detect.load_state_dict(torch.load('/kaggle/input/rsna-training-label-detection-sagt1/rsna-results/best_wll_model.pt',map_location=device,weights_only=True))\nsagt2_detect.load_state_dict(torch.load('/kaggle/input/rsna-training-label-detection-sagt2/rsna-results/best_wll_model.pt',map_location=device,weights_only=True))\n\n# axial_segment.load_state_dict(torch.load('/kaggle/input/fork-of-rsna2024-yolo-segment-mask-training-axial/rsna24-results/best_wll_model_Axial T2_fold-0.pt',map_location=device,weights_only=True))\n# sagt1_segment.load_state_dict(torch.load('/kaggle/input/fork-of-rsna2024-yolo-segment-mask-training-sagt1/rsna24-results/best_wll_model_Sagittal T1_fold-0.pt',map_location=device,weights_only=True))\n# sagt2_segment.load_state_dict(torch.load('/kaggle/input/fork-of-rsna2024-yolo-segment-mask-training-sagt2/rsna24-results/best_wll_model_Sagittal T2_STIR_fold-0.pt',map_location=device,weights_only=True))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:55:42.758595Z","iopub.execute_input":"2024-12-22T12:55:42.758941Z","iopub.status.idle":"2024-12-22T12:55:43.470708Z","shell.execute_reply.started":"2024-12-22T12:55:42.758901Z","shell.execute_reply":"2024-12-22T12:55:43.469746Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"axial = axial.eval().to(device)\nsagt1 = sagt1.eval().to(device)\nsagt2 = sagt2.eval().to(device)\n\naxial_detect = axial_detect.eval().to(device)\nsagt1_detect = sagt1_detect.eval().to(device)\nsagt2_detect = sagt2_detect.eval().to(device)\n\n# axial_segment = axial_segment.eval().to(device)\n# sagt1_segment = sagt1_segment.eval().to(device)\n# sagt2_segment = sagt2_segment.eval().to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T12:55:46.840602Z","iopub.execute_input":"2024-12-22T12:55:46.840995Z","iopub.status.idle":"2024-12-22T12:55:46.880131Z","shell.execute_reply.started":"2024-12-22T12:55:46.840962Z","shell.execute_reply":"2024-12-22T12:55:46.878943Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"autocast = torch.cuda.amp.autocast(enabled=USE_AMP, dtype=torch.half)\ny_preds = []\nrow_names = []\n\nwith tqdm(axial_test_dl, leave=True) as pbar:\n    with torch.no_grad():\n        for idx, (x, si) in enumerate(pbar):\n            x = x.to(device)\n            si = str(si.item())\n\n            with autocast:\n                y_detect = axial_detect(x)\n                y_detect = y_detect.reshape(2,-1).argmax(0)\n\n                if not y_detect.any():\n                    continue\n                label_indices = y_detect.nonzero().cpu()\n                pred_per_study = np.zeros((len(label_indices), 3))\n                if len(label_indices)==1:\n                    labels = [axial_condition_levels[label_indices]]\n                else:\n                    labels = axial_condition_levels[label_indices][:,0]\n\n                # y_segment = axial_segment(x)\n                # segmented_x = y_segment.argmax(1)*x\n                \n                # y = axial(segmented_x)[0].cpu()\n                y = axial(x)[0].cpu()\n                \n                for i, cond in enumerate(labels):\n                    if 'right' in cond:\n                        pred = y[3:]\n                    else:\n                        pred = y[:3]\n                    row_names.append(si + '_' + cond)\n                    \n                    y_pred=pred.float().softmax(0).cpu().numpy()\n                    pred_per_study[i]=y_pred\n                y_preds.append(pred_per_study)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T13:54:09.521948Z","iopub.execute_input":"2024-12-22T13:54:09.522359Z","iopub.status.idle":"2024-12-22T13:54:19.202046Z","shell.execute_reply.started":"2024-12-22T13:54:09.522326Z","shell.execute_reply":"2024-12-22T13:54:19.200831Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with tqdm(sagt1_test_dl, leave=True) as pbar:\n    with torch.no_grad():\n        for idx, (x, si) in enumerate(pbar):\n            x = x.to(device)\n            si = str(si.item())\n            pred_per_study = np.zeros((len(sagt1_condition_levels), 3))\n\n            with autocast:\n                y_detect = sagt1_detect(x)\n                y_detect = y_detect.reshape(2,-1).argmax(0)\n\n                if not y_detect.any():\n                    continue\n\n                label_indices = y_detect.nonzero().cpu()\n                pred_per_study = np.zeros((len(label_indices), 3))\n                if len(label_indices)==1:\n                    labels = [sagt1_condition_levels[label_indices]]\n                else:\n                    labels = sagt1_condition_levels[label_indices][:,0]\n                \n                # y_segment = sagt1_segment(x)\n                # segmented_x = y_segment.argmax(1)*x\n\n                # y = sagt1(segmented_x)[0].reshape(-1,3).cpu()\n                # y = y[label_indices][:,0]\n                y = sagt1(x)[0].reshape(-1,3).cpu()\n                y = y[label_indices][:,0]\n                \n                for i, cond in enumerate(labels):\n                    row_names.append(si + '_' + cond)\n                    \n                    pred = y[i]\n                    y_pred=pred.float().softmax(0).cpu().numpy()\n                    pred_per_study[i]=y_pred\n                y_preds.append(pred_per_study)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T13:55:25.184111Z","iopub.execute_input":"2024-12-22T13:55:25.184536Z","iopub.status.idle":"2024-12-22T13:55:31.336477Z","shell.execute_reply.started":"2024-12-22T13:55:25.184498Z","shell.execute_reply":"2024-12-22T13:55:31.334733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with tqdm(sagt2_test_dl, leave=True) as pbar:\n    with torch.no_grad():\n        for idx, (x, si) in enumerate(pbar):\n            x = x.to(device)\n            si = str(si.item())\n            pred_per_study = np.zeros((len(sagt2_condition_levels), 3))       \n\n            with autocast:\n                y_detect = sagt2_detect(x)\n                y_detect = y_detect.reshape(2,-1).argmax(0)\n\n                if not y_detect.any():\n                    continue\n\n                label_indices = y_detect.nonzero().cpu()\n                pred_per_study = np.zeros((len(label_indices), 3))\n                if len(label_indices)==1:\n                    labels = [sagt2_condition_levels[label_indices]]\n                else:\n                    labels = sagt2_condition_levels[label_indices][:,0]\n                \n                # y_segment = sagt2_segment(x)\n                # segmented_x = y_segment.argmax(1)*x\n                \n                # y = sagt2(segmented_x)[0].reshape(-1,3).cpu()\n                # y = y[label_indices][:,0]\n                y = sagt2(x)[0].reshape(-1,3).cpu()\n                y = y[label_indices][:,0]\n    \n                for i, cond in enumerate(labels):\n                    row_names.append(si + '_' + cond)\n                    \n                    pred = y[i]\n                    y_pred=pred.float().softmax(0).cpu().numpy()\n                    pred_per_study[i]=y_pred\n                y_preds.append(pred_per_study)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T13:57:07.372261Z","iopub.execute_input":"2024-12-22T13:57:07.372678Z","iopub.status.idle":"2024-12-22T13:57:11.692043Z","shell.execute_reply.started":"2024-12-22T13:57:07.372615Z","shell.execute_reply":"2024-12-22T13:57:11.690892Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CONDITIONS=['left_subarticular_stenosis','right_subarticular_stenosis',\n            'left_neural_foraminal_narrowing', 'right_neural_foraminal_narrowing',\n            'spinal_canal_stenosis'\n           ]\nlevels=[\n        'l1_l2',\n        'l2_l3',\n        'l3_l4',\n        'l4_l5',\n        'l5_s1'\n        ]\n#Ensures that if a label is missing, it will be replaced with default value\npred=np.array([[0.60,0.25,0.15]])\nfor st_id in study_ids:\n    for cond in CONDITIONS:\n        for lv in levels:\n            if f\"{st_id}_{cond}_{lv}\" not in row_names:\n                row_names.append(f\"{st_id}_{cond}_{lv}\")\n                y_preds.append(pred)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T14:00:13.046523Z","iopub.execute_input":"2024-12-22T14:00:13.046988Z","iopub.status.idle":"2024-12-22T14:00:13.05601Z","shell.execute_reply.started":"2024-12-22T14:00:13.04694Z","shell.execute_reply":"2024-12-22T14:00:13.054234Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_preds = np.concatenate(y_preds, axis=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T14:00:18.245293Z","iopub.execute_input":"2024-12-22T14:00:18.245621Z","iopub.status.idle":"2024-12-22T14:00:18.250933Z","shell.execute_reply.started":"2024-12-22T14:00:18.245589Z","shell.execute_reply":"2024-12-22T14:00:18.249444Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub = pd.DataFrame()\nsub['row_id'] = row_names\nsub[LABELS] = y_preds\nsub=sub.groupby('row_id').mean()#.reset_index()\nsub.head(25)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T14:00:20.581844Z","iopub.execute_input":"2024-12-22T14:00:20.582197Z","iopub.status.idle":"2024-12-22T14:00:20.613244Z","shell.execute_reply.started":"2024-12-22T14:00:20.582169Z","shell.execute_reply":"2024-12-22T14:00:20.611777Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub.to_csv('submission.csv')#,index=False)\npd.read_csv('submission.csv').head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T14:01:05.835946Z","iopub.execute_input":"2024-12-22T14:01:05.83635Z","iopub.status.idle":"2024-12-22T14:01:05.85331Z","shell.execute_reply.started":"2024-12-22T14:01:05.836317Z","shell.execute_reply":"2024-12-22T14:01:05.852161Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!rm -r /kaggle/working/test_images","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-22T14:01:09.698996Z","iopub.execute_input":"2024-12-22T14:01:09.699392Z","iopub.status.idle":"2024-12-22T14:01:09.880988Z","shell.execute_reply.started":"2024-12-22T14:01:09.699359Z","shell.execute_reply":"2024-12-22T14:01:09.87943Z"}},"outputs":[],"execution_count":null}]}