{"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":9063849,"sourceType":"datasetVersion","datasetId":5466169},{"sourceId":9113037,"sourceType":"datasetVersion","datasetId":5500484},{"sourceId":183439689,"sourceType":"kernelVersion"},{"sourceId":4534,"sourceType":"modelInstanceVersion","modelInstanceId":3326,"modelId":986}],"dockerImageVersionId":30747,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#!unzip -q /kaggle/input/rsna2024-lsdc-making-dataset/_output_.zip \nimport sys\nimport zipfile\nfrom pathlib import Path\n\ntimm_path = Path('/kaggle/input/timm-folder/timm')\nsys.path.append(str(timm_path))\nimport timm","metadata":{"execution":{"iopub.status.busy":"2024-08-05T19:51:17.953884Z","iopub.execute_input":"2024-08-05T19:51:17.954186Z","iopub.status.idle":"2024-08-05T19:51:24.471639Z","shell.execute_reply.started":"2024-08-05T19:51:17.954160Z","shell.execute_reply":"2024-08-05T19:51:24.470671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport timm\nimport torch.nn as nn\nimport glob\nimport numpy as np\nimport math\nimport pandas as pd\nfrom collections import OrderedDict\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\nimport random\nimport albumentations as A\nimport os\nfrom transformers import AutoImageProcessor, AutoModel\nimport re","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-08-05T19:51:24.473402Z","iopub.execute_input":"2024-08-05T19:51:24.473896Z","iopub.status.idle":"2024-08-05T19:51:37.930104Z","shell.execute_reply.started":"2024-08-05T19:51:24.473861Z","shell.execute_reply":"2024-08-05T19:51:37.929064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"OUTPUT_DIR = f'/kaggle/input/rsna2024-lsdc-training-baseline/rsna24-results'\n\ndevice = 'cuda:0' if torch.cuda.is_available() else 'cpu'\nN_WORKERS = os.cpu_count()\nSEED = 8620\n\nIMG_SIZE = [512, 512]\nIN_CHANS = 42\nN_FOLDS = [2, 3]\nN_LABELS = 25\nN_CLASSES = 3 * N_LABELS\nUSE_AMP = True\nBATCH_SIZE = 1\nKFold_models = True","metadata":{"execution":{"iopub.status.busy":"2024-08-05T19:51:37.931229Z","iopub.execute_input":"2024-08-05T19:51:37.931817Z","iopub.status.idle":"2024-08-05T19:51:37.968671Z","shell.execute_reply.started":"2024-08-05T19:51:37.931790Z","shell.execute_reply":"2024-08-05T19:51:37.967701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNA24Model(nn.Module):\n    def __init__(self, model_name, in_c=42, n_classes=75, pretrained=False, features_only=False, freeze_stages=2):\n        super().__init__()\n        self.model = timm.create_model(\n                                    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        for i, stage in enumerate(self.model.stages):\n            if i < freeze_stages:\n                for param in stage.parameters():\n                    param.requires_grad = False\n        \n        for stem in self.model.stem:\n            for param in stem.parameters():\n                param.requires_grad = False\n    \n    def forward(self, x):\n        y = self.model(x)\n        return y","metadata":{"execution":{"iopub.status.busy":"2024-08-05T19:51:37.971271Z","iopub.execute_input":"2024-08-05T19:51:37.972000Z","iopub.status.idle":"2024-08-05T19:51:38.001673Z","shell.execute_reply.started":"2024-08-05T19:51:37.971966Z","shell.execute_reply":"2024-08-05T19:51:38.000811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rd = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\ndf = pd.read_csv(f'{rd}/test_series_descriptions.csv')\nstudy_ids = list(df['study_id'].unique())\nsample_sub = pd.read_csv(f'{rd}/sample_submission.csv')\nLABELS = list(sample_sub.columns[1:])\n\nCONDITIONS = [\n    'spinal_canal_stenosis', \n    'left_neural_foraminal_narrowing', \n    'right_neural_foraminal_narrowing',\n    'left_subarticular_stenosis',\n    'right_subarticular_stenosis'\n]\n\nLEVELS = [\n    'l1_l2',\n    'l2_l3',\n    'l3_l4',\n    'l4_l5',\n    'l5_s1',\n]\n\ndef atoi(text):\n    return int(text) if text.isdigit() else text\n\ndef natural_keys(text):\n    return [ atoi(c) for c in re.split(r'(\\d+)', text) ]","metadata":{"execution":{"iopub.status.busy":"2024-08-05T19:51:38.002749Z","iopub.execute_input":"2024-08-05T19:51:38.003199Z","iopub.status.idle":"2024-08-05T19:51:38.033994Z","shell.execute_reply.started":"2024-08-05T19:51:38.003174Z","shell.execute_reply":"2024-08-05T19:51:38.033197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNA24TestDataset(Dataset):\n    def __init__(self, df, study_ids, phase='test', transform=None):\n        self.df = df\n        self.study_ids = study_ids\n        self.transform = transform\n        self.phase = phase\n    \n    def __len__(self):\n        return len(self.study_ids)\n    \n    def get_img_paths(self, study_id, series_desc):\n        pdf = self.df[self.df['study_id']==study_id]\n        pdf_ = pdf[pdf['series_description']==series_desc]\n        allimgs = []\n        for i, row in pdf_.iterrows():\n            pimgs = glob.glob(f'{rd}/test_images/{study_id}/{row[\"series_id\"]}/*.dcm')\n            pimgs = sorted(pimgs, key=natural_keys)\n            allimgs.extend(pimgs)\n            \n        return allimgs\n    \n    def read_dcm_ret_arr(self, 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        img = cv2.resize(image, (IMG_SIZE[0], IMG_SIZE[1]),interpolation=cv2.INTER_CUBIC)\n        assert img.shape==(IMG_SIZE[0], IMG_SIZE[1])\n        return img\n\n    def __getitem__(self, idx):\n        x = np.zeros((IMG_SIZE[0], IMG_SIZE[1], IN_CHANS), dtype=np.uint8)\n        st_id = self.study_ids[idx]        \n        \n        # Sagittal T1\n        allimgs_st1 = self.get_img_paths(st_id, 'Sagittal T1')\n        if len(allimgs_st1)==0:\n            print(st_id, ': Sagittal T1, has no images')\n        \n        else:\n            step = len(allimgs_st1) / 14.0\n            st = len(allimgs_st1)/2.0 - 6.0*step\n            end = len(allimgs_st1)+0.0001\n            for j, i in enumerate(np.arange(st, end, step)):\n                try:\n                    ind2 = max(0, int((i-0.5001).round()))\n                    img = self.read_dcm_ret_arr(allimgs_st1[ind2])\n                    x[..., j] = img.astype(np.uint8)\n                except:\n                    print(f'failed to load on {st_id}, Sagittal T1')\n                    pass\n            \n        # Sagittal T2/STIR\n        allimgs_st2 = self.get_img_paths(st_id, 'Sagittal T2/STIR')\n        if len(allimgs_st2)==0:\n            print(st_id, ': Sagittal T2/STIR, has no images')\n            \n        else:\n            step = len(allimgs_st2) / 14.0\n            st = len(allimgs_st2)/2.0 - 6.0*step\n            end = len(allimgs_st2)+0.0001\n            for j, i in enumerate(np.arange(st, end, step)):\n                try:\n                    ind2 = max(0, int((i-0.5001).round()))\n                    img = self.read_dcm_ret_arr(allimgs_st2[ind2])\n                    x[..., j+14] = img.astype(np.uint8)\n                except:\n                    print(f'failed to load on {st_id}, Sagittal T2/STIR')\n                    pass\n            \n        # Axial T2\n        allimgs_at2 = self.get_img_paths(st_id, 'Axial T2')\n        if len(allimgs_at2)==0:\n            print(st_id, ': Axial T2, has no images')\n            \n        else:\n            step = len(allimgs_at2) / 14.0\n            st = len(allimgs_at2)/2.0 - 6.0*step\n            end = len(allimgs_at2)+0.0001\n\n            for j, i in enumerate(np.arange(st, end, step)):\n                try:\n                    ind2 = max(0, int((i-0.5001).round()))\n                    img = self.read_dcm_ret_arr(allimgs_at2[ind2])\n                    x[..., j+28] = img.astype(np.uint8)\n                except:\n                    print(f'failed to load on {st_id}, Axial T2')\n                    pass  \n            \n            \n        if self.transform is not None:\n            x = self.transform(image=x)['image']\n\n        x = x.transpose(2, 0, 1)\n                \n        return x, str(st_id)","metadata":{"execution":{"iopub.status.busy":"2024-08-05T19:51:38.035407Z","iopub.execute_input":"2024-08-05T19:51:38.035742Z","iopub.status.idle":"2024-08-05T19:51:38.057648Z","shell.execute_reply.started":"2024-08-05T19:51:38.035717Z","shell.execute_reply":"2024-08-05T19:51:38.056682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-08-05T19:51:38.058940Z","iopub.execute_input":"2024-08-05T19:51:38.059674Z","iopub.status.idle":"2024-08-05T19:51:38.071785Z","shell.execute_reply.started":"2024-08-05T19:51:38.059639Z","shell.execute_reply":"2024-08-05T19:51:38.070890Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = RSNA24TestDataset(df, study_ids, transform=transforms_test)\ntest_dl = DataLoader(\n    test_ds, \n    batch_size=1, \n    shuffle=False)\n\nPATHS = '/kaggle/input/justright/goldilocks.pt'#[f'/kaggle/input/edgenext/edge-ultra-{num}.pt' for num in (N_FOLDS)]\nmodels = RSNA24Model('edgenext_base.in21k_ft_in1k') #for n in range(len(N_FOLDS))]\nstate_dicts = torch.load(PATHS)#[torch.load(path) for path in PATHS]\nmodels.load_state_dict(state_dicts)\nmodels.eval()\n#optuna_dicts = [OrderedDict() for f in range(N_FOLDS)]\n#for model, state_dict in zip(models, state_dicts):\n#    model.load_state_dict(state_dict)\n#    model.eval()","metadata":{"execution":{"iopub.status.busy":"2024-08-05T19:51:38.072908Z","iopub.execute_input":"2024-08-05T19:51:38.073459Z","iopub.status.idle":"2024-08-05T19:51:39.816377Z","shell.execute_reply.started":"2024-08-05T19:51:38.073428Z","shell.execute_reply":"2024-08-05T19:51:39.815369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"autocast = torch.cuda.amp.autocast(enabled=USE_AMP, dtype=torch.half)\ny_preds = []\nrow_names = []\n#models = [model.to(device) for model in models]\nmodels = [models.to(device)]\nwith tqdm(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            pred_per_study = np.zeros((25, 3))\n            \n            for cond in CONDITIONS:\n                for level in LEVELS:\n                    row_names.append(si[0] + '_' + cond + '_' + level)\n            \n            with autocast:\n                for m in models:\n                    y = m(x)[0]\n                    for col in range(N_LABELS):\n                        pred = y[col*3:col*3+3]\n                        y_pred = pred.float().softmax(0).cpu().numpy()\n                        pred_per_study[col] += y_pred / len(models)\n                y_preds.append(pred_per_study)\ny_preds = np.concatenate(y_preds, axis=0)\n    ","metadata":{"execution":{"iopub.status.busy":"2024-08-05T19:51:39.817578Z","iopub.execute_input":"2024-08-05T19:51:39.817905Z","iopub.status.idle":"2024-08-05T19:51:41.036130Z","shell.execute_reply.started":"2024-08-05T19:51:39.817879Z","shell.execute_reply":"2024-08-05T19:51:41.035197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.DataFrame()\nsub['row_id'] = row_names\nsub[LABELS] = y_preds\nsub.head(25)","metadata":{"execution":{"iopub.status.busy":"2024-08-05T19:51:41.038414Z","iopub.execute_input":"2024-08-05T19:51:41.038739Z","iopub.status.idle":"2024-08-05T19:51:41.064241Z","shell.execute_reply.started":"2024-08-05T19:51:41.038714Z","shell.execute_reply":"2024-08-05T19:51:41.063402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.to_csv('submission.csv', index=False)\npd.read_csv('submission.csv').head()","metadata":{"execution":{"iopub.status.busy":"2024-08-05T19:51:41.065149Z","iopub.execute_input":"2024-08-05T19:51:41.065399Z","iopub.status.idle":"2024-08-05T19:51:41.134516Z","shell.execute_reply.started":"2024-08-05T19:51:41.065377Z","shell.execute_reply":"2024-08-05T19:51:41.133664Z"},"trusted":true},"execution_count":null,"outputs":[]}]}