{"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":8902395,"sourceType":"datasetVersion","datasetId":5351950},{"sourceId":9072962,"sourceType":"datasetVersion","datasetId":5472877},{"sourceId":9108529,"sourceType":"datasetVersion","datasetId":5497275},{"sourceId":190549052,"sourceType":"kernelVersion"}],"dockerImageVersionId":30733,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## *Import Libraries*\n---\n","metadata":{}},{"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\nimport glob\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\n\nfrom collections import OrderedDict\n\nimport torch\nimport torch.nn.functional as F\nfrom torch import nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim import AdamW\n\nimport timm\nfrom timm.utils import ModelEmaV2\nfrom transformers import get_cosine_schedule_with_warmup\n\nimport albumentations as A\n\nfrom sklearn.model_selection import KFold\n\nimport re\nimport pydicom","metadata":{"execution":{"iopub.status.busy":"2024-08-05T08:28:58.849997Z","iopub.execute_input":"2024-08-05T08:28:58.850260Z","iopub.status.idle":"2024-08-05T08:29:07.573563Z","shell.execute_reply.started":"2024-08-05T08:28:58.850236Z","shell.execute_reply":"2024-08-05T08:29:07.572764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## *Setting & Variables*\n---\n","metadata":{}},{"cell_type":"code","source":"# OS\npath = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\n\n#OUTPUT_DIR = f'/kaggle/input/spine-degenerative-classification/pytorch/densenet201-10-fold/1'、\n#OUTPUT_DIR = f'/kaggle/input/rsna2024-training-resnet/rsna24-results'\n#OUTPUT_DIR = f'/kaggle/input/resnet101-rsna2024/rsna24-results'\n#OUTPUT_DIR = f'/kaggle/input/rsna24-resnet50'\n#resnet 34 \n#OUTPUT_DIR = f'/kaggle/input/resnet34'/kaggle/input/resnet/rsna24-results/\nOUTPUT_DIR = f'/kaggle/input/besties-rsna-notebook-training/rsna24-results'\n\ndevice = torch.device('cuda:0') if torch.cuda.is_available() else torch.device('cpu')\n#device = 'cuda:0' if torch.cuda.is_available() else 'cpu'\n\nN_WORKERS = os.cpu_count()\nUSE_AMP = True\nSEED = 8620\n\ndf = pd.read_csv(f'{path}/test_series_descriptions.csv')\nsample_sub = pd.read_csv(f'{path}/sample_submission.csv')\nstudy_ids = list(df['study_id'].unique())\nLABELS = list(sample_sub.columns[1:])\n\n# Image\nIMG_SIZE = [512, 512]\nIN_CHANS = 42 # This can change to 30\nN_LABELS = 25\nN_CLASSES = 3 * N_LABELS\n\n\n# K-fold\nN_FOLDS = 5\nBATCH_SIZE = 1\n\n\n\n# Model\n\n#MODEL_NAME = \"resnet50\"\n#MODEL_NAME = \"resnet101\"\n#MODEL_NAME = \"resnet34\"\nMODEL_NAME = \"vgg11\"\n\n# Variable\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]","metadata":{"execution":{"iopub.status.busy":"2024-08-05T08:29:07.575207Z","iopub.execute_input":"2024-08-05T08:29:07.575674Z","iopub.status.idle":"2024-08-05T08:29:07.632799Z","shell.execute_reply.started":"2024-08-05T08:29:07.575648Z","shell.execute_reply":"2024-08-05T08:29:07.632097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device","metadata":{"execution":{"iopub.status.busy":"2024-08-05T08:29:07.633807Z","iopub.execute_input":"2024-08-05T08:29:07.634077Z","iopub.status.idle":"2024-08-05T08:29:07.640445Z","shell.execute_reply.started":"2024-08-05T08:29:07.634054Z","shell.execute_reply":"2024-08-05T08:29:07.639517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## *Class for Test Dataset*\n---","metadata":{}},{"cell_type":"code","source":"# def for pimgs -> allimgs\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) ]\n\n#------------------------------------------------------------------------\n\n\nclass 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        '''get img paths of {Axial, Sagital T1, T2/STIR}'''\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'{path}/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        \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            \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            \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-05T08:29:07.643193Z","iopub.execute_input":"2024-08-05T08:29:07.643449Z","iopub.status.idle":"2024-08-05T08:29:07.665215Z","shell.execute_reply.started":"2024-08-05T08:29:07.643428Z","shell.execute_reply":"2024-08-05T08:29:07.664507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## *Class for timm Model*\n---","metadata":{}},{"cell_type":"code","source":"class RSNA24Model(nn.Module):\n    def __init__(self, model_name, in_c=42, n_classes=75, pretrained=True, features_only=False):\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    \n    def forward(self, x):\n        y = self.model(x)\n        return y","metadata":{"execution":{"iopub.status.busy":"2024-08-05T08:29:07.666370Z","iopub.execute_input":"2024-08-05T08:29:07.666687Z","iopub.status.idle":"2024-08-05T08:29:07.679256Z","shell.execute_reply.started":"2024-08-05T08:29:07.666657Z","shell.execute_reply":"2024-08-05T08:29:07.677971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## *Make test Instance & DataLoader*\n---","metadata":{}},{"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])\n\n\ntest_ds = RSNA24TestDataset(df, study_ids, transform=transforms_test)\ntest_dl = DataLoader(\n    test_ds, \n    batch_size=BATCH_SIZE, \n    shuffle=False,\n    num_workers=N_WORKERS,\n    pin_memory=True,\n    drop_last=False\n)\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-08-05T08:29:07.680545Z","iopub.execute_input":"2024-08-05T08:29:07.680874Z","iopub.status.idle":"2024-08-05T08:29:07.691226Z","shell.execute_reply.started":"2024-08-05T08:29:07.680840Z","shell.execute_reply":"2024-08-05T08:29:07.690468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## *Call Models & submit*\n---","metadata":{}},{"cell_type":"markdown","source":"### *Call best_wll_model_fold 1~10.pt*\n---","metadata":{}},{"cell_type":"code","source":"models = []\nimport glob\n#note_path = '/kaggle/input/besties-rsna-notebook-training/'\nnote_path = '/kaggle/input/vgg11-2-1/'\n\nCKPT_PATHS = glob.glob(f'{note_path}best_wll_model_fold-*.pt')\nCKPT_PATHS = sorted(CKPT_PATHS)\nCKPT_PATHS\n\n","metadata":{"execution":{"iopub.status.busy":"2024-08-05T08:29:07.692443Z","iopub.execute_input":"2024-08-05T08:29:07.692708Z","iopub.status.idle":"2024-08-05T08:29:07.706229Z","shell.execute_reply.started":"2024-08-05T08:29:07.692686Z","shell.execute_reply":"2024-08-05T08:29:07.705331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## *Run Model*\n---","metadata":{}},{"cell_type":"code","source":"\n\nfor i, cp in enumerate(CKPT_PATHS):\n    print(f'loading {cp}...')\n    model = RSNA24Model(MODEL_NAME, IN_CHANS, N_CLASSES, pretrained=False)\n    \n    \n    state_dict = torch.load(cp)\n    tmp_dict = OrderedDict()\n    for i, j in state_dict.items():   # 가중치의 모든 키 값 반복문\n        name = i.replace(\"embed_proj\",\"\")  # 매치되지 않는 키 값 변경\n        tmp_dict[name] = j\n    \n    \n    \n    \n    #model.load_state_dict(torch.load(cp))\n    model.eval()\n    model.half()\n    model.to(device)\n    models.append(model)\nautocast = torch.cuda.amp.autocast(enabled=USE_AMP, dtype=torch.half)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-08-05T08:29:07.707411Z","iopub.execute_input":"2024-08-05T08:29:07.707958Z","iopub.status.idle":"2024-08-05T08:29:19.486456Z","shell.execute_reply.started":"2024-08-05T08:29:07.707908Z","shell.execute_reply":"2024-08-05T08:29:19.485653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## *Submission*\n---","metadata":{}},{"cell_type":"code","source":"y_preds = []\nrow_names = []\n\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)\n\ny_preds = np.concatenate(y_preds, axis=0)\nsub = pd.DataFrame()\nsub['row_id'] = row_names\nsub[LABELS] = y_preds\nsub.head(25)\nsub.to_csv('submission.csv', index=False)\npd.read_csv('submission.csv').head(30)","metadata":{"execution":{"iopub.status.busy":"2024-08-05T08:29:19.487514Z","iopub.execute_input":"2024-08-05T08:29:19.487795Z","iopub.status.idle":"2024-08-05T08:29:21.909507Z","shell.execute_reply.started":"2024-08-05T08:29:19.487770Z","shell.execute_reply":"2024-08-05T08:29:21.908527Z"},"trusted":true},"execution_count":null,"outputs":[]}]}