{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9231709,"sourceType":"datasetVersion","datasetId":5583775},{"sourceId":9273089,"sourceType":"datasetVersion","datasetId":5611997}],"dockerImageVersionId":30762,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Import the Libraries","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\nfrom glob import glob\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import KFold\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\nfrom typing import Optional\nimport glob","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.516621Z","iopub.execute_input":"2024-09-07T12:24:25.517616Z","iopub.status.idle":"2024-09-07T12:24:25.525054Z","shell.execute_reply.started":"2024-09-07T12:24:25.517566Z","shell.execute_reply":"2024-09-07T12:24:25.523863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load & Edit The Data","metadata":{}},{"cell_type":"code","source":"rd = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\nOUTPUT_DIR = f'/kaggle/input/rsna2024-lsdc-training-baseline/rsna24-results'\ndevice = 'cuda:0' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.527005Z","iopub.execute_input":"2024-09-07T12:24:25.527580Z","iopub.status.idle":"2024-09-07T12:24:25.536526Z","shell.execute_reply.started":"2024-09-07T12:24:25.527535Z","shell.execute_reply":"2024-09-07T12:24:25.535583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Configuration\nEXP_NO = \"006\"\nMODEL_DIR = f\"/kaggle/input/rsna24-a-{EXP_NO}\"\nMODEL_NAME = \"tf_efficientnet_b5.ns_jft_in1k\"","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.537946Z","iopub.execute_input":"2024-09-07T12:24:25.538305Z","iopub.status.idle":"2024-09-07T12:24:25.544457Z","shell.execute_reply.started":"2024-09-07T12:24:25.538266Z","shell.execute_reply":"2024-09-07T12:24:25.543611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda:0' if torch.cuda.is_available() else 'cpu'\nN_WORKERS = os.cpu_count()\nUSE_AMP = True\nSEED = 42\nNUM_FOLDS = 5","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.545720Z","iopub.execute_input":"2024-09-07T12:24:25.546264Z","iopub.status.idle":"2024-09-07T12:24:25.552224Z","shell.execute_reply.started":"2024-09-07T12:24:25.546220Z","shell.execute_reply":"2024-09-07T12:24:25.551285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMG_SIZE = [512, 512]\nIN_CHANS = 30\nN_LABELS = 25\nN_CLASSES = 3 * N_LABELS\nBATCH_SIZE = 32","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.554172Z","iopub.execute_input":"2024-09-07T12:24:25.554547Z","iopub.status.idle":"2024-09-07T12:24:25.560482Z","shell.execute_reply.started":"2024-09-07T12:24:25.554515Z","shell.execute_reply":"2024-09-07T12:24:25.559562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rd = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\ndevice = torch.device('cuda:0') if torch.cuda.is_available() else torch.device('cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.563645Z","iopub.execute_input":"2024-09-07T12:24:25.564327Z","iopub.status.idle":"2024-09-07T12:24:25.571076Z","shell.execute_reply.started":"2024-09-07T12:24:25.564282Z","shell.execute_reply":"2024-09-07T12:24:25.570106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(f'{rd}/test_series_descriptions.csv')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.572150Z","iopub.execute_input":"2024-09-07T12:24:25.572472Z","iopub.status.idle":"2024-09-07T12:24:25.586857Z","shell.execute_reply.started":"2024-09-07T12:24:25.572425Z","shell.execute_reply":"2024-09-07T12:24:25.586045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study_ids = list(df['study_id'].unique())","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.587823Z","iopub.execute_input":"2024-09-07T12:24:25.588121Z","iopub.status.idle":"2024-09-07T12:24:25.592711Z","shell.execute_reply.started":"2024-09-07T12:24:25.588089Z","shell.execute_reply":"2024-09-07T12:24:25.591711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_sub = pd.read_csv(f'{rd}/sample_submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.593825Z","iopub.execute_input":"2024-09-07T12:24:25.594160Z","iopub.status.idle":"2024-09-07T12:24:25.601633Z","shell.execute_reply.started":"2024-09-07T12:24:25.594128Z","shell.execute_reply":"2024-09-07T12:24:25.600704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LABELS = list(sample_sub.columns[1:])\nLABELS","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.614769Z","iopub.execute_input":"2024-09-07T12:24:25.615116Z","iopub.status.idle":"2024-09-07T12:24:25.620463Z","shell.execute_reply.started":"2024-09-07T12:24:25.615084Z","shell.execute_reply":"2024-09-07T12:24:25.619601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Conditions and Levels\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-09-07T12:24:25.622148Z","iopub.execute_input":"2024-09-07T12:24:25.622482Z","iopub.status.idle":"2024-09-07T12:24:25.627544Z","shell.execute_reply.started":"2024-09-07T12:24:25.622449Z","shell.execute_reply":"2024-09-07T12:24:25.626590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Helper functions\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) ]\nprint(\"DONE\")","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.628620Z","iopub.execute_input":"2024-09-07T12:24:25.628943Z","iopub.status.idle":"2024-09-07T12:24:25.636633Z","shell.execute_reply.started":"2024-09-07T12:24:25.628908Z","shell.execute_reply":"2024-09-07T12:24:25.635646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Feature Engineering","metadata":{}},{"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) / 10.0\n            st = len(allimgs_st1)/2.0 - 4.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) / 10.0\n            st = len(allimgs_st2)/2.0 - 4.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+10] = 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) / 10.0\n            st = len(allimgs_at2)/2.0 - 4.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+20] = 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-09-07T12:24:25.638001Z","iopub.execute_input":"2024-09-07T12:24:25.638334Z","iopub.status.idle":"2024-09-07T12:24:25.658893Z","shell.execute_reply.started":"2024-09-07T12:24:25.638303Z","shell.execute_reply":"2024-09-07T12:24:25.658091Z"},"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-09-07T12:24:25.661158Z","iopub.execute_input":"2024-09-07T12:24:25.661444Z","iopub.status.idle":"2024-09-07T12:24:25.669561Z","shell.execute_reply.started":"2024-09-07T12:24:25.661414Z","shell.execute_reply":"2024-09-07T12:24:25.668580Z"},"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=32, \n    shuffle=False,\n    num_workers=N_WORKERS,\n    pin_memory=True,\n    drop_last=False\n)","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.670808Z","iopub.execute_input":"2024-09-07T12:24:25.671588Z","iopub.status.idle":"2024-09-07T12:24:25.678281Z","shell.execute_reply.started":"2024-09-07T12:24:25.671545Z","shell.execute_reply":"2024-09-07T12:24:25.677345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Build The Model","metadata":{}},{"cell_type":"code","source":"class RSNA24Model(nn.Module):\n    def __init__(\n        self,\n        model_name: str,\n        pretrained: bool,\n        features_only: bool,\n        in_chans: int,\n        n_classes: int,\n        n_labels: int,\n        loss_name: str,\n    ):\n        super().__init__()\n        self.model = timm.create_model(\n            model_name=model_name,\n            pretrained=pretrained, \n            features_only=features_only,\n            in_chans=in_chans,\n            num_classes=n_classes,\n            global_pool='avg'\n        )\n        self.loss_fn = loss_name\n        self.n_labels = n_labels\n    \n    def forward(\n        self,\n        x: torch.Tensor,\n        y: Optional[torch.Tensor],\n    ) -> dict[str, torch.Tensor]:\n        \n        logits = self.model(x)\n        \n        output = {\"logits\": logits}\n        if y is not None:\n            loss = 0\n            for col in range(self.n_labels):\n                pred = logits[:,col*3:col*3+3]\n                gt = y[:,col]\n                loss = loss + self.loss_fn(pred, gt) / self.n_labels\n            output[\"loss\"] = loss\n        \n        return output","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.679571Z","iopub.execute_input":"2024-09-07T12:24:25.680138Z","iopub.status.idle":"2024-09-07T12:24:25.688498Z","shell.execute_reply.started":"2024-09-07T12:24:25.680105Z","shell.execute_reply":"2024-09-07T12:24:25.687763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = []","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.689456Z","iopub.execute_input":"2024-09-07T12:24:25.689727Z","iopub.status.idle":"2024-09-07T12:24:25.710743Z","shell.execute_reply.started":"2024-09-07T12:24:25.689697Z","shell.execute_reply":"2024-09-07T12:24:25.709897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in range(NUM_FOLDS):\n    cp = f\"{MODEL_DIR}/{EXP_NO}-{i}/best_model.pth\"\n#     cp = f\"{MODEL_DIR}/best_model.pth\"\n    print(f'loading {cp}...')\n    model = RSNA24Model(MODEL_NAME, False, False, IN_CHANS, N_CLASSES, 25, 'dummy')\n    state_dict = torch.load(cp)\n\n    # 予期しないキーを削除\n    if 'loss_fn.weight' in state_dict:\n        del state_dict['loss_fn.weight']\n    model.load_state_dict(state_dict)\n    model.eval()\n    model.to(device)\n    models.append(model)","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:25.712475Z","iopub.execute_input":"2024-09-07T12:24:25.712828Z","iopub.status.idle":"2024-09-07T12:24:29.648294Z","shell.execute_reply.started":"2024-09-07T12:24:25.712786Z","shell.execute_reply":"2024-09-07T12:24:29.647520Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\nimport torch\nimport numpy as np\n\nautocast = torch.cuda.amp.autocast(enabled=USE_AMP, dtype=torch.half)\ny_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            try:\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, None)[\"logits\"][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            except Exception as e:\n                print(f\"Error processing index {idx}: {e}\")\n                # Optionally, append a placeholder or skip this sample\n                y_preds.append(np.zeros((25, 3)))\n\ny_preds = np.concatenate(y_preds, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-09-07T12:24:29.649471Z","iopub.execute_input":"2024-09-07T12:24:29.649779Z","iopub.status.idle":"2024-09-07T12:24:30.880337Z","shell.execute_reply.started":"2024-09-07T12:24:29.649745Z","shell.execute_reply":"2024-09-07T12:24:30.879250Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Release the output","metadata":{}},{"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-09-07T12:24:30.881866Z","iopub.execute_input":"2024-09-07T12:24:30.882236Z","iopub.status.idle":"2024-09-07T12:24:30.902552Z","shell.execute_reply.started":"2024-09-07T12:24:30.882198Z","shell.execute_reply":"2024-09-07T12:24:30.901479Z"},"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-09-07T12:24:30.903795Z","iopub.execute_input":"2024-09-07T12:24:30.904126Z","iopub.status.idle":"2024-09-07T12:24:30.919468Z","shell.execute_reply.started":"2024-09-07T12:24:30.904079Z","shell.execute_reply":"2024-09-07T12:24:30.918618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}