{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","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":9729124,"sourceType":"datasetVersion","datasetId":5953414}],"dockerImageVersionId":31040,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Import and Config","metadata":{}},{"cell_type":"code","source":"import gc\nimport wandb\nfrom pytorch_lightning.loggers import WandbLogger\nimport os\nimport yaml\nimport sys\nimport cv2\nimport random\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport torch\nfrom glob import glob\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom torch.optim import AdamW, Adam\nimport torch.nn as nn\nimport pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping, TQDMProgressBar\nimport torchvision.transforms as T\nimport albumentations as A\nimport pandas.api.types\nimport sklearn.metrics\nimport timm\nimport scipy\nimport albumentations as A\nfrom torchvision.transforms import v2\nfrom torchvision import models\nfrom tqdm.auto import tqdm\nfrom joblib import Parallel, delayed\nfrom torch.utils.data import default_collate\nimport pydicom as dcm\nimport transformers","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:08.807985Z","iopub.execute_input":"2025-06-07T01:05:08.808284Z","iopub.status.idle":"2025-06-07T01:05:08.814219Z","shell.execute_reply.started":"2025-06-07T01:05:08.808264Z","shell.execute_reply":"2025-06-07T01:05:08.813350Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.listdir('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/44036939/2828203845')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:08.815345Z","iopub.execute_input":"2025-06-07T01:05:08.816008Z","iopub.status.idle":"2025-06-07T01:05:08.835951Z","shell.execute_reply.started":"2025-06-07T01:05:08.815989Z","shell.execute_reply":"2025-06-07T01:05:08.835296Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Anatomy & image Visual for more\nhttps://www.kaggle.com/code/abhinavsuri/anatomy-image-visualization-overview-rsna-raids#RSNA-Lumbar-Spine-Challenge.\n","metadata":{}},{"cell_type":"code","source":"# Đường dẫn đến folder chứa ảnh\nfolder = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/44036939/2828203845'\n# Liệt kê các file DICOM trong thư mục\nfiles = [f for f in os.listdir(folder) if f.endswith('.dcm')]\n\n# Đọc một file (ví dụ: ảnh đầu tiên)\ndicom_path = os.path.join(folder, files[0])\ndemo_dcm = dcm.dcmread(dicom_path)\n\n# Chuyển pixel data sang numpy array\nimg = demo_dcm.pixel_array\n\n# Hiển thị ảnh\nplt.imshow(img, cmap='gray')\nplt.title(f\"File: {files[0]}\")\nplt.axis('off')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:08.836614Z","iopub.execute_input":"2025-06-07T01:05:08.836840Z","iopub.status.idle":"2025-06-07T01:05:08.992166Z","shell.execute_reply.started":"2025-06-07T01:05:08.836824Z","shell.execute_reply":"2025-06-07T01:05:08.991267Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 1710 # My birth day\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True # Fix the network according to random seed\n    print('Finish seeding with seed {}'.format(seed))\n\nseed_everything(SEED)\nprint('Training on device {}'.format(device))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:08.993862Z","iopub.execute_input":"2025-06-07T01:05:08.994054Z","iopub.status.idle":"2025-06-07T01:05:09.003872Z","shell.execute_reply.started":"2025-06-07T01:05:08.994038Z","shell.execute_reply":"2025-06-07T01:05:09.003256Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile config.yaml \ndata_path : '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\nout_put_dir : '/kaggle/working/models'\n\nseed : 1101\ndebug : False\ntrain_bs : 4\nvalid_bs : 4\ntest_bs : 8\nworker : 1\n\nprogress_bar_refresh_rate : 1\n\npseudo_train : 0\n\nsave_topk : 1\nfold : 5 # Cross Validation \n\ntask:\n    kind: 'detect' # coordinate x, y\n    #kind : 'classify' # severity -> label\n    #kind : 'depth'\n    condition: 'nfn'\n    #condition: 'scs'\n    #condition: 'scs'\n    #condition: 'all'\n    #direction: 'satg2'\n    direction: 'ax'\n    #direction: 'sagt1'\n    position:\n        - 'L1/L2'\n        - 'L2/L3'\n        - 'L3/L4'\n        - 'L4/L5'\n        - 'L5/S1'\n\nin_chans : 3\n\nimage_size : 384\n\nmodel:","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:09.004599Z","iopub.execute_input":"2025-06-07T01:05:09.004807Z","iopub.status.idle":"2025-06-07T01:05:09.014779Z","shell.execute_reply.started":"2025-06-07T01:05:09.004791Z","shell.execute_reply":"2025-06-07T01:05:09.014205Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os \nos.listdir('/kaggle/input/')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:09.015455Z","iopub.execute_input":"2025-06-07T01:05:09.015661Z","iopub.status.idle":"2025-06-07T01:05:09.025268Z","shell.execute_reply.started":"2025-06-07T01:05:09.015641Z","shell.execute_reply":"2025-06-07T01:05:09.024757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open(\"config.yaml\", \"r\") as file_obj:\n    config = yaml.safe_load(file_obj)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:09.025948Z","iopub.execute_input":"2025-06-07T01:05:09.026181Z","iopub.status.idle":"2025-06-07T01:05:09.037672Z","shell.execute_reply.started":"2025-06-07T01:05:09.026167Z","shell.execute_reply":"2025-06-07T01:05:09.037078Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Sử dụng train data trong mở debug mode\nif config['debug']:\n    IMAGE_PATH = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/'\n    series = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv')\nelse:\n    IMAGE_PATH = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/'\n    series = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:09.039841Z","iopub.execute_input":"2025-06-07T01:05:09.040051Z","iopub.status.idle":"2025-06-07T01:05:09.061716Z","shell.execute_reply.started":"2025-06-07T01:05:09.040032Z","shell.execute_reply":"2025-06-07T01:05:09.060985Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(IMAGE_PATH)\nseries.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:09.062389Z","iopub.execute_input":"2025-06-07T01:05:09.062640Z","iopub.status.idle":"2025-06-07T01:05:09.081477Z","shell.execute_reply.started":"2025-06-07T01:05:09.062623Z","shell.execute_reply":"2025-06-07T01:05:09.080819Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Stage 0: intial (create meta file)\n* tạo metafile từ meta data của ảnh dicom: không chỉ chứa ảnh mà nhiều meta data khác\n* tạo một data frame - file chứa metadata quan trọng từ một series ảnh thuộc một study_id (phiên chụp bệnh nhân)","metadata":{}},{"cell_type":"code","source":"def create_dcm_df(study_id, series_id, series_description):\n    try:\n        # Find all .dcm to extract Instance_number từ tên file 1, 2, 3,...  /* để tìm all file dcm\n        path_list = glob(IMAGE_PATH + f'{study_id}/{series_id}/*.dcm')\n        in_list = sorted([int(s.split('/')[-1].split('.')[0]) for s in path_list])\n\n        # read file dicom by pydicom\n        dcm_list = []\n        for i in in_list:\n            dcm_list.append(dcm.dcmread(IMAGE_PATH + f'{study_id}/{series_id}/{i}.dcm'))\n            \n        '''\n        Extract metadata chính:\n        ipp : vị trí x,y,z của ảnh trong cơ thể\n        ioo : hướng trục của ảnh trong không gian\n        '''\n        ipp = np.asarray([d.ImagePositionPatient for d in dcm_list]).astype('float')\n        iop = [d.ImageOrientationPatient for d in dcm_list]\n        iop = [[float(d[0]), float(d[1]), float(d[2]), \n                float(d[3]), float(d[4]), float(d[5])] for d in iop]\n        ipp_x = ipp[:, 0]\n        ipp_y = ipp[:, 1]\n        ipp_z = ipp[:, 2]\n\n        # Extract độ phân giải (Pixel Spacing) và kích thước\n        shape = np.array([d.pixel_array.shape for d in dcm_list])\n        sbs = np.asarray([d.SpacingBetweenSlices for d in dcm_list]).astype('float')\n        ps = np.asarray([d.PixelSpacing for d in dcm_list]).astype('float') # x, y\n        ps_x = ps[:, 0]\n        ps_y = ps[:, 1]\n\n        # create dataframe chứa metadata\n        meta_dict = {\n            'instance_number' : in_list,\n            'ipp_x' : ipp_x,\n            'ipp_y' : ipp_y,\n            'ipp_z' : ipp_z,\n            'sbs' : sbs,\n            'ps_x' : ps_x, \n            'ps_y' : ps_y\n        }\n        meta_df = pd.DataFrame(meta_dict)\n\n        #add other fearture\n        meta_df['series_id'] = series_id\n        meta_df['study_id'] = study_id\n        meta_df['series_description'] = series_description\n        meta_df['height'] = shape[:, 0]\n        meta_df['width'] = shape[:, 1]\n        meta_df['iop'] = pd.Series(iop)\n        \n        # xoá dữ liệu giải phóng bộ nhớ\n        del dcm_list, ipp, iop, sbs, ps\n        gc.collect()\n        return meta_df[['study_id', 'series_id', 'series_description', 'instance_number', 'height', 'width', 'ipp_x', 'ipp_y', 'ipp_z', 'iop', 'sbs', 'ps_x', 'ps_y']]\n    except:\n        print(study_id, series_id, series_description)\n        return None\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:09.082396Z","iopub.execute_input":"2025-06-07T01:05:09.082642Z","iopub.status.idle":"2025-06-07T01:05:09.091255Z","shell.execute_reply.started":"2025-06-07T01:05:09.082623Z","shell.execute_reply":"2025-06-07T01:05:09.090625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"create_dcm_df(44036939, 2828203845, 'Sagittal T1').head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:09.092136Z","iopub.execute_input":"2025-06-07T01:05:09.092385Z","iopub.status.idle":"2025-06-07T01:05:09.904230Z","shell.execute_reply.started":"2025-06-07T01:05:09.092363Z","shell.execute_reply":"2025-06-07T01:05:09.903469Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Turn into .parquet file**","metadata":{}},{"cell_type":"code","source":"%%time\nif config['debug']:\n    meta_df_list = []\n    \n    # Parallbel song song các \n    meta_df_list = Parallel(n_jobs=-1)([delayed(create_dcm_df)(row.study_id, row.series_id, row.series_description) for _, row in series.iterrows()])\n\n    # merge to big dataframe\n    meta_df = pd.concat(meta_df_list)\n    del meta_df_list\n    gc.collect()\n\n    #turn to file parquet\n    meta_df.to_parquet('meta.parquet')\nelse:\n    meta_df_list = []\n    meta_df_list = Parallel(n_job = -1)([delayed(create_dcm_df)(row.study_id, row.series_id, row.series_description) for _, row in series.iterrows()])\n\n    meta_df = pd.concat(meta_df_list)\n    del meta_df_list\n    gc.collect()\n    meta_df.to_parquet('meta.parquet')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:09.904878Z","iopub.execute_input":"2025-06-07T01:05:09.905128Z","iopub.status.idle":"2025-06-07T01:05:12.535926Z","shell.execute_reply.started":"2025-06-07T01:05:09.905106Z","shell.execute_reply":"2025-06-07T01:05:12.535295Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Kiểm tra file parquet: file lớn thích hợp để lưu với nhiều dữ liệu hơn csv\nimport os \nos.listdir('/kaggle/working/')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:12.536779Z","iopub.execute_input":"2025-06-07T01:05:12.537165Z","iopub.status.idle":"2025-06-07T01:05:12.541866Z","shell.execute_reply.started":"2025-06-07T01:05:12.537138Z","shell.execute_reply":"2025-06-07T01:05:12.541263Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# First stage (depth inference)\n* infer - suy luận depth từ ảnh sagittal T1 và T2\n* căn chỉnh độ sâu suy luận theo luật thuật toán - align infered depth using rule base algorithims","metadata":{}},{"cell_type":"markdown","source":"## Custom dataset\n* Hẹp lỗ liên hợp trái (left nfn)\n* Hẹp lỗ liên hợp phải (right nfn)\n* Hẹp dưới khớp trái (left ss)\n* Hẹp dưới khớp phải (right ss)\n* Hẹp ống sống (scs)\n","metadata":{}},{"cell_type":"markdown","source":"condition: xác định loại ảnh MRI nào dùng, gồm:\n\n'scs': Sagittal T2/STIR\n\n'nfn': Sagittal T1","metadata":{}},{"cell_type":"code","source":"# # Bước đầu của pipeline \n# class DepthDetectDataset(Dataset):\n#     def __init__(self, meta, condition, usage = 'sub'):\n#         if condition == 'scs':\n#             meta = meta.loc[meta.series_description == 'Sagittal T2/STIR']\n#         else:\n#             meta = meta.loc[meta.series_description == 'Sagittal T1']\n#         self.id = list(meta.study_id.unique())\n#         if 3637444890 in self.id:\n#             self.id.remove(3637444890)\n#         self.meta = meta\n#         self.condition = condition\n#         self.usage = usage\n#         self.resize = v2.Resize(384, 384)\n\n#     def for_scs(self, study_id):\n#         depth = 32\n#         #lọc dữ liệu với 1 study_id\n#         meta = self.meta.loc[(self.meta.study_id == study_id) & (self.meta.series_decription == 'Sagittal T2/STIR')]\n        \n#         #sắp xếp trục theo không gian x\n#         meta = meta.sort_values('ipp_x', ascending = True).reset_index(drop = True)\n        \n#         #đọc ảnh dicom thành mảng 2d numpy\n#         img = [self.load_dicom(IMAGE + f'{row.study_id}/{row.series_id}/{row.instance_number}.dcm') for _, row in meta.iterrows()]\n        \n#         #resize, chuyển ảnh thành tensor rồi stack lại\n#         volume = self.normalize(torch.cat([self.resize(torch.tensor(i.astype(np.float32))[None, ...]).to(torch.float32) for i in img]).contiguous())\n        \n#         #chuẩn hoá độ sâu depth\n#         if volume.shape[0] < depth:\n#             volume = torch.cat([volume, torch.zeros(depth-volume.shape[0], volume.shape[1], volume.shape[2])])\n#         elif volume.shape[0] > depth:\n#             volume = torch.nn.fuctional.interpolate(volume[None, None, ...], (depth, volume.shape[1], volume.shape[2])).squeeze()\n#         return volume.to(torch.float32)\n            \n#     def for_nfn(self, study_id):\n#         depth = 32\n#         #lọc dữ liệu với 1 study_id\n#         meta = self.meta.loc[(self.meta.study_id == study_id) & (self.meta.series_decription == 'Sagittal T1')]\n        \n#         #sắp xếp trục theo không gian x\n#         meta = meta.sort_values('ipp_x', ascending = True).reset_index(drop = True)\n        \n#         #đọc ảnh dicom thành mảng 2d numpy\n#         img = [self.load_dicom(IMAGE + f'{row.study_id}/{row.series_id}/{row.instance_number}.dcm') for _, row in meta.iterrows()]\n        \n#         #resize, chuyển ảnh thành tensor rồi stack lại\n#         volume = self.normalize(torch.cat([self.resize(torch.tensor(i.astype(np.float32))[None, ...]).to(torch.float32) for i in img]).contiguous())\n        \n#         #chuẩn hoá độ sâu depth\n#         if volume.shape[0] < depth:\n#             volume = torch.cat([volume, torch.zeros(depth-volume.shape[0], volume.shape[1], volume.shape[2])])\n#         elif volume.shape[0] > depth:\n#             volume = torch.nn.fuctional.interpolate(volume[None, None, ...], (depth, volume.shape[1], volume.shape[2])).squeeze()\n#         return volume.to(torch.float32)\n    \n#     def __getitem__(self, index):\n#         study_id = self.id[index]\n\n#         if self.condition == 'scs':\n#             volume = self.for_scs(study_id)\n#         elif self.condition == 'nfn':\n#             try:\n#                 volume = self.for_nfn(study_id)\n#             except:\n#                 print(study_id)\n#         return volume, torch.tensor([study_id])\n\n#     def normalize(self, x):\n#         upper = torch.quantile(x, torch.tensor([0.99]))\n#         lower = torch.quantile(x, torch.tensor([0.01]))\n#         x = torch.clip(x, lower, upper)\n#         x = x - torch.min(x)\n#         x = x / (torch.max(x)+1e-6)\n\n#     def __len__(self):\n#         return len(self.id)\n\n#     def load_dcm(self, path):\n#         dicom = dcm.read_file(path)\n#         data = dicom.pixel_array\n#         return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:12.542648Z","iopub.execute_input":"2025-06-07T01:05:12.543430Z","iopub.status.idle":"2025-06-07T01:05:12.553024Z","shell.execute_reply.started":"2025-06-07T01:05:12.543409Z","shell.execute_reply":"2025-06-07T01:05:12.552357Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DepthDetectDataset(Dataset):\n    def __init__(self, meta, condition, usage='sub'):\n        if condition == 'scs': \n            meta = meta.loc[meta.series_description=='Sagittal T2/STIR']\n        else: \n            meta = meta.loc[meta.series_description=='Sagittal T1']\n        self.id = list(meta.study_id.unique())\n        if 3637444890 in self.id: \n            self.id.remove(3637444890)\n        self.meta = meta\n        self.condition = condition\n        self.usage = usage\n        \n        self.resize = v2.Resize((384, 384))\n        \n    def __getitem__(self, index):\n        study_id = self.id[index]\n        #print(study_id)\n        #try:\n        if self.condition == 'scs':\n            volume = self.for_scs(study_id)\n        elif self.condition == 'nfn':\n            try: \n                volume = self.for_nfn(study_id)\n            except: \n                print(study_id)\n        return volume, torch.tensor([study_id])\n\n    def for_scs(self, study_id):\n        depth = 32\n        meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Sagittal T2/STIR')]\n        meta = meta.sort_values('ipp_x', ascending=True).reset_index(drop=True)\n        img = [self.load_dicom(IMAGE_PATH + f'{row.study_id}/{row.series_id}/{row.instance_number}.dcm') for _, row in meta.iterrows()]\n        volume = self.normalize(torch.cat([self.resize(torch.tensor(i.astype(np.float32))[None, ...]).to(torch.float32) for i in img]).contiguous())\n        if volume.shape[0] < depth:\n            volume = torch.cat([volume, torch.zeros(depth-volume.shape[0], volume.shape[1], volume.shape[2])])\n        elif volume.shape[0] > depth:\n            volume = torch.nn.functional.interpolate(volume[None, None, ...], (depth, volume.shape[1], volume.shape[2])).squeeze()\n        return volume.to(torch.float32)\n\n    def for_nfn(self, study_id):\n        depth = 32\n        meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Sagittal T1')]\n        meta = meta.sort_values('ipp_x', ascending=True).reset_index(drop=True)\n        img = [self.load_dicom(IMAGE_PATH + f'{row.study_id}/{row.series_id}/{row.instance_number}.dcm') for _, row in meta.iterrows()]\n        volume = self.normalize(torch.cat([self.resize(torch.tensor(i.astype(np.float32))[None, ...]).to(torch.float32) for i in img]).contiguous())\n        if volume.shape[0] < depth:\n            volume = torch.cat([volume, torch.zeros(depth-volume.shape[0], volume.shape[1], volume.shape[2])])\n        elif volume.shape[0] > depth:\n            volume = torch.nn.functional.interpolate(volume[None, None, ...], (depth, volume.shape[1], volume.shape[2])).squeeze()\n        return volume.to(torch.float32)\n    \n    def normalize(self, x):\n        upper = torch.quantile(x, torch.tensor([0.99]))\n        lower = torch.quantile(x, torch.tensor([0.01]))\n        x = torch.clip(x, lower, upper)\n        x = x - torch.min(x)\n        x = x / (torch.max(x)+1e-6)\n        return x\n\n    def __len__(self):\n        return len(self.id)\n\n    def load_dicom(self, path):\n        dicom = dcm.dcmread(path)\n        data = dicom.pixel_array\n        return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:12.555748Z","iopub.execute_input":"2025-06-07T01:05:12.555935Z","iopub.status.idle":"2025-06-07T01:05:12.573133Z","shell.execute_reply.started":"2025-06-07T01:05:12.555920Z","shell.execute_reply":"2025-06-07T01:05:12.572416Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"demo_path = '/kaggle/working/meta.parquet'\nmeta_df = pd.read_parquet(demo_path)\ndataset_demo = DepthDetectDataset(meta=meta_df, condition = 'nfn')\nprint(len(dataset_demo))\ndataset_demo.__getitem__(0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:12.573863Z","iopub.execute_input":"2025-06-07T01:05:12.574091Z","iopub.status.idle":"2025-06-07T01:05:13.920033Z","shell.execute_reply.started":"2025-06-07T01:05:12.574069Z","shell.execute_reply":"2025-06-07T01:05:13.919283Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model\nsử dụng 3D ConvNext","metadata":{}},{"cell_type":"code","source":"from torchvision.ops import StochasticDepth \nfrom typing import List, Dict\nfrom torch import Tensor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:13.920828Z","iopub.execute_input":"2025-06-07T01:05:13.921116Z","iopub.status.idle":"2025-06-07T01:05:13.924911Z","shell.execute_reply.started":"2025-06-07T01:05:13.921092Z","shell.execute_reply":"2025-06-07T01:05:13.924218Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Stem block: để sử lý đầu vào ảnh 3D và chuẩn hoá\n# 3D Conv - dùng kernel_size [D, H, W] - (1, 2, 2) - giữ nguyên depth, chỉ tích chập trên H, W. Tương tự với stride\n# Output h/2, w/2 với stride = 2\nclass ConvNextStem(nn.Sequential):\n    def __init__(self, in_features: int, out_features : int):\n        super().__init__(\n            nn.Conv3d(in_features, out_features, kernel_size = (1, 2, 2), stride = (1, 2, 2)),\n            nn.GroupNorm(num_groups = 1, num_channels = out_features)\n        )\n        \n# Tránh vanish: nhân hệ số x vơi gamma (learnable scaling) học được\nclass LayerScaler(nn.Module):\n    def __init__(self, init_value : float, dimensions : int):\n        super.__init__()\n        self.gamma = nn.Parameter(init_value * torch.ones((dimensions)), requires_grad = True)\n\n    def forward(self, x):\n        return self.gamma[None, ..., None, None]* x\n\n# Residual bottleneck block trong các kiến trúc hiện đại: \n# Học được nhiều đặc trưng hơn thông qua: expansion, ...\nclass BottleNeckBlock(nn.Module):\n    def __init__(\n        self, \n        in_features : int,\n        out_features: int,\n        expansion : int = 4,\n        drop_p : float = .0,\n        layer_scaler_init_value : float = 1e-6,\n    ):\n        super().__init__()\n        expanded_features = out_features * expansion \n        \n        # 3 lần Conv để lấy được nhiều feature đặc trưng hơn\n        self.block = nn.Sequential(\n            nn.Conv3d(\n                in_features, in_features, kernel_size = (2, 7, 7), padding = 'same', bias = False, groups = in_features\n            ),\n\n            # chuẩn hoá theo channel\n            nn.GroupNorm(num_groups = in_features, num_channels = in_features),\n            nn.Conv3d(in_features, expanded_features, kernel_size = 1),\n            nn.GELU(),\n            nn.Conv3d(expanded_features, out_features, kernel_size = 1),\n        )\n\n    def forward(self, x : Tensor) -> Tensor:\n        res = x\n        x = self.block(x)\n        x += res\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:13.925850Z","iopub.execute_input":"2025-06-07T01:05:13.926092Z","iopub.status.idle":"2025-06-07T01:05:13.938569Z","shell.execute_reply.started":"2025-06-07T01:05:13.926072Z","shell.execute_reply":"2025-06-07T01:05:13.937928Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# tăng channel không đồng nghĩa với học được nhiều feature hữu ích\nclass ConvNexStage(nn.Sequential):\n    def __init__(\n        self, in_features: int, out_features: int, depth: int, **kwargs\n    ):\n        super().__init__(\n            # Downsampler giảm độ phân giải không gian (D, H, W), tăng số lượng channel.\n            nn.Sequential(\n                nn.GroupNorm(num_groups=in_features, num_channels=in_features),\n                nn.Conv3d(in_features, out_features, kernel_size=(2, 2, 2), stride=(2, 2, 2))\n            ),\n            # BottleNeckBlock: áp dụng nhiều khối tích chập (Conv3D) nâng cao biểu diễn đặc trưng.\n            *[\n                BottleNeckBlock(out_features, out_features, **kwargs)\n                for _ in range(depth)\n            ],\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:13.939203Z","iopub.execute_input":"2025-06-07T01:05:13.939360Z","iopub.status.idle":"2025-06-07T01:05:13.952977Z","shell.execute_reply.started":"2025-06-07T01:05:13.939348Z","shell.execute_reply":"2025-06-07T01:05:13.952368Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ConvNextEncoder(nn.Module):\n    def __init__(\n        self,\n        in_channels: int,\n        stem_features: int,\n        depths: List[int],\n        widths: List[int],\n        drop_p: float = .0, # regularization\n    ):\n        super().__init__()\n        # giảm khích thước ảnh đầu vào (spatial downsampling), tăng số lượng features - channel\n        self.stem = ConvNextStem(in_channels, stem_features)\n\n        in_out_widths = list(zip(widths, widths[1:]))\n        # create drop paths probabilities (one for each stage)\n        drop_probs = [x.item() for x in torch.linspace(0, drop_p, sum(depths))]\n\n        # sử dụng các NexStage để trích xuất được feature trừu tượng\n        self.stages = nn.ModuleList(\n            [\n                ConvNexStage(stem_features, widths[0], depths[0], drop_p=drop_probs[0]),\n                *[\n                    ConvNexStage(in_features, out_features, depth, drop_p=drop_p)\n                    for (in_features, out_features), depth, drop_p in zip(\n                        in_out_widths, depths[1:], drop_probs[1:]\n                    )\n                ],\n            ]\n        )\n\n\n    def forward(self, x):\n        # qua stem để tiền sử lý\n        x = self.stem(x)\n\n        #rồi đi qua từng stage để extract\n        for stage in self.stages:\n            x = stage(x)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:13.953642Z","iopub.execute_input":"2025-06-07T01:05:13.953801Z","iopub.status.idle":"2025-06-07T01:05:13.965990Z","shell.execute_reply.started":"2025-06-07T01:05:13.953788Z","shell.execute_reply":"2025-06-07T01:05:13.965264Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ClassificationHead(nn.Sequential):\n    def __init__(self):\n        super().__init__(\n            nn.AdaptiveAvgPool3d((1, 1, 1)), # (B, C, 1, 1, 1)\n            nn.Flatten(1), # (B, C)\n            nn.LayerNorm(512),\n            nn.Linear(512, 3)\n        )\nclass Flatten(nn.Sequential):\n    def __init__(self):\n        super().__init__(\n            nn.AdaptiveAvgPool3d((1, 1, 1)),\n            nn.Flatten(1),\n            nn.LayerNorm(512)\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:13.966744Z","iopub.execute_input":"2025-06-07T01:05:13.966953Z","iopub.status.idle":"2025-06-07T01:05:13.981063Z","shell.execute_reply.started":"2025-06-07T01:05:13.966930Z","shell.execute_reply":"2025-06-07T01:05:13.980263Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# #test phase\n# x = torch.randn(4, 512, 32, 384, 384)\n# test = ClassificationHead()\n# print(test(x).shape) \n# run to long must to gpu\nx = torch.randn(4, 512, 8, 8, 8)\ntest = Flatten()\nprint(test(x).shape) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:13.981714Z","iopub.execute_input":"2025-06-07T01:05:13.981901Z","iopub.status.idle":"2025-06-07T01:05:14.008679Z","shell.execute_reply.started":"2025-06-07T01:05:13.981887Z","shell.execute_reply":"2025-06-07T01:05:14.008015Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# dự đoán độ sâu của đốt sống bên trái và bên phải\nclass ConvNextSSDepthDetect(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # Extract feature từ 3D ConvNextEncoder \n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=32, depths=[3,3,9,3], widths=[64, 128, 256, 512])\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(512))\n        # Output là 10 branches: mỗi branch là 1 đoạn đốt sống\n        self.ll1 = nn.Linear(512, 96)\n        self.ll2 = nn.Linear(512, 96)\n        self.ll3 = nn.Linear(512, 96)\n        self.ll4 = nn.Linear(512, 96)\n        self.ll5 = nn.Linear(512, 96)\n        self.rl1 = nn.Linear(512, 96)\n        self.rl2 = nn.Linear(512, 96)\n        self.rl3 = nn.Linear(512, 96)\n        self.rl4 = nn.Linear(512, 96)\n        self.rl5 = nn.Linear(512, 96)\n        # trả về 10 keys: mỗi key là tensor [B, 96]\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        ll1 = self.ll1(x)\n        ll2 = self.ll2(x)\n        ll3 = self.ll3(x)\n        ll4 = self.ll4(x)\n        ll5 = self.ll5(x)\n        rl1 = self.rl1(x)\n        rl2 = self.rl2(x)\n        rl3 = self.rl3(x)\n        rl4 = self.rl4(x)\n        rl5 = self.rl5(x)\n        # [B, 96]\n        return {'left_L1/L2': ll1,'left_L2/L3': ll2,'left_L3/L4': ll3, 'left_L4/L5': ll4, 'left_L5/S1': ll5,\n                'right_L1/L2': rl1, 'right_L2/L3': rl2, 'right_L3/L4': rl3, 'right_L4/L5': rl4, 'right_L5/S1': rl5}\n\n\n# Bổ sung thông tin về vị trí trong chuỗi vào embedding vector\nclass PositionalEncoding(nn.Module):\n    \n    def __init__(self, d_model: int, dropout: float = 0.1, max_len: int = 5000):\n        super().__init__()\n        self.dropout = nn.Dropout(p=dropout)\n        \n        # tạo vector position từ 0 - max_len-1\n        position = torch.arange(max_len).unsqueeze(1)\n        \n        # tấn suất sin/cos \n        div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))\n        \n        # Ánh xạ vào sin cos\n        pe = torch.zeros(max_len, 1, d_model)\n        pe[:, 0, 0::2] = torch.sin(position * div_term)\n        pe[:, 0, 1::2] = torch.cos(position * div_term)\n        \n        # lưu pe và buffer \n        self.register_buffer('pe', pe)\n\n    def forward(self, x: Tensor) -> Tensor:\n        \"\"\"\n        Args:\n            x: Tensor, shape [batch_size, seq_len, embedding_dim]\n        \"\"\"\n        # Thêm encoding vào x. Sau đó chuyển lại về [B, T, D]\n        x = x.permute(1, 0, 2)\n        x = x + self.pe[:x.size(0)]\n        return self.dropout(x.permute(1, 0, 2))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:14.009351Z","iopub.execute_input":"2025-06-07T01:05:14.009598Z","iopub.status.idle":"2025-06-07T01:05:14.020382Z","shell.execute_reply.started":"2025-06-07T01:05:14.009577Z","shell.execute_reply":"2025-06-07T01:05:14.019681Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## pipeline 3 use pretrain model \n* Ở AttentionSSDepthDetect có thể bỏ unsqueeze trong forward","metadata":{}},{"cell_type":"code","source":"# pipeline 3 usepretrain model \nclass AttentionSSDepthDetect(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # load pretrained model ConvNeXt-base\n        self.encoder = timm.create_model('convnext_base.fb_in22k_ft_in1k_384', in_chans=3, pretrained=True, num_classes=0)\n        \n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(512))\n        \n        self.in_features = self.encoder.num_features\n\n        # tạo position encoding dạng sin cos\n        self.rpe = PositionalEncoding(self.in_features, dropout=0., max_len=64)\n        \n        self.transformer0 = nn.TransformerEncoderLayer(d_model=self.in_features, nhead=8, activation='gelu',dropout=0.1, batch_first=True)\n        self.transformer1 = nn.TransformerEncoderLayer(d_model=self.in_features, nhead=8, activation='gelu',dropout=0.1, batch_first=True)\n\n        # L1/L2 -> L5/S1\n        self.ll1 = nn.Linear(512, 64)\n        self.ll2 = nn.Linear(512, 64)\n        self.ll3 = nn.Linear(512, 64)\n        self.ll4 = nn.Linear(512, 64)\n        self.ll5 = nn.Linear(512, 64)\n        self.rl1 = nn.Linear(512, 64)\n        self.rl2 = nn.Linear(512, 64)\n        self.rl3 = nn.Linear(512, 64)\n        self.rl4 = nn.Linear(512, 64)\n        self.rl5 = nn.Linear(512, 64)\n    def forward(self, x, label=None):\n        # Đảm bảo là x có 3 kênh đầu vào, nếu chỉ có 1 thì thêm 1 dims vào\n        # hơi dư thừa nếu về mặt Conv2D\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        ll1 = self.ll1(x) # [B, 64]\n        ll2 = self.ll2(x)\n        ll3 = self.ll3(x)\n        ll4 = self.ll4(x)\n        ll5 = self.ll5(x)\n        rl1 = self.rl1(x)\n        rl2 = self.rl2(x)\n        rl3 = self.rl3(x)\n        rl4 = self.rl4(x)\n        rl5 = self.rl5(x)\n\n        # Trả về dictionary chứa 10 output embeddings\n        return {'left_L1/L2': ll1,'left_L2/L3': ll2,'left_L3/L4': ll3, 'left_L4/L5': ll4, 'left_L5/S1': ll5,\n                'right_L1/L2': rl1, 'right_L2/L3': rl2, 'right_L3/L4': rl3, 'right_L4/L5': rl4, 'right_L5/S1': rl5}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:14.021033Z","iopub.execute_input":"2025-06-07T01:05:14.021237Z","iopub.status.idle":"2025-06-07T01:05:14.033897Z","shell.execute_reply.started":"2025-06-07T01:05:14.021223Z","shell.execute_reply":"2025-06-07T01:05:14.033147Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ConvNeXtNFNDepthDetect sử dụng backbone của ConvNeXt 3D để phát hiện bệnh nfn cột sống ở thắt lưng\nclass ConvNextNFNDepthDetect(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # backbone (encoder)\n        # Chức năng chính là trích xuất các đặc trưng từ 1 ảnh - volumn 3D MRI kích thước [B, 1, D, H, W]\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=32, depths=[3,3,9,3], widths=[64, 128, 256, 512])\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(512))\n        # left l1 -> l5\n        self.ll1 = nn.Linear(512, 32)\n        self.ll2 = nn.Linear(512, 32)\n        self.ll3 = nn.Linear(512, 32)\n        self.ll4 = nn.Linear(512, 32)\n        self.ll5 = nn.Linear(512, 32)\n\n        # right l1 -> l5\n        self.rl1 = nn.Linear(512, 32)\n        self.rl2 = nn.Linear(512, 32)\n        self.rl3 = nn.Linear(512, 32)\n        self.rl4 = nn.Linear(512, 32)\n        self.rl5 = nn.Linear(512, 32)\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        ll1 = self.ll1(x)\n        ll2 = self.ll2(x)\n        ll3 = self.ll3(x)\n        ll4 = self.ll4(x)\n        ll5 = self.ll5(x)\n        rl1 = self.rl1(x)\n        rl2 = self.rl2(x)\n        rl3 = self.rl3(x)\n        rl4 = self.rl4(x)\n        rl5 = self.rl5(x)\n        # 10 Linear header dự đoán đặc trung cho từng kh đốt sống\n        return {'left_L1/L2': ll1,'left_L2/L3': ll2,'left_L3/L4': ll3, 'left_L4/L5': ll4, 'left_L5/S1': ll5,\n                'right_L1/L2': rl1, 'right_L2/L3': rl2, 'right_L3/L4': rl3, 'right_L4/L5': rl4, 'right_L5/S1': rl5}\n\nclass ConvNextNFNDepthDetect(nn.Module):\n    def __init__(self, widths):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=widths[0]//2, depths=[3,3,9,3], widths=widths)\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(widths[-1]))\n        self.ll1 = nn.Linear(widths[-1], 32)\n        self.ll2 = nn.Linear(widths[-1], 32)\n        self.ll3 = nn.Linear(widths[-1], 32)\n        self.ll4 = nn.Linear(widths[-1], 32)\n        self.ll5 = nn.Linear(widths[-1], 32)\n        self.rl1 = nn.Linear(widths[-1], 32)\n        self.rl2 = nn.Linear(widths[-1], 32)\n        self.rl3 = nn.Linear(widths[-1], 32)\n        self.rl4 = nn.Linear(widths[-1], 32)\n        self.rl5 = nn.Linear(widths[-1], 32)\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        ll1 = self.ll1(x)\n        ll2 = self.ll2(x)\n        ll3 = self.ll3(x)\n        ll4 = self.ll4(x)\n        ll5 = self.ll5(x)\n        rl1 = self.rl1(x)\n        rl2 = self.rl2(x)\n        rl3 = self.rl3(x)\n        rl4 = self.rl4(x)\n        rl5 = self.rl5(x)\n        return {'left_L1/L2': ll1,'left_L2/L3': ll2,'left_L3/L4': ll3, 'left_L4/L5': ll4, 'left_L5/S1': ll5,\n                'right_L1/L2': rl1, 'right_L2/L3': rl2, 'right_L3/L4': rl3, 'right_L4/L5': rl4, 'right_L5/S1': rl5}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:14.034564Z","iopub.execute_input":"2025-06-07T01:05:14.034776Z","iopub.status.idle":"2025-06-07T01:05:14.048688Z","shell.execute_reply.started":"2025-06-07T01:05:14.034753Z","shell.execute_reply":"2025-06-07T01:05:14.048114Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ConvNextSCSDepthDetect(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=32, depths=[3,3,9,3], widths=[64, 128, 256, 512])\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                        nn.Flatten(1),\n                                        nn.LayerNorm(512))\n        self.l1 = nn.Linear(512, 32)\n        self.l2 = nn.Linear(512, 32)\n        self.l3 = nn.Linear(512, 32)\n        self.l4 = nn.Linear(512, 32)\n        self.l5 = nn.Linear(512, 32)\n    def forward(self, x, label = None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        x = self.l1(x)\n        x = self.l2(x)\n        x = self.l3(x)\n        x = self.l4(x)\n        x = self.l5(x)\n        \n        return {'L1/L2': l1, 'L2/L3': l2, 'L3/L4': l3, 'L4/L5': l4, 'L5/S1': l5}\n        \nclass ConvNextSCSDepthDetect(nn.Module):\n    def __init__(self, widths):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=widths[0]//2, depths=[3,3,9,3], widths=widths)\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                        nn.Flatten(1),\n                                        nn.LayerNorm(512))\n        self.l1 = nn.Linear(widths[-1], 32)\n        self.l2 = nn.Linear(widths[-1], 32)\n        self.l3 = nn.Linear(widths[-1], 32)\n        self.l4 = nn.Linear(widths[-1], 32)\n        self.l5 = nn.Linear(widths[-1], 32)\n    def forward(self, x, label = None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        x = self.l1(x)\n        x = self.l2(x)\n        x = self.l3(x)\n        x = self.l4(x)\n        x = self.l5(x)\n        \n        return {'L1/L2': l1, 'L2/L3': l2, 'L3/L4': l3, 'L4/L5': l4, 'L5/S1': l5}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:14.049255Z","iopub.execute_input":"2025-06-07T01:05:14.049486Z","iopub.status.idle":"2025-06-07T01:05:14.064776Z","shell.execute_reply.started":"2025-06-07T01:05:14.049466Z","shell.execute_reply":"2025-06-07T01:05:14.064079Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RegConvNextNFNDepthDetect(nn.Module):\n    def __init__(self, widths):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=widths[0]//2, depths=[3,3,9,3], widths=widths)\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(1024))\n        self.ll1 = nn.Linear(1024, 3)\n        self.ll2 = nn.Linear(1024, 3)\n        self.ll3 = nn.Linear(1024, 3)\n        self.ll4 = nn.Linear(1024, 3)\n        self.ll5 = nn.Linear(1024, 3)\n        self.rl1 = nn.Linear(1024, 3)\n        self.rl2 = nn.Linear(1024, 3)\n        self.rl3 = nn.Linear(1024, 3)\n        self.rl4 = nn.Linear(1024, 3)\n        self.rl5 = nn.Linear(1024, 3)\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        ll1 = self.ll1(x).sigmoid()\n        ll2 = self.ll2(x).sigmoid()\n        ll3 = self.ll3(x).sigmoid()\n        ll4 = self.ll4(x).sigmoid()\n        ll5 = self.ll5(x).sigmoid()\n        rl1 = self.rl1(x).sigmoid()\n        rl2 = self.rl2(x).sigmoid()\n        rl3 = self.rl3(x).sigmoid()\n        rl4 = self.rl4(x).sigmoid()\n        rl5 = self.rl5(x).sigmoid()\n        return {'left_L1/L2': ll1,'left_L2/L3': ll2,'left_L3/L4': ll3, 'left_L4/L5': ll4, 'left_L5/S1': ll5,\n                'right_L1/L2': rl1, 'right_L2/L3': rl2, 'right_L3/L4': rl3, 'right_L4/L5': rl4, 'right_L5/S1': rl5}\n\n\nclass RegConvNextSCSDepthDetect(nn.Module):\n    def __init__(self, widths):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=widths[0]//2, depths=[3,3,9,3], widths=widths)\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(1024),\n                                     )\n        self.l1 = nn.Linear(1024, 3)\n        self.l2 = nn.Linear(1024, 3)\n        self.l3 = nn.Linear(1024, 3)\n        self.l4 = nn.Linear(1024, 3)\n        self.l5 = nn.Linear(1024, 3)\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        l1 = self.l1(x).sigmoid()\n        l2 = self.l2(x).sigmoid()\n        l3 = self.l3(x).sigmoid()\n        l4 = self.l4(x).sigmoid()\n        l5 = self.l5(x).sigmoid()\n        return {'L1/L2': l1, 'L2/L3': l2, 'L3/L4': l3, 'L4/L5': l4, 'L5/S1': l5}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:14.065421Z","iopub.execute_input":"2025-06-07T01:05:14.065656Z","iopub.status.idle":"2025-06-07T01:05:14.079234Z","shell.execute_reply.started":"2025-06-07T01:05:14.065638Z","shell.execute_reply":"2025-06-07T01:05:14.078526Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Base Models\n\n\nclass ConvNextStem(nn.Sequential):\n    def __init__(self, in_features: int, out_features: int):\n        super().__init__(\n            nn.Conv3d(in_features, out_features, kernel_size=(1, 2, 2), stride=(1, 2, 2)),\n            nn.GroupNorm(num_groups=1, num_channels=out_features)\n        )\n\nclass LayerScaler(nn.Module):\n    def __init__(self, init_value: float, dimensions: int):\n        super().__init__()\n        self.gamma = nn.Parameter(init_value * torch.ones((dimensions)),\n                                    requires_grad=True)\n\n    def forward(self, x):\n        return self.gamma[None,...,None,None] * x\n\nclass BottleNeckBlock(nn.Module):\n    def __init__(\n        self,\n        in_features: int,\n        out_features: int,\n        expansion: int = 4,\n        drop_p: float = .0,\n        layer_scaler_init_value: float = 1e-6,\n    ):\n        super().__init__()\n        expanded_features = out_features * expansion\n        self.block = nn.Sequential(\n            # narrow -> wide (with depth-wise and bigger kernel)\n            nn.Conv3d(\n                in_features, in_features, kernel_size=(2, 7, 7), padding='same', bias=False, groups=in_features\n            ),\n            # GroupNorm with num_groups=1 is the same as LayerNorm but works for 2D data\n            nn.GroupNorm(num_groups=in_features, num_channels=in_features),\n            # wide -> wide\n            nn.Conv3d(in_features, expanded_features, kernel_size=1),\n            nn.GELU(),\n            # wide -> narrow\n            nn.Conv3d(expanded_features, out_features, kernel_size=1),\n        )\n        #self.layer_scaler = LayerScaler(layer_scaler_init_value, out_features)\n        #self.drop_path = StochasticDepth(drop_p, mode=\"batch\")\n\n\n    def forward(self, x: Tensor) -> Tensor:\n        res = x\n        x = self.block(x)\n        #x = self.layer_scaler(x)\n        #x = self.drop_path(x)\n        x += res\n        return x\n\nclass ConvNexStage(nn.Sequential):\n    def __init__(\n        self, in_features: int, out_features: int, depth: int, **kwargs\n    ):\n        super().__init__(\n            # add the downsampler\n            nn.Sequential(\n                nn.GroupNorm(num_groups=in_features, num_channels=in_features),\n                nn.Conv3d(in_features, out_features, kernel_size=(2, 2, 2), stride=(2, 2, 2))\n            ),\n            *[\n                BottleNeckBlock(out_features, out_features, **kwargs)\n                for _ in range(depth)\n            ],\n        )\n\nclass ConvNextEncoder(nn.Module):\n    def __init__(\n        self,\n        in_channels: int,\n        stem_features: int,\n        depths: List[int],\n        widths: List[int],\n        drop_p: float = .0,\n    ):\n        super().__init__()\n        self.stem = ConvNextStem(in_channels, stem_features)\n\n        in_out_widths = list(zip(widths, widths[1:]))\n        # create drop paths probabilities (one for each stage)\n        drop_probs = [x.item() for x in torch.linspace(0, drop_p, sum(depths))]\n\n        self.stages = nn.ModuleList(\n            [\n                ConvNexStage(stem_features, widths[0], depths[0], drop_p=drop_probs[0]),\n                *[\n                    ConvNexStage(in_features, out_features, depth, drop_p=drop_p)\n                    for (in_features, out_features), depth, drop_p in zip(\n                        in_out_widths, depths[1:], drop_probs[1:]\n                    )\n                ],\n            ]\n        )\n\n\n    def forward(self, x):\n        x = self.stem(x)\n        for stage in self.stages:\n            x = stage(x)\n        return x\n\nclass ClassificationHead(nn.Sequential):\n    def __init__(self):\n        super().__init__(\n            nn.AdaptiveAvgPool3d((1, 1, 1)),\n            nn.Flatten(1),\n            nn.LayerNorm(512),\n            nn.Linear(512, 3)\n        )\nclass Flatten(nn.Sequential):\n    def __init__(self):\n        super().__init__(\n            nn.AdaptiveAvgPool3d((1, 1, 1)),\n            nn.Flatten(1),\n            nn.LayerNorm(512)\n        )\n\n\nclass ConvNextSSDepthDetect(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=32, depths=[3,3,9,3], widths=[64, 128, 256, 512])\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(512))\n        self.ll1 = nn.Linear(512, 96)\n        self.ll2 = nn.Linear(512, 96)\n        self.ll3 = nn.Linear(512, 96)\n        self.ll4 = nn.Linear(512, 96)\n        self.ll5 = nn.Linear(512, 96)\n        self.rl1 = nn.Linear(512, 96)\n        self.rl2 = nn.Linear(512, 96)\n        self.rl3 = nn.Linear(512, 96)\n        self.rl4 = nn.Linear(512, 96)\n        self.rl5 = nn.Linear(512, 96)\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        ll1 = self.ll1(x)\n        ll2 = self.ll2(x)\n        ll3 = self.ll3(x)\n        ll4 = self.ll4(x)\n        ll5 = self.ll5(x)\n        rl1 = self.rl1(x)\n        rl2 = self.rl2(x)\n        rl3 = self.rl3(x)\n        rl4 = self.rl4(x)\n        rl5 = self.rl5(x)\n        return {'left_L1/L2': ll1,'left_L2/L3': ll2,'left_L3/L4': ll3, 'left_L4/L5': ll4, 'left_L5/S1': ll5,\n                'right_L1/L2': rl1, 'right_L2/L3': rl2, 'right_L3/L4': rl3, 'right_L4/L5': rl4, 'right_L5/S1': rl5}\nclass PositionalEncoding(nn.Module):\n\n    def __init__(self, d_model: int, dropout: float = 0.1, max_len: int = 5000):\n        super().__init__()\n        self.dropout = nn.Dropout(p=dropout)\n\n        position = torch.arange(max_len).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))\n        pe = torch.zeros(max_len, 1, d_model)\n        pe[:, 0, 0::2] = torch.sin(position * div_term)\n        pe[:, 0, 1::2] = torch.cos(position * div_term)\n        self.register_buffer('pe', pe)\n\n    def forward(self, x: Tensor) -> Tensor:\n        \"\"\"\n        Args:\n            x: Tensor, shape [batch_size, seq_len, embedding_dim]\n        \"\"\"\n        x = x.permute(1, 0, 2)\n        x = x + self.pe[:x.size(0)]\n        return self.dropout(x.permute(1, 0, 2))\n\nclass AttentionSSDepthDetect(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = timm.create_model('convnext_base.fb_in22k_ft_in1k_384', in_chans=3, pretrained=True, num_classes=0)\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(512))\n        self.in_features = self.encoder.num_features\n        self.rpe = PositionalEncoding(self.in_features, dropout=0., max_len=64)\n        self.transformer0 = nn.TransformerEncoderLayer(d_model=self.in_features, nhead=8, activation='gelu',dropout=0.1, batch_first=True)\n        self.transformer1 = nn.TransformerEncoderLayer(d_model=self.in_features, nhead=8, activation='gelu',dropout=0.1, batch_first=True)\n\n        self.ll1 = nn.Linear(512, 64)\n        self.ll2 = nn.Linear(512, 64)\n        self.ll3 = nn.Linear(512, 64)\n        self.ll4 = nn.Linear(512, 64)\n        self.ll5 = nn.Linear(512, 64)\n        self.rl1 = nn.Linear(512, 64)\n        self.rl2 = nn.Linear(512, 64)\n        self.rl3 = nn.Linear(512, 64)\n        self.rl4 = nn.Linear(512, 64)\n        self.rl5 = nn.Linear(512, 64)\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        ll1 = self.ll1(x)\n        ll2 = self.ll2(x)\n        ll3 = self.ll3(x)\n        ll4 = self.ll4(x)\n        ll5 = self.ll5(x)\n        rl1 = self.rl1(x)\n        rl2 = self.rl2(x)\n        rl3 = self.rl3(x)\n        rl4 = self.rl4(x)\n        rl5 = self.rl5(x)\n        return {'left_L1/L2': ll1,'left_L2/L3': ll2,'left_L3/L4': ll3, 'left_L4/L5': ll4, 'left_L5/S1': ll5,\n                'right_L1/L2': rl1, 'right_L2/L3': rl2, 'right_L3/L4': rl3, 'right_L4/L5': rl4, 'right_L5/S1': rl5}\n\nclass ConvNextNFNDepthDetect(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=32, depths=[3,3,9,3], widths=[64, 128, 256, 512])\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(512))\n        self.ll1 = nn.Linear(512, 32)\n        self.ll2 = nn.Linear(512, 32)\n        self.ll3 = nn.Linear(512, 32)\n        self.ll4 = nn.Linear(512, 32)\n        self.ll5 = nn.Linear(512, 32)\n        self.rl1 = nn.Linear(512, 32)\n        self.rl2 = nn.Linear(512, 32)\n        self.rl3 = nn.Linear(512, 32)\n        self.rl4 = nn.Linear(512, 32)\n        self.rl5 = nn.Linear(512, 32)\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        ll1 = self.ll1(x)\n        ll2 = self.ll2(x)\n        ll3 = self.ll3(x)\n        ll4 = self.ll4(x)\n        ll5 = self.ll5(x)\n        rl1 = self.rl1(x)\n        rl2 = self.rl2(x)\n        rl3 = self.rl3(x)\n        rl4 = self.rl4(x)\n        rl5 = self.rl5(x)\n        return {'left_L1/L2': ll1,'left_L2/L3': ll2,'left_L3/L4': ll3, 'left_L4/L5': ll4, 'left_L5/S1': ll5,\n                'right_L1/L2': rl1, 'right_L2/L3': rl2, 'right_L3/L4': rl3, 'right_L4/L5': rl4, 'right_L5/S1': rl5}\n\nclass ConvNextSCSDepthDetect(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=32, depths=[3,3,9,3], widths=[64, 128, 256, 512])\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(512))\n        self.l1 = nn.Linear(512, 32)\n        self.l2 = nn.Linear(512, 32)\n        self.l3 = nn.Linear(512, 32)\n        self.l4 = nn.Linear(512, 32)\n        self.l5 = nn.Linear(512, 32)\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        l1 = self.l1(x)\n        l2 = self.l2(x)\n        l3 = self.l3(x)\n        l4 = self.l4(x)\n        l5 = self.l5(x)\n        return {'L1/L2': l1, 'L2/L3': l2, 'L3/L4': l3, 'L4/L5': l4, 'L5/S1': l5}\n    \nclass ConvNextSCSDepthDetect(nn.Module):\n    def __init__(self, widths):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=widths[0]//2, depths=[3,3,9,3], widths=widths)\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(widths[-1]),\n                                     )\n        self.l1 = nn.Linear(widths[-1], 32)\n        self.l2 = nn.Linear(widths[-1], 32)\n        self.l3 = nn.Linear(widths[-1], 32)\n        self.l4 = nn.Linear(widths[-1], 32)\n        self.l5 = nn.Linear(widths[-1], 32)\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        l1 = self.l1(x)\n        l2 = self.l2(x)\n        l3 = self.l3(x)\n        l4 = self.l4(x)\n        l5 = self.l5(x)\n        return {'L1/L2': l1, 'L2/L3': l2, 'L3/L4': l3, 'L4/L5': l4, 'L5/S1': l5}\n\nclass ConvNextNFNDepthDetect(nn.Module):\n    def __init__(self, widths):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=widths[0]//2, depths=[3,3,9,3], widths=widths)\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(widths[-1]))\n        self.ll1 = nn.Linear(widths[-1], 32)\n        self.ll2 = nn.Linear(widths[-1], 32)\n        self.ll3 = nn.Linear(widths[-1], 32)\n        self.ll4 = nn.Linear(widths[-1], 32)\n        self.ll5 = nn.Linear(widths[-1], 32)\n        self.rl1 = nn.Linear(widths[-1], 32)\n        self.rl2 = nn.Linear(widths[-1], 32)\n        self.rl3 = nn.Linear(widths[-1], 32)\n        self.rl4 = nn.Linear(widths[-1], 32)\n        self.rl5 = nn.Linear(widths[-1], 32)\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        ll1 = self.ll1(x)\n        ll2 = self.ll2(x)\n        ll3 = self.ll3(x)\n        ll4 = self.ll4(x)\n        ll5 = self.ll5(x)\n        rl1 = self.rl1(x)\n        rl2 = self.rl2(x)\n        rl3 = self.rl3(x)\n        rl4 = self.rl4(x)\n        rl5 = self.rl5(x)\n        return {'left_L1/L2': ll1,'left_L2/L3': ll2,'left_L3/L4': ll3, 'left_L4/L5': ll4, 'left_L5/S1': ll5,\n                'right_L1/L2': rl1, 'right_L2/L3': rl2, 'right_L3/L4': rl3, 'right_L4/L5': rl4, 'right_L5/S1': rl5}\n    \nclass RegConvNextNFNDepthDetect(nn.Module):\n    def __init__(self, widths):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=widths[0]//2, depths=[3,3,9,3], widths=widths)\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(1024))\n        self.ll1 = nn.Linear(1024, 3)\n        self.ll2 = nn.Linear(1024, 3)\n        self.ll3 = nn.Linear(1024, 3)\n        self.ll4 = nn.Linear(1024, 3)\n        self.ll5 = nn.Linear(1024, 3)\n        self.rl1 = nn.Linear(1024, 3)\n        self.rl2 = nn.Linear(1024, 3)\n        self.rl3 = nn.Linear(1024, 3)\n        self.rl4 = nn.Linear(1024, 3)\n        self.rl5 = nn.Linear(1024, 3)\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        ll1 = self.ll1(x).sigmoid()\n        ll2 = self.ll2(x).sigmoid()\n        ll3 = self.ll3(x).sigmoid()\n        ll4 = self.ll4(x).sigmoid()\n        ll5 = self.ll5(x).sigmoid()\n        rl1 = self.rl1(x).sigmoid()\n        rl2 = self.rl2(x).sigmoid()\n        rl3 = self.rl3(x).sigmoid()\n        rl4 = self.rl4(x).sigmoid()\n        rl5 = self.rl5(x).sigmoid()\n        return {'left_L1/L2': ll1,'left_L2/L3': ll2,'left_L3/L4': ll3, 'left_L4/L5': ll4, 'left_L5/S1': ll5,\n                'right_L1/L2': rl1, 'right_L2/L3': rl2, 'right_L3/L4': rl3, 'right_L4/L5': rl4, 'right_L5/S1': rl5}\n\nclass RegConvNextSCSDepthDetect(nn.Module):\n    def __init__(self, widths):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=widths[0]//2, depths=[3,3,9,3], widths=widths)\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(1024),\n                                     )\n        self.l1 = nn.Linear(1024, 3)\n        self.l2 = nn.Linear(1024, 3)\n        self.l3 = nn.Linear(1024, 3)\n        self.l4 = nn.Linear(1024, 3)\n        self.l5 = nn.Linear(1024, 3)\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)\n        x = self.flatten(x)\n        l1 = self.l1(x).sigmoid()\n        l2 = self.l2(x).sigmoid()\n        l3 = self.l3(x).sigmoid()\n        l4 = self.l4(x).sigmoid()\n        l5 = self.l5(x).sigmoid()\n        return {'L1/L2': l1, 'L2/L3': l2, 'L3/L4': l3, 'L4/L5': l4, 'L5/S1': l5}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:14.080078Z","iopub.execute_input":"2025-06-07T01:05:14.080308Z","iopub.status.idle":"2025-06-07T01:05:14.124485Z","shell.execute_reply.started":"2025-06-07T01:05:14.080285Z","shell.execute_reply":"2025-06-07T01:05:14.123894Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Lightning Module","metadata":{}},{"cell_type":"code","source":"# Là module wrapper cho các model khác, sử dụng pytorch_lightning \nclass DepthDetectModule(pl.LightningModule):\n    def __init__(self, condition, widths=None, model_type='regression'):\n        super().__init__()\n        self.config = config\n        if condition == 'scs':\n            if model_type == 'regression': \n                self.model = RegConvNextSCSDepthDetect(widths)\n            else: \n                self.model = ConvNextSCSDepthDetect(widths)\n        elif condition == 'nfn':\n            if model_type == 'regression': \n                self.model = RegConvNextNFNDepthDetect(widths)\n            else: \n                self.model = ConvNextNFNDepthDetect(widths)\n    def forward(self, batch):\n        preds = self.model(batch)\n        return preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:14.125189Z","iopub.execute_input":"2025-06-07T01:05:14.125423Z","iopub.status.idle":"2025-06-07T01:05:14.138250Z","shell.execute_reply.started":"2025-06-07T01:05:14.125396Z","shell.execute_reply":"2025-06-07T01:05:14.137702Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"markdown","source":"### instance prediction\n* dự đoán độ sâu (depth) cho từng instance / study\n* Đây là bước inference (suy diễn) từ các checkpoint .ckpt đã được huấn luyện trước đó","metadata":{}},{"cell_type":"code","source":"%%time\nprefix = ''\nimport warnings\nwarnings.filterwarnings(\"ignore\")\ndepth_predict = {'scs': {\n                     'L1/L2':[], \n                     'L2/L3': [], \n                     'L3/L4': [], \n                     'L4/L5': [], \n                     'L5/S1': []\n                     }, \n                 'nfn': {\n                     'left_L1/L2': [], \n                     'left_L2/L3': [], \n                     'left_L3/L4': [], \n                     'left_L4/L5': [], \n                     'left_L5/S1': [], \n                     'right_L1/L2': [], \n                     'right_L2/L3': [], \n                     'right_L3/L4': [], \n                     'right_L4/L5': [], \n                     'right_L5/S1': [], \n                     }\n                    }\nmodel_path_dict = {\n    'scs': [\n        '/kaggle/input/rsna-spine-final-models/scs_depth_1024_ssr_0.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_1024_ssr_1.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_1024_ssr_2.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_1024_ssr_3.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_1024_ssr_4.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_0.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_1.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_2.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_3.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_4.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_1024_ssr_l1_0.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_1024_ssr_l1_1.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_1024_ssr_l1_2.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_1024_ssr_l1_3.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_depth_1024_ssr_l1_4.ckpt', \n    ], \n    'nfn': [\n        '/kaggle/input/rsna-spine-final-models/nfn_depth_1024_ssr_0.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_1024_ssr_1.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_1024_ssr_2.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_1024_ssr_3.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_1024_ssr_4.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_0.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_1.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_2.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_3.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_4.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_1024_ssr_l1_0.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_1024_ssr_l1_1.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_1024_ssr_l1_2.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_1024_ssr_l1_3.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_depth_1024_ssr_l1_4.ckpt', \n    ]\n}\n##############DEPTH DETECT#########################\nfor condition in ['nfn', 'scs']:\n    print(condition)\n    model_path_list = model_path_dict[condition]\n    for model_path in model_path_list:\n        _meta_df = meta_df.copy()\n        #_series = series.copy()\n        dataset_test = DepthDetectDataset(_meta_df, condition, 'sub')\n        data_loader_test = DataLoader(\n            dataset_test,\n            batch_size=config[\"test_bs\"], \n            shuffle=False,\n            num_workers=4,\n            pin_memory=False\n        )\n        model_name = model_path.split('/')[-1]\n        if '1024' in model_name: \n            widths = [128, 256, 512, 1024]\n        else: \n            widths = [64, 128, 256, 512]\n        if 'l1' in model_name: \n            model_type = 'regression'\n        else: \n            model_type = 'classification'\n        model = DepthDetectModule.load_from_checkpoint(model_path, condition=condition, widths=widths, model_type=model_type)\n        model.eval()\n        model.zero_grad()\n        model.to(device)\n\n        pred_temp = {}\n        for k in depth_predict[condition].keys(): \n            pred_temp[k] = []\n        study_id_list = []\n        with torch.no_grad():\n            for data in tqdm(data_loader_test, total=len(data_loader_test)):\n                images, study_id = data\n                images = images.to(device)\n                preds = model.forward(images)\n                #print(preds)\n                if model_type == 'regression': \n                    for k, v in preds.items(): \n                        pred_temp[k].append((v[:, -1]*32).to('cpu').detach().numpy())\n                else: \n                    for k, v in preds.items(): \n                        pred_temp[k].append(torch.argmax(v, dim=1).to('cpu').detach().numpy())\n                study_id_list.append(study_id.to('cpu').reshape(-1).detach().numpy())\n                del images, study_id, preds\n                gc.collect()\n        for k, v in pred_temp.items(): \n            depth_predict[condition][k].append(np.concatenate(v))\n        study_id = np.concatenate(study_id_list)\n        del pred_temp, study_id_list\n        gc.collect()\n        \n    for k, v in depth_predict[condition].items(): \n        depth_predict[condition][k] = np.median(np.array(depth_predict[condition][k]), axis=0)\n    depth_predict[condition]['study_id'] = study_id\n    del study_id\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:05:14.138876Z","iopub.execute_input":"2025-06-07T01:05:14.139088Z","iopub.status.idle":"2025-06-07T01:10:04.647652Z","shell.execute_reply.started":"2025-06-07T01:05:14.139063Z","shell.execute_reply":"2025-06-07T01:10:04.647004Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Create label coordinate % align","metadata":{}},{"cell_type":"code","source":"def create_label_ins(study_id, depth, level, condition, desc): \n    coor_dict = {'study_id': [], 'series_id': [], 'instance_number': []}\n    _meta = meta_df.loc[meta_df.series_description==desc]\n    for s, d in zip(study_id, depth): \n        sub_meta = _meta.loc[_meta.study_id==s]\n        sub_meta = sub_meta.sort_values('ipp_x', ascending=True).reset_index(drop=True)\n        if len(sub_meta) > 32: \n            d = (d/32)*len(sub_meta)\n        try: \n            row = sub_meta.iloc[round(d)]\n        except: \n            if condition == 'Spinal Canal Stenosis': \n                row = sub_meta.iloc[int(len(sub_meta)//2)]\n            elif condition == 'Left Neural Foraminal Narrowing': \n                row = sub_meta.iloc[int(2*(len(sub_meta)//3))]\n            elif condition == 'Right Neural Foraminal Narrowing': \n                row = sub_meta.iloc[int(len(sub_meta)//3)]\n            print(s)\n        coor_dict['study_id'].append(s)\n        coor_dict['series_id'].append(row.series_id)\n        coor_dict['instance_number'].append(row.instance_number)\n    coor_dict['condition'] = condition\n    coor_dict['level'] = level.split('_')[-1]\n    return pd.DataFrame(coor_dict)\n\nscs_study_id = depth_predict['scs']['study_id']\nscs_coor_list = []\nfor k, v in depth_predict['scs'].items(): \n    if k != 'study_id': \n        scs_coor_list.append(create_label_ins(scs_study_id, v, k, 'Spinal Canal Stenosis', 'Sagittal T2/STIR'))\n\nnfn_study_id = depth_predict['nfn']['study_id']\nnfn_coor_list = []\nfor k, v in depth_predict['nfn'].items(): \n    if k != 'study_id': \n        if k.split('_')[0] == 'left': \n            condition = 'Left Neural Foraminal Narrowing'\n        else: \n            condition = 'Right Neural Foraminal Narrowing'\n        nfn_coor_list.append(create_label_ins(nfn_study_id, v, k, condition, 'Sagittal T1'))\nscs_coor = pd.concat(scs_coor_list)\nnfn_coor = pd.concat(nfn_coor_list)\npred_coor = pd.concat([scs_coor, nfn_coor]).sort_values(['study_id', 'series_id', 'level'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:10:04.648408Z","iopub.execute_input":"2025-06-07T01:10:04.648627Z","iopub.status.idle":"2025-06-07T01:10:04.682004Z","shell.execute_reply.started":"2025-06-07T01:10:04.648608Z","shell.execute_reply":"2025-06-07T01:10:04.681331Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del scs_coor, nfn_coor, scs_coor_list, nfn_coor_list, depth_predict\ngc.collect()\npred_coor.head()\npred_coor.to_csv('stage1_coor.csv', index=False)\nos.listdir('/kaggle/working/')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:10:04.682906Z","iopub.execute_input":"2025-06-07T01:10:04.683156Z","iopub.status.idle":"2025-06-07T01:10:04.952769Z","shell.execute_reply.started":"2025-06-07T01:10:04.683132Z","shell.execute_reply":"2025-06-07T01:10:04.951931Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Second Stage (xy inference)\n* infer xy-coordinate of locations of sagittal t1 & t2\n* ensemble or align (rule base)","metadata":{}},{"cell_type":"markdown","source":"## Coordinate prediction dataset","metadata":{}},{"cell_type":"code","source":"class CoorDetectDataset(Dataset):\n    def __init__(self, coor, meta, condition, usage='train'):\n        if condition == 'scs':\n            coor = coor.loc[coor.condition=='Spinal Canal Stenosis']\n        elif condition == 'ss':\n            coor = coor.loc[(coor.condition=='Left Subarticular Stenosis') | (coor.condition=='Right Subarticular Stenosis')]\n        elif condition == 'nfn':\n            coor = coor.loc[(coor.condition=='Right Neural Foraminal Narrowing') | (coor.condition=='Left Neural Foraminal Narrowing')]\n        #g_coor = coor.groupby('study_id').count()\n        #if condition == 'scs':\n        #    self.id = g_coor.loc[g_coor.series_id==5].reset_index().study_id.unique()\n        #else:\n        #    self.id = g_coor.loc[g_coor.series_id==10].reset_index().study_id.unique()\n        self.id = coor.study_id.unique()\n        self.coor = coor\n        self.meta = meta\n        self.condition = condition\n        self.usage = usage\n        if 3637444890 in self.id: \n            self.id.remove(3637444890)\n        #self.id = [2773343225]\n        #self.id = [1782095928]\n\n        self.resize = v2.Resize((384, 384))\n        \n    def __getitem__(self, index):\n        study_id = self.id[index]\n        #print(study_id)\n        #try:\n        if self.condition == 'scs':\n            volume = self.for_scs(study_id)\n        elif self.condition == 'nfn':\n            volume = self.for_nfn(study_id)\n        if self.condition == 'ss':\n            volume = self.for_ss(study_id)\n        return volume, torch.tensor(study_id)\n\n    def for_scs(self, study_id):\n        meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Sagittal T2/STIR')]\n        meta = meta.sort_values('ipp_x', ascending=True).reset_index(drop=True)\n        #img = [self.normalize(self.load_dicom(f'/content/train_images/{row.study_id}/{row.series_id}/{row.instance_number}.dcm')) for _, row in meta.iterrows()]\n        coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Spinal Canal Stenosis')]\n        meta_list = []\n        for _, row in coor.iterrows():\n            series_id, instance_number = row.series_id, row.instance_number\n            meta_list.append(meta.loc[(meta.series_id==series_id) & (meta.instance_number==instance_number)])\n        sub_meta = pd.concat(meta_list)\n        idx = meta.loc[meta.ipp_x == sub_meta.ipp_x.median()].index[0]\n        #print(old_idx)\n        img_row = meta.iloc[idx]\n        before_img_row = meta.iloc[idx-1]\n        after_img_row = meta.iloc[idx+1]\n        img = self.normalize(self.load_dicom(IMAGE_PATH + f'{img_row.study_id}/{img_row.series_id}/{img_row.instance_number}.dcm'))\n        bimg = self.normalize(self.load_dicom(IMAGE_PATH + f'{before_img_row.study_id}/{before_img_row.series_id}/{before_img_row.instance_number}.dcm'))\n        aimg = self.normalize(self.load_dicom(IMAGE_PATH + f'{after_img_row.study_id}/{after_img_row.series_id}/{after_img_row.instance_number}.dcm'))\n        img = self.resize(torch.tensor(img[None, ...]))\n        bimg = self.resize(torch.tensor(bimg[None, ...]))\n        aimg = self.resize(torch.tensor(aimg[None, ...]))\n        img = torch.cat([bimg, img, aimg]).to(torch.float32)\n        return img\n    def for_ss(self, study_id):\n        meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Axial T2')]\n        meta = meta.sort_values('ipp_z', ascending=False).reset_index(drop=True)\n        img = [self.normalize(self.load_dicom(f'/content/train_images/{row.study_id}/{row.series_id}/{row.instance_number}.dcm')) for _, row in meta.iterrows()]\n        coor = self.coor.loc[(self.coor.study_id==study_id)]\n        coor_dict = {}\n        for _, row in coor.iterrows():\n            series_id, instance_number = row.series_id, row.instance_number\n            target_row = meta.loc[(meta.series_id==series_id) & (meta.instance_number==instance_number)]\n            idx = target_row.index[0]\n            #print(row.level, idx, idx/len(img))\n            #plt.title(row.level)\n            #plt.imshow(img[idx])\n            #mask = torch.zeros(img[idx].shape)\n            #mask[int(row.y)-10:int((row.y))+10, int(row.x)-10:int((row.x))+10] = 1\n            #plt.imshow(mask, alpha=0.5)\n            #plt.show()\n            height, width = img[idx].shape\n            z = idx/depth if len(img) < depth else idx/len(img)\n            x = row.x/width\n            y = row.y/height\n            if row.condition == 'Right Subarticular Stenosis':\n                coor_dict['right_' + row.level] = torch.tensor([x, y, z]).to(torch.float32)\n            else:\n                coor_dict['left_' + row.level] = torch.tensor([x, y, z]).to(torch.float32)\n        volume = torch.cat([self.resize(torch.tensor(i)[None, ...]).to(torch.float32) for i in img]).contiguous()\n        if volume.shape[0] < depth:\n            volume = torch.cat([volume, torch.zeros(depth-volume.shape[0], volume.shape[1], volume.shape[2])])\n        elif volume.shape[0] > depth:\n            volume = torch.nn.functional.interpolate(volume[None, None, ...], (depth, volume.shape[1], volume.shape[2])).squeeze()\n        return volume, coor_dict\n\n    def for_nfn(self, study_id):\n        meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Sagittal T1')]\n        meta = meta.sort_values('ipp_x', ascending=True).reset_index(drop=True)\n        #img = [self.normalize(self.load_dicom(f'/content/train_images/{row.study_id}/{row.series_id}/{row.instance_number}.dcm')) for _, row in meta.iterrows()]\n        coor = self.coor.loc[(self.coor.study_id==study_id)]\n        right_meta_list = []\n        left_meta_list = []\n        for _, row in coor.iterrows():\n            series_id, instance_number = row.series_id, row.instance_number\n            if row.condition == 'Right Neural Foraminal Narrowing':\n                right_meta_list.append(meta.loc[(meta.series_id==series_id) & (meta.instance_number==instance_number)])\n            else: \n                left_meta_list.append(meta.loc[(meta.series_id==series_id) & (meta.instance_number==instance_number)])\n\n        right_sub_meta = pd.concat(right_meta_list)\n        left_sub_meta = pd.concat(left_meta_list)\n        ridx = meta.loc[meta.ipp_x == right_sub_meta.ipp_x.median()].index[0]\n        lidx = meta.loc[meta.ipp_x == left_sub_meta.ipp_x.median()].index[0]\n        right_img_row = meta.iloc[min(max(ridx, 0), len(meta)-1)]\n        #display(right_img_row)\n        right_before_img_row = meta.iloc[min(max(ridx-1, 0), len(meta)-1)]\n        rightafter_img_row = meta.iloc[min(max(ridx+1, 0), len(meta)-1)]\n        left_img_row = meta.iloc[min(max(lidx, 0), len(meta)-1)]\n        left_before_img_row = meta.iloc[min(max(lidx-1, 0), len(meta)-1)]\n        leftafter_img_row = meta.iloc[min(max(lidx+1, 0), len(meta)-1)]\n        rimg = self.normalize(self.load_dicom(IMAGE_PATH + f'{right_img_row.study_id}/{right_img_row.series_id}/{right_img_row.instance_number}.dcm'))\n        rbimg = self.normalize(self.load_dicom(IMAGE_PATH + f'{right_before_img_row.study_id}/{right_before_img_row.series_id}/{right_before_img_row.instance_number}.dcm'))\n        raimg = self.normalize(self.load_dicom(IMAGE_PATH + f'{rightafter_img_row.study_id}/{rightafter_img_row.series_id}/{rightafter_img_row.instance_number}.dcm'))\n        limg = self.normalize(self.load_dicom(IMAGE_PATH + f'{left_img_row.study_id}/{left_img_row.series_id}/{left_img_row.instance_number}.dcm'))\n        lbimg = self.normalize(self.load_dicom(IMAGE_PATH + f'{left_before_img_row.study_id}/{left_before_img_row.series_id}/{left_before_img_row.instance_number}.dcm'))\n        laimg = self.normalize(self.load_dicom(IMAGE_PATH + f'{leftafter_img_row.study_id}/{leftafter_img_row.series_id}/{leftafter_img_row.instance_number}.dcm'))\n              \n        rimg = torch.cat([self.resize(torch.tensor(i)[None, ...]).to(torch.float32) for i in [rbimg, rimg, raimg]])\n        limg = torch.cat([self.resize(torch.tensor(i)[None, ...]).to(torch.float32) for i in [lbimg, limg, laimg]])\n        img = torch.stack([limg, rimg]).to(torch.float32).contiguous()\n        return img\n\n    def normalize(self, x):\n        lower, upper = np.percentile(x, (1, 99))\n        x = np.clip(x, lower, upper)\n        x = x - np.min(x)\n        x = x / np.max(x)\n        return x\n\n    def __len__(self):\n        return len(self.id)\n\n    def load_dicom(self, path):\n        dicom = dcm.dcmread(path)\n        data = dicom.pixel_array\n        return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:10:04.953654Z","iopub.execute_input":"2025-06-07T01:10:04.953900Z","iopub.status.idle":"2025-06-07T01:10:04.980154Z","shell.execute_reply.started":"2025-06-07T01:10:04.953883Z","shell.execute_reply":"2025-06-07T01:10:04.979422Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Coordinate prediction models","metadata":{}},{"cell_type":"code","source":"class ConvNextSCSDetect(nn.Module):\n    def __init__(self, encoder):\n        super().__init__()\n        #self.size = 384\n        if encoder == 'convnext': \n            self.encoder = timm.create_model('convnext_base.fb_in22k_ft_in1k_384', in_chans=3, pretrained=False, num_classes=0)\n        elif encoder == 'efficientnetv2-l': \n            self.encoder = timm.create_model('tf_efficientnetv2_l.in21k_ft_in1k', in_chans=3, pretrained=False, num_classes=0, drop_rate=0.)\n        self.in_features = self.encoder.num_features\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n                                    nn.Flatten(1),\n                                    #nn.LayerNorm(self.in_features)\n                                    )\n        self.l1 = nn.Linear(self.in_features, 2)\n        self.l2 = nn.Linear(self.in_features, 2)\n        self.l3 = nn.Linear(self.in_features, 2)\n        self.l4 = nn.Linear(self.in_features, 2)\n        self.l5 = nn.Linear(self.in_features, 2)\n    def forward(self, x, label=None):\n        #for loc, img in x.items():\n            #print(img.shape)\n        #    img = self.encoder.forward_features(img)\n        #    img = self.flatten(img)\n        #    x[loc] = img\n        x = self.encoder.forward_features(x)\n        x = self.flatten(x)\n        l1 = self.l1(x)\n        l2 = self.l2(x)\n        l3 = self.l3(x)\n        l4 = self.l4(x)\n        l5 = self.l5(x)\n        return {'L1/L2': l1.sigmoid(), 'L2/L3': l2.sigmoid(), 'L3/L4': l3.sigmoid(), 'L4/L5': l4.sigmoid(), 'L5/S1': l5.sigmoid()}\n\nclass ConvNextNFNDetect(nn.Module):\n    def __init__(self, encoder):\n        super().__init__()\n        if encoder == 'convnext': \n            self.encoder = timm.create_model('convnext_base.fb_in22k_ft_in1k_384', in_chans=3, pretrained=False, num_classes=0)\n        elif encoder == 'efficientnetv2-l': \n            self.encoder = timm.create_model('tf_efficientnetv2_l.in21k_ft_in1k', in_chans=3, pretrained=False, num_classes=0, drop_rate=0.)\n        self.in_features = self.encoder.num_features\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n                                    nn.Flatten(1),\n                                    #nn.LayerNorm(self.in_features)\n                                    )\n        self.ll1 = nn.Linear(self.in_features, 2)\n        self.ll2 = nn.Linear(self.in_features, 2)\n        self.ll3 = nn.Linear(self.in_features, 2)\n        self.ll4 = nn.Linear(self.in_features, 2)\n        self.ll5 = nn.Linear(self.in_features, 2)\n        self.rl1 = nn.Linear(self.in_features, 2)\n        self.rl2 = nn.Linear(self.in_features, 2)\n        self.rl3 = nn.Linear(self.in_features, 2)\n        self.rl4 = nn.Linear(self.in_features, 2)\n        self.rl5 = nn.Linear(self.in_features, 2)\n    def forward(self, x, label=None):\n        shape = x.shape\n        x = x.reshape(shape[0]*shape[1], 3, shape[-2], shape[-1])\n        x = self.encoder.forward_features(x)\n        x = self.flatten(x)\n        x = x.reshape(shape[0], shape[1], -1)\n        x_left = x[:, 0, :]\n        x_right = x[:, 1, :]\n        ll1 = self.ll1(x_left)\n        ll2 = self.ll2(x_left)\n        ll3 = self.ll3(x_left)\n        ll4 = self.ll4(x_left)\n        ll5 = self.ll5(x_left)\n        rl1 = self.rl1(x_right)\n        rl2 = self.rl2(x_right)\n        rl3 = self.rl3(x_right)\n        rl4 = self.rl4(x_right)\n        rl5 = self.rl5(x_right)\n        return {'left_L1/L2': ll1.sigmoid(),'left_L2/L3': ll2.sigmoid(),'left_L3/L4': ll3.sigmoid(), 'left_L4/L5': ll4.sigmoid(), 'left_L5/S1': ll5.sigmoid(),\n                'right_L1/L2': rl1.sigmoid(), 'right_L2/L3': rl2.sigmoid(), 'right_L3/L4': rl3.sigmoid(), 'right_L4/L5': rl4.sigmoid(), 'right_L5/S1': rl5.sigmoid()}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:10:04.980909Z","iopub.execute_input":"2025-06-07T01:10:04.981099Z","iopub.status.idle":"2025-06-07T01:10:05.001708Z","shell.execute_reply.started":"2025-06-07T01:10:04.981084Z","shell.execute_reply":"2025-06-07T01:10:05.000998Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Coordinate detection lightning module","metadata":{}},{"cell_type":"code","source":"class DetectModule(pl.LightningModule):\n    def __init__(self, condition, encoder):\n        super().__init__()\n        self.config = condition\n        if condition == 'scs':\n            self.model = ConvNextSCSDetect(encoder)\n        elif condition == 'nfn':\n            self.model = ConvNextNFNDetect(encoder)\n        elif  condition == 'ss': \n            pass\n        #self.ema = ExponentialMovingAverage(self.model.parameters(), decay=0.995)\n        #self.ema.to(device)\n\n        #self.model = torch.optim.swa_utils.AveragedModel(self.model,\n        #                                                 multi_avg_fn=torch.optim.swa_utils.get_ema_multi_avg_fn(0.999))\n\n    def forward(self, batch):\n        preds = self.model(batch)\n        return preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:10:05.005473Z","iopub.execute_input":"2025-06-07T01:10:05.005738Z","iopub.status.idle":"2025-06-07T01:10:05.026145Z","shell.execute_reply.started":"2025-06-07T01:10:05.005722Z","shell.execute_reply":"2025-06-07T01:10:05.025452Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Coordinate inference","metadata":{}},{"cell_type":"code","source":"%%time\nprefix = ''\nimport warnings\nwarnings.filterwarnings(\"ignore\")\ncoor_predict = {'scs': {\n                     'L1/L2':[], \n                     'L2/L3': [], \n                     'L3/L4': [], \n                     'L4/L5': [], \n                     'L5/S1': []\n                     }, \n                 'nfn': {\n                     'left_L1/L2': [], \n                     'left_L2/L3': [], \n                     'left_L3/L4': [], \n                     'left_L4/L5': [], \n                     'left_L5/S1': [], \n                     'right_L1/L2': [], \n                     'right_L2/L3': [], \n                     'right_L3/L4': [], \n                     'right_L4/L5': [], \n                     'right_L5/S1': [], \n                     }\n                    }\nmodel_path_dict = {\n    'scs': [\n        '/kaggle/input/rsna-spine-final-models/scs_detect_pre_0.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_detect_pre_1.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_detect_pre_2.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_detect_pre_3.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_detect_pre_4.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_detect_pre_effv2l_0.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_detect_pre_effv2l_1.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_detect_pre_effv2l_2.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_detect_pre_effv2l_3.ckpt', \n        '/kaggle/input/rsna-spine-final-models/scs_detect_pre_effv2l_4.ckpt', \n    ], \n    'nfn': [\n        '/kaggle/input/rsna-spine-final-models/nfn_detect_pre_0.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_detect_pre_1.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_detect_pre_2.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_detect_pre_3.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_detect_pre_4.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_detect_pre_effv2l_0.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_detect_pre_effv2l_1.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_detect_pre_effv2l_2.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_detect_pre_effv2l_3.ckpt', \n        '/kaggle/input/rsna-spine-final-models/nfn_detect_pre_effv2l_4.ckpt', \n    ]\n}\n##############COOR DETECT#########################\nfor condition in ['nfn', 'scs']:\n    print(condition)\n    model_path_list = model_path_dict[condition]\n    for path in model_path_list:\n        if 'effv2l' in path.split('/')[-1]: \n            encoder = 'efficientnetv2-l'\n        else: \n            encoder = 'convnext'\n        _meta_df = meta_df.copy()\n        _coor = pred_coor.copy()\n        dataset_test = CoorDetectDataset(_coor, _meta_df, condition, 'sub')\n        data_loader_test = DataLoader(\n            dataset_test,\n            batch_size=config[\"test_bs\"],\n            shuffle=False,\n            num_workers=4,\n            pin_memory=False\n        )\n        print(path, encoder)\n        model = DetectModule.load_from_checkpoint(path, condition=condition, encoder=encoder)\n        model.eval()\n        model.zero_grad()\n        model.to(device)\n\n        pred_temp = {}\n        for k in coor_predict[condition].keys(): \n            pred_temp[k] = []\n        study_id_list = []\n        with torch.no_grad():\n            for data in tqdm(data_loader_test, total=len(data_loader_test)):\n                images, study_id = data\n                images = images.to(device)\n                preds = model.forward(images)\n                #print(preds)\n                for k, v in preds.items(): \n                    pred_temp[k].append(v.to('cpu').detach().numpy())\n                study_id_list.append(study_id.to('cpu').reshape(-1).detach().numpy())\n        for k, v in pred_temp.items(): \n            coor_predict[condition][k].append(np.concatenate(v))\n        del pred_temp\n        gc.collect()\n        study_id = np.concatenate(study_id_list)\n    for k, v in coor_predict[condition].items(): \n        coor_predict[condition][k] = np.mean(np.array(coor_predict[condition][k]), axis=0)\n    coor_predict[condition]['study_id'] = study_id\n    del study_id\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:10:05.026867Z","iopub.execute_input":"2025-06-07T01:10:05.027060Z","iopub.status.idle":"2025-06-07T01:13:47.245271Z","shell.execute_reply.started":"2025-06-07T01:10:05.027038Z","shell.execute_reply":"2025-06-07T01:13:47.244458Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_coor.head(1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:13:47.246439Z","iopub.execute_input":"2025-06-07T01:13:47.246727Z","iopub.status.idle":"2025-06-07T01:13:47.254944Z","shell.execute_reply.started":"2025-06-07T01:13:47.246705Z","shell.execute_reply":"2025-06-07T01:13:47.254352Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\ndef create_label_coor(study_id, coor_df, coor, level, condition, desc): \n    _meta = meta_df.loc[meta_df.series_description==desc].copy()\n    _coor = coor_df.loc[coor_df.condition == condition]\n    _coor_df = {'study_id': [], 'series_id': [], 'x': [], 'y': []}\n    for s, c in zip(study_id, coor): \n        sub_meta = _meta.loc[_meta.study_id == s]\n        sub_coor = _coor.loc[(_coor.study_id==s) & (_coor.level==level.split('_')[-1])].squeeze(axis=0)\n        #display(sub_coor)\n        meta_row = sub_meta.loc[(sub_meta.instance_number==sub_coor.instance_number) & (sub_meta.series_id==sub_coor.series_id)].squeeze(axis=0)\n        x = round(meta_row.width * c[0])\n        y = round(meta_row.height * c[1])\n        _coor_df['study_id'].append(s)\n        _coor_df['series_id'].append(sub_coor.series_id)\n        _coor_df['x'].append(x)\n        _coor_df['y'].append(y)\n    _coor_df['level'] = level.split('_')[-1]\n    _coor_df['condition'] = condition\n    del _meta, _coor, sub_meta, sub_coor, meta_row\n    return pd.DataFrame(_coor_df)\n\nscs_study_id = coor_predict['scs']['study_id']\nscs_coor_list = []\nfor k, v in coor_predict['scs'].items(): \n    if k != 'study_id': \n        scs_coor_list.append(create_label_coor(scs_study_id, pred_coor, v, k, 'Spinal Canal Stenosis', 'Sagittal T2/STIR'))\n\nnfn_study_id = coor_predict['nfn']['study_id']\nnfn_coor_list = []\nfor k, v in coor_predict['nfn'].items(): \n    if k != 'study_id': \n        if k.split('_')[0] == 'left': \n            condition = 'Left Neural Foraminal Narrowing'\n        else: \n            condition = 'Right Neural Foraminal Narrowing'\n        nfn_coor_list.append(create_label_coor(nfn_study_id, pred_coor, v, k, condition, 'Sagittal T1'))\nscs_coor = pd.concat(scs_coor_list)\nnfn_coor = pd.concat(nfn_coor_list)\n_pred_coor = pd.concat([scs_coor, nfn_coor]).sort_values(['study_id', 'series_id', 'level'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:13:47.255699Z","iopub.execute_input":"2025-06-07T01:13:47.255951Z","iopub.status.idle":"2025-06-07T01:13:47.305306Z","shell.execute_reply.started":"2025-06-07T01:13:47.255926Z","shell.execute_reply":"2025-06-07T01:13:47.304671Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_coor_stage2 = pd.merge(pred_coor, _pred_coor, on=['study_id', 'series_id', 'level', 'condition'], how='inner')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:13:47.306032Z","iopub.execute_input":"2025-06-07T01:13:47.306280Z","iopub.status.idle":"2025-06-07T01:13:47.316312Z","shell.execute_reply.started":"2025-06-07T01:13:47.306260Z","shell.execute_reply":"2025-06-07T01:13:47.315668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"display(pred_coor_stage2.head())\npred_coor_stage2.to_csv('stage2_coor.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:13:47.316985Z","iopub.execute_input":"2025-06-07T01:13:47.317158Z","iopub.status.idle":"2025-06-07T01:13:47.325884Z","shell.execute_reply.started":"2025-06-07T01:13:47.317144Z","shell.execute_reply":"2025-06-07T01:13:47.325390Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.listdir('/kaggle/working/')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:13:47.326571Z","iopub.execute_input":"2025-06-07T01:13:47.326866Z","iopub.status.idle":"2025-06-07T01:13:47.338270Z","shell.execute_reply.started":"2025-06-07T01:13:47.326840Z","shell.execute_reply":"2025-06-07T01:13:47.337728Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Third Stage (calc. location of axial t2)¶\n* calcurate depth of axial t2 for each location roughly, using xyz-coordinate (refered to @hengck's transformation from sagittal t2 to axial t2)\n* roughly separate each locations\n* infer instance number\n* infer xy-coordinate","metadata":{}},{"cell_type":"markdown","source":"### Calculate axial slice","metadata":{}},{"cell_type":"code","source":"def project_to_3d(row):\n    sx, sy, sz = row.ipp_x, row.ipp_y, row.ipp_z\n    x, y = row.x, row.y\n    o0, o1, o2, o3, o4, o5 = row.iop\n    delx, dely = row.ps_x, row.ps_y\n    xx = o0 * delx * x + o3 * dely * y + sx\n    yy = o1 * delx * x + o4 * dely * y + sy\n    zz = o2 * delx * x + o5 * dely * y + sz\n    return xx,yy,zz\n\ndef sag_to_ax(sub_coor, sub_meta): \n    point = sub_coor[['ipp_x', 'ipp_y', 'ipp_z']].values #2d\n    level_list = sub_coor.level.tolist()\n    # here we project 2d to 3d\n    center=[] \n    for _, row in sub_coor.iterrows():\n        xx,yy,zz = project_to_3d(row)\n        center.append([xx,yy,zz])\n    center = np.array(center) #3d\n\n    # == 2. we get closest axial slices to the CSC points =================\n    #df = valid_data[0].axial_t2[0].df\n\n    orientation = np.array(sub_meta.iop.values.tolist())\n    position= np.array(sub_meta[['ipp_x', 'ipp_y', 'ipp_z']].values.tolist())\n    ox = orientation[:, :3]\n    oy = orientation[:, 3:]\n    oz = np.cross(ox,oy)\n    t = center.reshape(-1,1,3) - position.reshape(1,-1,3)\n    dis = (oz.reshape(1,-1,3) * t).sum(-1)  # np.dot(point-s,oz)\n    dis = np.fabs(dis)\n    closest = dis.argmin(-1)\n    closest_df = sub_meta.iloc[closest]\n    closest_df['level'] = level_list#['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']\n    closest_df['x'] = 0\n    closest_df['y'] = 0\n    #closest_df = pd.concat([closest_df, closest_df])\n    #closest_df['condition'] = ['Left Subarticular Stenosis']*5 + ['Right Subarticular Stenosis']*5\n    return closest_df[['study_id', 'series_id', 'instance_number', 'level']]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:13:47.338976Z","iopub.execute_input":"2025-06-07T01:13:47.339214Z","iopub.status.idle":"2025-06-07T01:13:47.350003Z","shell.execute_reply.started":"2025-06-07T01:13:47.339193Z","shell.execute_reply":"2025-06-07T01:13:47.349434Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# sagittal t2 => axial t2\nscs_coor = pred_coor_stage2.loc[pred_coor_stage2.condition=='Spinal Canal Stenosis'].copy()\nscs_coor = scs_coor.merge(meta_df, on=['study_id', 'series_id', 'instance_number'], how='left')\nstudy_id = scs_coor.study_id.unique()\nax_meta  = meta_df.loc[(meta_df.series_description=='Axial T2')]\nclosest_ax_list = []\nfor s in tqdm(study_id, total=len(study_id)): \n    sub_coor = scs_coor.loc[scs_coor.study_id==s]\n    sub_meta = ax_meta.loc[ax_meta.study_id==s]\n    closest_ax_list.append(sag_to_ax(sub_coor, sub_meta)) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:13:47.350781Z","iopub.execute_input":"2025-06-07T01:13:47.350988Z","iopub.status.idle":"2025-06-07T01:13:47.384257Z","shell.execute_reply.started":"2025-06-07T01:13:47.350973Z","shell.execute_reply":"2025-06-07T01:13:47.383545Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"closest_ax = pd.concat(closest_ax_list)\nclosest_ax.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:13:47.384992Z","iopub.execute_input":"2025-06-07T01:13:47.385182Z","iopub.status.idle":"2025-06-07T01:13:47.391940Z","shell.execute_reply.started":"2025-06-07T01:13:47.385167Z","shell.execute_reply":"2025-06-07T01:13:47.391265Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Subarticular stenosis (ss) - hẹp ống dưới coordinate prediction dataset","metadata":{}},{"cell_type":"code","source":"class SSDetectDataset(Dataset):\n    def __init__(self, ax, usage='train'):\n        self.ax = ax\n        self.id = ax.study_id.unique()\n        self.usage = usage\n        self.id = list(set(self.id) - set([3637444890]))\n        #self.id = [2773343225]\n        #self.id = [1782095928]\n\n        self.resize = v2.Resize((384, 384))\n        \n    def __getitem__(self, index):\n        study_id = self.id[index]\n        volume = self.for_ss(study_id)\n        return volume, torch.tensor(study_id)\n\n    def for_ss(self, study_id):\n        ax = self.ax.loc[self.ax.study_id==study_id]\n        img_dict = {}\n        for _, row in ax.iterrows():\n            series_id, instance_number = row.series_id, row.instance_number\n            img = self.load_dicom(IMAGE_PATH + f'{study_id}/{series_id}/{instance_number}.dcm').astype(np.float32)\n            img = self.resize(torch.tensor(img)[None, ...])\n            img = self.normalize(img)\n            img_dict[row.level] = img\n        img_list = []\n        for k in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']: \n            img_list.append(img_dict[k])\n        volume = torch.stack(img_list).contiguous()\n        return volume\n\n    def normalize(self, x):\n        upper = torch.quantile(x, torch.tensor([0.99]))\n        lower = torch.quantile(x, torch.tensor([0.01]))\n        x = torch.clip(x, lower, upper)\n        x = x - torch.min(x)\n        x = x / (torch.max(x)+1e-6)\n        return x\n\n    def __len__(self):\n        return len(self.id)\n\n    def load_dicom(self, path):\n        dicom = dcm.dcmread(path)\n        data = dicom.pixel_array\n        return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:13:47.392699Z","iopub.execute_input":"2025-06-07T01:13:47.392906Z","iopub.status.idle":"2025-06-07T01:13:47.405232Z","shell.execute_reply.started":"2025-06-07T01:13:47.392882Z","shell.execute_reply":"2025-06-07T01:13:47.404448Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## subarticular stenosis coordinate dectection model","metadata":{}},{"cell_type":"code","source":"class SSDetect(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = timm.create_model('convnext_base.fb_in22k_ft_in1k_384', in_chans=1, pretrained=False, num_classes=0)\n        self.in_features = self.encoder.num_features\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n                                    nn.Flatten(1),\n                                    #nn.LayerNorm(self.in_features)\n                                    )\n        self.left = nn.Linear(self.in_features, 2)\n        self.right = nn.Linear(self.in_features, 2)\n    def forward(self, x, label=None):\n        shape = x.shape\n        x = x.reshape(shape[0]*shape[1], 1, shape[-2], shape[-1])\n        x = self.encoder.forward_features(x)\n        x = self.flatten(x)\n        x = x.reshape(shape[0], shape[1], -1)\n        x_left = x\n        x_right = x\n        left = self.left(x_left)\n        right = self.right(x_right)\n        return {'left_L1/L2': left[:, 0, :].sigmoid(),'left_L2/L3': left[:, 1, :].sigmoid(),'left_L3/L4': left[:, 2, :].sigmoid(), 'left_L4/L5': left[:, 3, :].sigmoid(), 'left_L5/S1': left[:, 4, :].sigmoid(),\n                'right_L1/L2': right[:, 0, :].sigmoid(), 'right_L2/L3': right[:, 1, :].sigmoid(), 'right_L3/L4': right[:, 2, :].sigmoid(), 'right_L4/L5': right[:, 3, :].sigmoid(), 'right_L5/S1': right[:, 4, :].sigmoid()}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:13:47.406216Z","iopub.execute_input":"2025-06-07T01:13:47.406472Z","iopub.status.idle":"2025-06-07T01:13:47.420428Z","shell.execute_reply.started":"2025-06-07T01:13:47.406446Z","shell.execute_reply":"2025-06-07T01:13:47.419727Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Lightning module","metadata":{}},{"cell_type":"code","source":"class SSDetectModule(pl.LightningModule):\n    def __init__(self):\n        super().__init__()\n        self.model = SSDetect()\n    def forward(self, batch):\n        preds = self.model(batch)\n        return preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:13:47.421311Z","iopub.execute_input":"2025-06-07T01:13:47.421627Z","iopub.status.idle":"2025-06-07T01:13:47.434280Z","shell.execute_reply.started":"2025-06-07T01:13:47.421609Z","shell.execute_reply":"2025-06-07T01:13:47.433669Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## SS Coordinate inference","metadata":{}},{"cell_type":"code","source":"%%time\nprefix = ''\nimport warnings\nwarnings.filterwarnings(\"ignore\")\ncoor_predict = {\n             'left_L1/L2': [], \n             'left_L2/L3': [], \n             'left_L3/L4': [], \n             'left_L4/L5': [], \n             'left_L5/S1': [], \n             'right_L1/L2': [], \n             'right_L2/L3': [], \n             'right_L3/L4': [], \n             'right_L4/L5': [], \n             'right_L5/S1': [], \n                }\n\n##############COOR DETECT#########################\nfor i in [0, 1, 2, 3, 4]:\n    _meta_df = meta_df.copy()\n    _series = series.copy()\n    _coor = pred_coor.copy()\n    dataset_test = SSDetectDataset(closest_ax, 'sub')\n    data_loader_test = DataLoader(\n        dataset_test,\n        batch_size=config[\"test_bs\"],\n        shuffle=False,\n        num_workers=4,\n        pin_memory=False\n    )\n\n    model = SSDetectModule.load_from_checkpoint(f'/kaggle/input/rsna-spine-final-models/ss_detect_{i}.ckpt')\n    model.eval()\n    model.zero_grad()\n    model.to(device)\n\n    pred_temp = {}\n    for k in coor_predict.keys(): \n        pred_temp[k] = []\n    study_id_list = []\n    with torch.no_grad():\n        for data in tqdm(data_loader_test, total=len(data_loader_test)):\n            images, study_id = data\n            images = images.to(device)\n            preds = model.forward(images)\n            #print(preds)\n            for k, v in preds.items(): \n                pred_temp[k].append(v.to('cpu').detach().numpy())\n            study_id_list.append(study_id.to('cpu').reshape(-1).detach().numpy())\n    for k, v in pred_temp.items(): \n        coor_predict[k].append(np.concatenate(v))\n    del pred_temp\n    gc.collect()\n    study_id = np.concatenate(study_id_list)\nfor k, v in coor_predict.items(): \n    coor_predict[k] = np.mean(np.array(coor_predict[k]), axis=0)\ncoor_predict['study_id'] = study_id\ndel study_id\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:13:47.434846Z","iopub.execute_input":"2025-06-07T01:13:47.435014Z","iopub.status.idle":"2025-06-07T01:14:39.950366Z","shell.execute_reply.started":"2025-06-07T01:13:47.435001Z","shell.execute_reply":"2025-06-07T01:14:39.949559Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"study_id = coor_predict['study_id']\ncoor_dict = {'study_id': [], 'x': [], 'y': [], 'condition': [], 'level': []}\nfor k, v in coor_predict.items(): \n    if k == 'study_id': \n        continue\n    _lr, location = k.split('_')\n    if _lr == 'left': \n        lr = 'Left Subarticular Stenosis'\n    else: \n        lr = 'Right Subarticular Stenosis'\n    coor_dict['study_id'].extend(list(coor_predict['study_id']))\n    coor_dict['x'].extend(list(v[:, 0]))\n    coor_dict['y'].extend(list(v[:, 1]))\n    coor_dict['condition'].extend([lr]*len(v))\n    coor_dict['level'].extend([location]*len(v))\nax_coor_pred = pd.DataFrame(coor_dict)\nax_coor_pred = pd.merge(closest_ax, ax_coor_pred, on=['study_id', 'level'], how='left')\nprint(ax_coor_pred.shape)\nax_coor_pred.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:14:39.951364Z","iopub.execute_input":"2025-06-07T01:14:39.951665Z","iopub.status.idle":"2025-06-07T01:14:39.968402Z","shell.execute_reply.started":"2025-06-07T01:14:39.951638Z","shell.execute_reply":"2025-06-07T01:14:39.967565Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ax_coor_pred = ax_coor_pred.merge(meta_df[['study_id', 'series_id', 'instance_number', 'height', 'width']], on=['study_id', 'series_id', 'instance_number'], how='left')\nax_coor_pred['x'] = ax_coor_pred['x']*ax_coor_pred['width']\nax_coor_pred['y'] = ax_coor_pred['y']*ax_coor_pred['height']\nax_coor_pred['x'] = ax_coor_pred['x'].apply(lambda x: round(x))\nax_coor_pred['y'] = ax_coor_pred['y'].apply(lambda x: round(x))\nax_coor_pred = ax_coor_pred.drop(['height', 'width'], axis=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:14:39.969373Z","iopub.execute_input":"2025-06-07T01:14:39.969650Z","iopub.status.idle":"2025-06-07T01:14:39.990216Z","shell.execute_reply.started":"2025-06-07T01:14:39.969626Z","shell.execute_reply":"2025-06-07T01:14:39.989517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_coor_stage3 = pd.concat([pred_coor_stage2, ax_coor_pred])\ndisplay(pred_coor_stage3)\npred_coor_stage3.to_csv('stage3_coor.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:14:39.991008Z","iopub.execute_input":"2025-06-07T01:14:39.991276Z","iopub.status.idle":"2025-06-07T01:14:40.010464Z","shell.execute_reply.started":"2025-06-07T01:14:39.991251Z","shell.execute_reply":"2025-06-07T01:14:40.009801Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del pred_coor_stage2, ax_coor_pred, pred_coor, closest_ax\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:14:40.011284Z","iopub.execute_input":"2025-06-07T01:14:40.012097Z","iopub.status.idle":"2025-06-07T01:14:40.277555Z","shell.execute_reply.started":"2025-06-07T01:14:40.012074Z","shell.execute_reply":"2025-06-07T01:14:40.276673Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Stage 4: predict each severity","metadata":{}},{"cell_type":"markdown","source":"## Serverity prediction datasets","metadata":{}},{"cell_type":"code","source":"os.listdir('/kaggle/working/')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:14:40.278394Z","iopub.execute_input":"2025-06-07T01:14:40.278918Z","iopub.status.idle":"2025-06-07T01:14:40.291692Z","shell.execute_reply.started":"2025-06-07T01:14:40.278893Z","shell.execute_reply":"2025-06-07T01:14:40.291066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ClassDataset(Dataset):\n    def __init__(self, coor, meta, condition, channel, usage='sub'):\n        self.coor = coor\n        self.meta = meta\n        self.condition = condition\n        self.usage = usage\n        self.sag_window = channel\n        self.ax_window = channel\n        self.wide_resize = v2.Resize((128, 224))\n        #self.wide_resize = v2.Resize((224, 224))\n        self.rec_resize = v2.Resize((256, 256))\n        self.resize = v2.Resize((128, 128))\n        self.resize_3d = v2.Resize((256, 256))\n        self.pre_resize = v2.Resize((512, 512))\n        self.id = list(meta.study_id.unique())\n        if 3637444890 in self.id: \n            self.id.remove(3637444890)\n    def __getitem__(self, index):\n        study_id = self.id[index]\n        #print(study_id)\n        res = {}\n        #try:\n        if self.condition == 'scs':\n            sagt2_img, ax_img = self.for_scs(study_id)\n            res['sagt2'] = sagt2_img.to(torch.float32)\n            res['ax'] = ax_img.to(torch.float32)\n            #res['sagt1'] = sagt1_img.to(torch.float32)\n        elif self.condition == 'nfn':\n            ax_img, sagt1_img = self.for_nfn(study_id)\n            res['ax'] = ax_img.to(torch.float32)\n            res['sagt1'] = sagt1_img.to(torch.float32)\n        if self.condition == 'ss':\n            ax_img = self.for_ss(study_id)\n            #ax_img, sagt1_img, sagt2_img = self.for_ss(study_id)\n            res['ax'] = ax_img.to(torch.float32)\n        return res, torch.tensor(study_id)\n\n    def crop(self, image, x, y, z, x_left, x_right, y_bottom, y_top, wide):\n        size = [image[i].shape for i in z]\n        #print([self.pre_resize(torch.tensor(image[i])[None, ...]).squeeze() for i, shape in zip(z, size)][0].shape)\n        data = torch.stack([torch.tensor(self.pre_resize(torch.tensor(image[i])[None, ...]).squeeze()[max(int((y/shape[0])*512-y_top), 0):int((y/shape[0])*512+y_bottom), max(int((x/shape[1])*512-x_left), 0): int((x/shape[1])*512+x_right)]) for i, shape in zip(z, size)])\n\n        if wide:\n            data = self.wide_resize(data)\n        else:\n            data = self.rec_resize(data)\n\n        return data\n    def for_scs(self, study_id):\n        sagt2_meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Sagittal T2/STIR')]\n        #display(sagt2_meta)\n        ax_meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Axial T2')]\n        sagt1_meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Sagittal T1')]\n        #display(ax_meta)\n        sagt2_meta = sagt2_meta.sort_values('ipp_x', ascending=True).reset_index(drop=True)\n        ax_meta = ax_meta.sort_values('ipp_z', ascending=False).reset_index(drop=True)\n        sagt1_meta = sagt1_meta.sort_values('ipp_x', ascending=True).reset_index(drop=True)\n        sagt2_img = [self.normalize(self.load_dicom(IMAGE_PATH + f'{row.study_id}/{row.series_id}/{row.instance_number}.dcm')) for _, row in sagt2_meta.iterrows()]\n        ax_img = [self.normalize(self.load_dicom(IMAGE_PATH + f'{row.study_id}/{row.series_id}/{row.instance_number}.dcm')) for _, row in ax_meta.iterrows()]\n        sagt1_img = [self.normalize(self.load_dicom(IMAGE_PATH + f'{row.study_id}/{row.series_id}/{row.instance_number}.dcm')) for _, row in sagt1_meta.iterrows()]\n        sagt1_img = [img if (img.shape[0]> 1 and img.shape[1] > 1) else np.zeros((512, 512)) for img in sagt1_img]\n        sagt2_sub_coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Spinal Canal Stenosis')]\n        ax_right_sub_coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Right Subarticular Stenosis')]\n        ax_left_sub_coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Left Subarticular Stenosis')]\n        sagt2_dict = {}\n        for _, row in sagt2_sub_coor.iterrows():\n            try: \n                y_shift = 0\n                x_shift = 0\n                z_shift = 0\n                #display(sagt2_meta)\n                #display(row)\n                mid = sagt2_meta.loc[(sagt2_meta.series_id==row.series_id)&(sagt2_meta.instance_number==row.instance_number)].index[0]\n                if row.level == 'L5/S1':\n                    ushift = 20\n                else:\n                    ushift = 0\n                z = [min(max(mid+w+z_shift, 0), len(sagt2_meta)-1) for w in range(-(self.sag_window-1)//2, ((self.sag_window-1)//2)+1)]\n                sagt2_dict[row.level] = self.crop(sagt2_img, row.x+x_shift, row.y+y_shift, z, 96, 32, 40+ushift, 40-ushift, wide=True)\n            except: \n                pass\n                \n        # AXIAL T2\n        #in_list = ax_meta.instance_number.tolist()\n        ax_dict = {}\n        if np.random.choice([0, 1]) == 0:\n            ax_sub_coor = ax_right_sub_coor\n            lrshift = +20\n        else:\n            ax_sub_coor = ax_left_sub_coor\n            lrshift = -20\n        for _, row in ax_sub_coor.iterrows():\n            try: \n                y_shift = 0\n                x_shift = 0\n                z_shift = 0\n                ax_meta_sub = ax_meta.loc[(ax_meta.series_id==row.series_id)]\n                ax_meta_sub_original_idx = ax_meta_sub.index.tolist()\n                ax_meta_sub = ax_meta_sub.reset_index(drop=True)\n                mid = ax_meta_sub.loc[(ax_meta_sub.instance_number==row.instance_number)].index[0]\n                z = [min(max(mid+w+z_shift, 0), len(ax_meta_sub)-1) for w in range(-(self.ax_window-1)//2, ((self.ax_window-1)//2)+1)]\n                ax_dict[row.level] = self.crop([ax_img[i] for i in range(len(ax_img)) if i in ax_meta_sub_original_idx], row.x+x_shift+lrshift, row.y+y_shift, z, 96, 96, 96, 96, wide=False)\n            except: \n                pass\n        sagt2_img = [sagt2_dict.get(l, torch.zeros((self.sag_window, 128, 224))) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']]\n        ax_img = [ax_dict.get(l, torch.zeros((self.ax_window, 256, 256))) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']]\n        return torch.stack(sagt2_img).contiguous(), torch.stack(ax_img).contiguous()#, torch.stack(sagt1_img).contiguous()\n    def for_ss(self, study_id):\n        #display(sagt2_meta)\n        ax_meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Axial T2')]\n        sagt1_meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Sagittal T1')]\n        sagt2_meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Sagittal T2/STIR')]\n        #display(ax_meta)\n        ax_meta = ax_meta.sort_values('ipp_z', ascending=False).reset_index(drop=True)\n        ax_img = [self.normalize(self.load_dicom(IMAGE_PATH + f'{row.study_id}/{row.series_id}/{row.instance_number}.dcm')) for _, row in ax_meta.iterrows()]\n        ax_right_sub_coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Right Subarticular Stenosis')]\n        ax_left_sub_coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Left Subarticular Stenosis')]\n\n        # AXIAL T2\n        #in_list = ax_meta.instance_number.tolist()\n        ax_right_dict = {}\n        for _, row in ax_right_sub_coor.iterrows():\n            try: \n                y_shift = 0\n                x_shift = 0\n                z_shift = 0\n                ax_meta_sub = ax_meta.loc[(ax_meta.series_id==row.series_id)]\n                ax_meta_sub_original_idx = ax_meta_sub.index.tolist()\n                ax_meta_sub = ax_meta_sub.reset_index(drop=True)\n                mid = ax_meta_sub.loc[(ax_meta_sub.instance_number==row.instance_number)].index[0]\n                z = [min(max(mid+w+z_shift, 0), len(ax_meta_sub)-1) for w in range(-(self.ax_window-1)//2, ((self.ax_window-1)//2)+1)]\n                ax_right_dict[row.level] = self.crop([ax_img[i] for i in range(len(ax_img)) if i in ax_meta_sub_original_idx], row.x+x_shift, row.y+y_shift, z, 160-16, 32+16, 64+32, 64+32, wide=False)\n            except: \n                pass\n        ax_left_dict = {}\n        for _, row in ax_left_sub_coor.iterrows():\n            try: \n                y_shift = 0\n                x_shift = 0\n                z_shift = 0\n                ax_meta_sub = ax_meta.loc[(ax_meta.series_id==row.series_id)]\n                ax_meta_sub_original_idx = ax_meta_sub.index.tolist()\n                ax_meta_sub = ax_meta_sub.reset_index(drop=True)\n                mid = ax_meta_sub.loc[(ax_meta_sub.instance_number==row.instance_number)].index[0]\n                z = [min(max(mid+w+z_shift, 0), len(ax_meta_sub)-1) for w in range(-(self.ax_window-1)//2, ((self.ax_window-1)//2)+1)]\n                ax_left_dict[row.level] = self.crop([ax_img[i] for i in range(len(ax_img)) if i in ax_meta_sub_original_idx], row.x+x_shift, row.y+y_shift, z, 32+16, 160-16, 64+32, 64+32, wide=False)\n            except: \n                pass\n        ax_right_img = [ax_right_dict.get(l, torch.zeros((self.ax_window, 256, 256))) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']]\n        ax_left_img = [ax_left_dict.get(l, torch.zeros((self.ax_window, 256, 256))) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']]\n        ax_img = ax_left_img + ax_right_img\n        return torch.stack(ax_img).contiguous()#, torch.stack(sagt1_img).contiguous(), torch.stack(sagt2_img).contiguous()\n\n    def for_nfn(self, study_id):\n        ax_meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Axial T2')]\n        sagt1_meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Sagittal T1')]\n\n        ax_meta = ax_meta.sort_values('ipp_z', ascending=False).reset_index(drop=True)\n        sagt1_meta = sagt1_meta.sort_values('ipp_x', ascending=True).reset_index(drop=True)\n        ax_img = [self.normalize(self.load_dicom(IMAGE_PATH + f'{row.study_id}/{row.series_id}/{row.instance_number}.dcm')) for _, row in ax_meta.iterrows()]\n        sagt1_img = [self.normalize(self.load_dicom(IMAGE_PATH + f'{row.study_id}/{row.series_id}/{row.instance_number}.dcm')) for _, row in sagt1_meta.iterrows()]\n        ax_right_sub_coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Right Subarticular Stenosis')]\n        ax_left_sub_coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Left Subarticular Stenosis')]\n        sagt1_right_sub_coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Right Neural Foraminal Narrowing')]\n        sagt1_left_sub_coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Left Neural Foraminal Narrowing')]\n\n        # SAGITTAL T2\n        # not implemented\n\n        # AXIAL T2\n        #in_list = ax_meta.instance_number.tolist()\n        ax_right_dict = {}\n        for _, row in ax_right_sub_coor.iterrows():\n            try: \n                y_shift = 0\n                x_shift = 0\n                z_shift = 0\n                ax_meta_sub = ax_meta.loc[(ax_meta.series_id==row.series_id)]\n                ax_meta_sub_original_idx = ax_meta_sub.index.tolist()\n                ax_meta_sub = ax_meta_sub.reset_index(drop=True)\n                mid = ax_meta_sub.loc[(ax_meta_sub.instance_number==row.instance_number)].index[0]\n                z = [min(max(mid+w+z_shift, 0), len(ax_meta_sub)-1) for w in range(-(self.ax_window-1)//2, ((self.ax_window-1)//2)+1)]\n                ax_right_dict[row.level] = self.crop([ax_img[i] for i in range(len(ax_img)) if i in ax_meta_sub_original_idx], row.x+x_shift, row.y+y_shift, z, 160-16, 32+16, 64+32, 64+32, wide=False)\n            except: \n                pass\n        ax_left_dict = {}\n        for _, row in ax_left_sub_coor.iterrows():\n            try: \n                y_shift = 0\n                x_shift = 0\n                z_shift = 0\n                ax_meta_sub = ax_meta.loc[(ax_meta.series_id==row.series_id)]\n                ax_meta_sub_original_idx = ax_meta_sub.index.tolist()\n                ax_meta_sub = ax_meta_sub.reset_index(drop=True)\n                mid = ax_meta_sub.loc[(ax_meta_sub.instance_number==row.instance_number)].index[0]\n                z = [min(max(mid+w+z_shift, 0), len(ax_meta_sub)-1) for w in range(-(self.ax_window-1)//2, ((self.ax_window-1)//2)+1)]\n                ax_left_dict[row.level] = self.crop([ax_img[i] for i in range(len(ax_img)) if i in ax_meta_sub_original_idx], row.x+x_shift, row.y+y_shift, z, 32+16, 160-16, 64+32, 64+32, wide=False)\n            except: \n                pass\n        # SAGITTAL T1\n        sagt1_right_dict = {}\n        for _, row in sagt1_right_sub_coor.iterrows():\n            try: \n                y_shift = 0\n                x_shift = 0\n                z_shift = 0\n                if row.level == 'L5/S1':\n                    ushift = 10\n                else:\n                    ushift = 0\n                sagt1_meta_sub = sagt1_meta.loc[(sagt1_meta.series_id==row.series_id)]\n                sagt1_meta_sub_original_idx = sagt1_meta_sub.index.tolist()\n                sagt1_meta_sub = sagt1_meta_sub.reset_index(drop=True)\n                #display(sagt2_meta)\n                #display(row)\n                mid = sagt1_meta_sub.loc[(sagt1_meta_sub.instance_number==row.instance_number)].index[0]\n                z = [min(max(mid+w+z_shift, 0), len(sagt1_meta_sub)-1) for w in range(-(self.sag_window-1)//2, ((self.sag_window-1)//2)+1)]\n                sagt1_right_dict[row.level] = self.crop([sagt1_img[i] for i in range(len(sagt1_img)) if i in sagt1_meta_sub_original_idx], row.x+x_shift, row.y+y_shift, z, 96, 64, 32+ushift, 32-ushift, wide=True)\n            except: \n                pass\n        sagt1_left_dict = {}\n        for _, row in sagt1_left_sub_coor.iterrows():\n            try: \n                y_shift = 0\n                x_shift = 0\n                z_shift = 0\n                if row.level == 'L5/S1':\n                    ushift = 10\n                else:\n                    ushift = 0\n                sagt1_meta_sub = sagt1_meta.loc[(sagt1_meta.series_id==row.series_id)]\n                sagt1_meta_sub_original_idx = sagt1_meta_sub.index.tolist()\n                sagt1_meta_sub = sagt1_meta_sub.reset_index(drop=True)\n                mid = sagt1_meta_sub.loc[(sagt1_meta_sub.instance_number==row.instance_number)].index[0]\n                z = [min(max(mid+w+z_shift, 0), len(sagt1_meta_sub)-1) for w in range(-(self.sag_window-1)//2, ((self.sag_window-1)//2)+1)]\n                sagt1_left_dict[row.level] = self.crop([sagt1_img[i] for i in range(len(sagt1_img)) if i in sagt1_meta_sub_original_idx], row.x+x_shift, row.y+y_shift, z, 96, 64, 32+ushift, 32-ushift, wide=True)\n            except: \n                pass\n        ax_right_img = [ax_right_dict.get(l, torch.zeros((self.ax_window, 256, 256))) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']]\n        ax_left_img = [ax_left_dict.get(l, torch.zeros((self.ax_window, 256, 256))) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']]\n        sagt1_right_img = [sagt1_right_dict.get(l, torch.zeros((self.sag_window, 128, 224))) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']]\n        sagt1_left_img = [sagt1_left_dict.get(l, torch.zeros((self.sag_window, 128, 224))) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']]\n        ax_img = ax_left_img + ax_right_img\n        sagt1_img = sagt1_left_img + sagt1_right_img\n        return torch.stack(ax_img).contiguous(), torch.stack(sagt1_img).contiguous()\n\n\n    def normalize(self, x):\n        lower, upper = np.percentile(x, (1, 99))\n        x = np.clip(x, lower, upper)\n        x = x - np.min(x)\n        x = x / np.max(x)\n        return x\n\n    def __len__(self):\n        return len(self.id)\n\n    def load_dicom(self, path):\n        dicom = dcm.dcmread(path)\n        data = dicom.pixel_array\n        return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:14:40.292602Z","iopub.execute_input":"2025-06-07T01:14:40.292805Z","iopub.status.idle":"2025-06-07T01:14:40.336133Z","shell.execute_reply.started":"2025-06-07T01:14:40.292790Z","shell.execute_reply":"2025-06-07T01:14:40.335542Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Severity predict models","metadata":{}},{"cell_type":"code","source":"class Flatten(nn.Sequential):\n    def __init__(self):\n        super().__init__(\n            nn.AdaptiveAvgPool2d((1, 1)),\n            nn.Flatten(1),\n            #nn.LayerNorm(512)\n        )\n        \nclass ConConvnextSCS(nn.Module):\n    def __init__(self, direction='sagt2'):\n        super().__init__()\n        self.direction = direction\n        self.ax = timm.create_model('convnext_base.fb_in22k_ft_in1k_384', in_chans=3, pretrained=False, num_classes=0)\n        self.sagt2 = timm.create_model('convnext_base.fb_in22k_ft_in1k_384', in_chans=3, pretrained=False, num_classes=0)\n        self.num_features = self.ax.num_features\n        self.flatten_ax = Flatten()\n        self.flatten_sagt2 = Flatten()\n        self.lin_ax = nn.Linear(self.num_features, 512)\n        self.lin_sagt2 = nn.Linear(self.num_features, 512)\n        self.aux_ax = nn.Linear(512, 3)\n        self.aux_sagt2 = nn.Linear(512, 3)\n        self.lin = nn.Linear(512*2, 512)\n        self.out = nn.Linear(512, 3)\n        self.dropout = nn.Dropout(0.0)\n    def forward(self, ax, sagt2, sagt1=None, label=None):\n        shape = ax.shape\n        ax = ax.reshape(shape[0]*shape[1], 3, shape[-2], shape[-1])\n        shape=sagt2.shape\n        sagt2 = sagt2.reshape(shape[0]*shape[1], 3, shape[-2], shape[-1])\n        ax = nn.functional.leaky_relu(self.lin_ax(self.flatten_ax(self.ax.forward_features(ax))))\n        sagt2 = nn.functional.leaky_relu(self.lin_sagt2(self.flatten_sagt2(self.sagt2.forward_features(sagt2))))\n        x = torch.cat([ax, sagt2], dim=1)\n        x = self.lin(x)\n        x = nn.functional.leaky_relu(x)\n        x = self.dropout(x)\n        x = self.out(x)\n        return x\n\nclass ConConvnextNFN(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.ax = timm.create_model('convnext_base.fb_in22k_ft_in1k_384', in_chans=3, pretrained=False, num_classes=0)\n        self.sagt1 = timm.create_model('convnext_base.fb_in22k_ft_in1k_384', in_chans=3, pretrained=False, num_classes=0)\n        self.num_features = self.ax.num_features\n        self.flatten_ax = Flatten()\n        self.flatten_sagt1 = Flatten()\n        self.lin_ax = nn.Linear(self.num_features, 512)\n        self.lin_sagt1 = nn.Linear(self.num_features, 512)\n        self.aux_ax = nn.Linear(512, 3)\n        self.aux_sagt1 = nn.Linear(512, 3)\n        self.lin = nn.Linear(512*2, 512)\n        self.out = nn.Linear(512, 3)\n        self.dropout = nn.Dropout(0.0)\n    def forward(self, ax, sagt1, sagt2=None, label=None):\n        shape = ax.shape\n        ax = ax.reshape(shape[0]*shape[1], 3, shape[-2], shape[-1])\n        shape=sagt1.shape\n        sagt1 = sagt1.reshape(shape[0]*shape[1], 3, shape[-2], shape[-1])\n        ax = nn.functional.leaky_relu(self.lin_ax(self.flatten_ax(self.ax.forward_features(ax))))\n        sagt1 = nn.functional.leaky_relu(self.lin_sagt1(self.flatten_sagt1(self.sagt1.forward_features(sagt1))))\n        x = torch.cat([ax, sagt1], dim=1)\n        x = self.lin(x)\n        x = nn.functional.leaky_relu(x)\n        x = self.dropout(x)\n        x = self.out(x)\n        return x\n    \nclass ConvnextSS(nn.Module):\n    def __init__(self, direction='ax'):\n        super().__init__()\n        self.direction = direction\n        self.encoder = timm.create_model('convnext_large.fb_in22k_ft_in1k_384', in_chans=3, pretrained=False, num_classes=0)\n        self.in_features = self.encoder.num_features\n        self.flatten = Flatten()\n        self.lin = nn.Linear(self.in_features, 512)\n        self.out = nn.Linear(512, 3)\n        self.dropout = nn.Dropout(0.0)\n    def forward(self, ax, sagt1=None, sagt2=None, label=None):\n        if self.direction == 'ax':\n            shape = ax.shape\n            x = ax.reshape(shape[0]*shape[1], 3, shape[-2], shape[-1])\n        elif self.direction == 'sagt1':\n            shape = sagt1.shape\n            x = sagt1.reshape(shape[0]*shape[1], 3, shape[-2], shape[-1])\n        elif self.direction == 'sagt2':\n            shape = sagt2.shape\n            x = sagt2.reshape(shape[0]*shape[1], 3, shape[-2], shape[-1])\n        x = self.flatten(self.encoder.forward_features(x))\n        x = self.lin(x)\n        x = nn.functional.leaky_relu(x)\n        x = self.dropout(x)\n        x = self.out(x)\n        return x#, ax, sagt1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:14:40.336766Z","iopub.execute_input":"2025-06-07T01:14:40.336946Z","iopub.status.idle":"2025-06-07T01:14:40.352687Z","shell.execute_reply.started":"2025-06-07T01:14:40.336933Z","shell.execute_reply":"2025-06-07T01:14:40.351949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AttentionMIL(nn.Module):\n    def __init__(self, input_dim, hidden_dim, num_classes, condition):\n        super(AttentionMIL, self).__init__()\n        self.condition = condition\n        if condition == 'nfn': \n            self.attention = nn.Sequential(\n            nn.Linear(input_dim, hidden_dim),\n            nn.Tanh(),\n            nn.Linear(hidden_dim, 1)\n            )\n        else: \n            self.lin = nn.Linear(input_dim, hidden_dim)\n            self.attn_score = nn.Linear(hidden_dim, 1)\n            self.act = nn.Tanh()\n    def forward(self, bags):\n        \"\"\"\n        Args:\n            bags: (batch_size, num_instances, input_dim)\n\n        Returns:\n            logits: (batch_size, num_classes)\n        \"\"\"\n        batch_size, num_instances, input_dim = bags.size()\n\n        # Attention mechanism\n        if self.condition=='nfn': \n            attn_scores = self.attention(bags).squeeze(-1)  # (batch_size, num_instances)\n        else: \n            x = self.lin(bags)\n            attn_scores = self.attn_score(self.act(x)).squeeze(-1)\n        attn_weights = torch.softmax(attn_scores, dim=-1)  # (batch_size, num_instances)\n        # Weighted sum of instances\n        weighted_instances = torch.bmm(attn_weights.unsqueeze(1), bags).squeeze(1)  # (batch_size, input_dim)\n\n        # Classification\n        #logits = self.classifier(weighted_instances)\n        return weighted_instances, attn_scores\nclass SelfAttentionMIL(nn.Module):\n    def __init__(self, input_dim, hidden_dim, num_classes, is_layer_norm=False):\n        super(SelfAttentionMIL, self).__init__()\n        self.is_layer_norm = is_layer_norm\n\n        # Self-Attention層\n        self.self_attn = nn.MultiheadAttention(input_dim, num_heads=8, batch_first=True)\n\n\n        # バッグレベルの分類器\n        #self.bag_classifier = nn.Sequential(\n        self.layer_norm = nn.LayerNorm(input_dim)\n        self.dropout = nn.Dropout(p=0.0)\n        self.lin = nn.Linear(input_dim, hidden_dim)\n        self.act = nn.Tanh()\n        self.calc_attn_score = nn.Linear(hidden_dim, 1)  # バッグレベルのスコア\n        #)\n\n    def forward(self, bags):\n        # Self-Attention\n        attn_output, _ = self.self_attn(bags, bags, bags)\n        x = attn_output + bags\n        if self.is_layer_norm: \n            x = self.layer_norm(x)\n        # バッグレベルのAttentionスコアを計算\n        #bag_attn_scores = self.bag_classifier(attn_output).squeeze(-1)\n        x = self.lin(x)\n        bag_attn_scores = self.calc_attn_score(self.act(x)).squeeze(-1)\n        bag_attn_weights = torch.softmax(bag_attn_scores, dim=-1)\n\n        # Attention重み付き平均でバッグレベルの特徴量を計算\n        bag_features = torch.bmm(bag_attn_weights.unsqueeze(1), attn_output).squeeze(1)\n\n        return bag_features, bag_attn_scores\n    \nclass LSTMMIL(nn.Module):\n    def __init__(self, input_dim, hidden_dim, num_classes):\n        super(LSTMMIL, self).__init__()\n        #self.attention = nn.Sequential(\n        #    nn.Linear(input_dim, hidden_dim),\n        #    nn.Tanh(),\n        #    nn.Linear(hidden_dim, 1)\n        #)\n        #self.classifier = nn.Linear(input_dim, num_classes)\n        self.lstm = nn.LSTM(input_dim, input_dim//2, num_layers=2, batch_first=True, dropout=0.1, bidirectional=True)\n        #self.lin = nn.Linear(input_dim, hidden_dim)\n        self.aux_attention = nn.Sequential(\n            nn.Tanh(),\n            nn.Linear(input_dim, 1)\n        )\n        self.attention = nn.Sequential(\n            nn.Tanh(),\n            nn.Linear(input_dim, 1)\n        )\n    def forward(self, bags):\n        \"\"\"\n        Args:\n            bags: (batch_size, num_instances, input_dim)\n\n        Returns:\n            logits: (batch_size, num_classes)\n        \"\"\"\n        batch_size, num_instances, input_dim = bags.size()\n\n        # Attention mechanism\n        #attn_scores = self.attention(bags).squeeze(-1)  # (batch_size, num_instances)\n        bags_lstm, _ = self.lstm(bags)\n        attn_scores = self.attention(bags_lstm).squeeze(-1)\n        aux_attn_scores = self.aux_attention(bags_lstm).squeeze(-1)\n        attn_weights = torch.softmax(attn_scores, dim=-1)  # (batch_size, num_instances)\n        #aux_attn_weights = torch.softmax(aux_attn_scores, dim=-1)\n        # Weighted sum of instances\n        weighted_instances = torch.bmm(attn_weights.unsqueeze(1), bags_lstm).squeeze(1)  # (batch_size, input_dim)\n        # Classification\n        #logits = self.classifier(weighted_instances)\n        return weighted_instances, aux_attn_scores\n    \nclass SCSMIL(nn.Module):\n    def __init__(self, model_name):\n        super().__init__()\n        if 'convnext' in model_name: \n            self.sagt2_encoder = timm.create_model('convnext_small.fb_in22k_ft_in1k_384', in_chans=1, pretrained=False, num_classes=0)\n            self.ax_encoder = timm.create_model('convnext_small.fb_in22k_ft_in1k_384', in_chans=1, pretrained=False, num_classes=0)\n        elif 'effv2s' in model_name: \n            self.sagt2_encoder = timm.create_model('tf_efficientnetv2_s.in21k_ft_in1k', in_chans=1, pretrained=False, num_classes=0)\n            self.ax_encoder = timm.create_model('tf_efficientnetv2_s.in21k_ft_in1k', in_chans=1, pretrained=False, num_classes=0)\n        self.sagt2_flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n                                    nn.Flatten(1))\n        #self.sagt1_flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n        #                            nn.Flatten(1))\n        self.ax_flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n                                    nn.Flatten(1))\n        self.sagt2_num_features = self.sagt2_encoder.num_features\n        #self.sagt1_num_features = self.sagt1_encoder.num_features\n        self.ax_num_features = self.ax_encoder.num_features\n        self.sagt2_head = LSTMMIL(self.sagt2_num_features, 512, 3)\n        #self.sagt1_head = AttentionMIL(self.sagt1_num_features, 512, 3)\n        self.ax_head = LSTMMIL(self.ax_num_features, 512, 3)\n        \n        self.out = nn.Linear(self.sagt2_num_features + self.ax_num_features, 3)\n        self.aux_out = nn.Linear(self.sagt2_num_features, 3)\n        self.dropout = nn.Dropout(0.0)\n    def forward(self, ax, sagt2, sagt1=None):\n        ax_shape = ax.shape\n        ax = ax.reshape(ax_shape[0]*ax_shape[1]*ax_shape[2], 1, ax_shape[-2], ax_shape[-1])\n        ax = self.ax_encoder.forward_features(ax)\n        ax = self.ax_flatten(ax)\n        ax = ax.reshape(ax_shape[0]*ax_shape[1], ax_shape[2], -1)\n        ax_weighted_sum, ax_attn = self.ax_head(ax)\n        ax_attn = ax_attn.reshape(ax_shape[0], ax_shape[1], -1)\n\n        sagt2_shape = sagt2.shape\n        sagt2 = sagt2.reshape(sagt2_shape[0]*sagt2_shape[1]*sagt2_shape[2], 1, sagt2_shape[-2], sagt2_shape[-1])\n        sagt2 = self.sagt2_encoder.forward_features(sagt2)\n        sagt2 = self.sagt2_flatten(sagt2)\n        sagt2 = sagt2.reshape(sagt2_shape[0]*sagt2_shape[1], sagt2_shape[2], -1)\n        sagt2_weighted_sum, sagt2_attn = self.sagt2_head(sagt2)\n        sagt2_attn = sagt2_attn.reshape(sagt2_shape[0], sagt2_shape[1], -1)\n\n        out = torch.cat([ax_weighted_sum, sagt2_weighted_sum], dim=1)\n        out = self.out(out)\n        sagt2_out = self.aux_out(sagt2_weighted_sum)\n        ax_out = self.aux_out(ax_weighted_sum)\n        #print(sagt2_attn.shape, ax_attn.shape)\n        ax_attn = {'L1/L2': ax_attn[:, 0, :], 'L2/L3': ax_attn[:, 1, :], 'L3/L4': ax_attn[:, 2, :], 'L4/L5': ax_attn[:, 3, :], 'L5/S1': ax_attn[:, 4, :]}\n        sagt2_attn = {'L1/L2': sagt2_attn[:, 0, :], 'L2/L3': sagt2_attn[:, 1, :], 'L3/L4': sagt2_attn[:, 2, :], 'L4/L5': sagt2_attn[:, 3, :], 'L5/S1': sagt2_attn[:, 4, :]}\n        #sagt1_attn = {'L1/L2': sagt1_attn[:, 0, :], 'L2/L3': sagt1_attn[:, 1, :], 'L3/L4': sagt1_attn[:, 2, :], 'L4/L5': sagt1_attn[:, 3, :], 'L5/S1': sagt1_attn[:, 4, :]}\n        return out\n\n\nclass NFNMIL(nn.Module):\n    def __init__(self, model_name):\n        super().__init__()\n        if 'convnext' in model_name: \n            self.sagt1_encoder = timm.create_model('convnext_small.fb_in22k_ft_in1k_384', in_chans=1, pretrained=False, num_classes=0)\n            self.ax_encoder = timm.create_model('convnext_small.fb_in22k_ft_in1k_384', in_chans=1, pretrained=False, num_classes=0)\n        elif 'effv2s' in model_name: \n            self.sagt1_encoder = timm.create_model('tf_efficientnetv2_s.in21k_ft_in1k', in_chans=1, pretrained=False, num_classes=0)\n            self.ax_encoder = timm.create_model('tf_efficientnetv2_s.in21k_ft_in1k', in_chans=1, pretrained=False, num_classes=0)\n        self.sagt1_flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n                                    nn.Flatten(1))\n        self.ax_flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n                                    nn.Flatten(1))\n        self.sagt1_num_features = self.sagt1_encoder.num_features\n        self.ax_num_features = self.ax_encoder.num_features\n        self.sagt1_head = LSTMMIL(self.sagt1_num_features, 512, 3)\n        self.ax_head = LSTMMIL(self.ax_num_features, 512, 3)\n        self.out = nn.Linear(self.sagt1_num_features+self.ax_num_features, 3)\n        self.aux_out = nn.Linear(self.sagt1_num_features, 3)\n        self.dropout = nn.Dropout(0.0)\n    def forward(self, ax, sagt1):\n        ax_shape = ax.shape\n        #print(ax.shape, sagt2.shape)\n        ax = ax.reshape(ax_shape[0]*ax_shape[1]*ax_shape[2], 1, ax_shape[-2], ax_shape[-1])\n        ax = self.ax_encoder.forward_features(ax)\n        ax = self.ax_flatten(ax)\n        ax = ax.reshape(ax_shape[0]*ax_shape[1], ax_shape[2], -1)\n        ax_weighted_sum, ax_attn = self.ax_head(ax)\n        ax_attn = ax_attn.reshape(ax_shape[0], ax_shape[1], -1)\n        #ax_attn = ax_attn.transpose(1, 2)\n        sagt1_shape = sagt1.shape\n        sagt1 = sagt1.reshape(sagt1_shape[0]*sagt1_shape[1]*sagt1_shape[2], 1, sagt1_shape[-2], sagt1_shape[-1])\n        sagt1 = self.sagt1_encoder.forward_features(sagt1)\n        sagt1 = self.sagt1_flatten(sagt1)\n        sagt1 = sagt1.reshape(sagt1_shape[0]*sagt1_shape[1], sagt1_shape[2], -1)\n        sagt1_weighted_sum, sagt1_attn = self.sagt1_head(sagt1)\n        sagt1_attn = sagt1_attn.reshape(sagt1_shape[0], sagt1_shape[1], -1)\n        x = torch.cat([ax_weighted_sum, sagt1_weighted_sum], dim=1)\n        out = self.out(x)\n        return out\n\nclass SSMIL(nn.Module):\n    def __init__(self, model_name):\n        super().__init__()\n        if 'convnext' in model_name: \n            self.ax_encoder = timm.create_model('convnext_small.fb_in22k_ft_in1k_384', in_chans=1, pretrained=False, num_classes=0)\n        elif 'effv2s' in model_name: \n            self.ax_encoder = timm.create_model('tf_efficientnetv2_s.in21k_ft_in1k', in_chans=1, pretrained=False, num_classes=0)\n        \n        self.ax_flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n                                    nn.Flatten(1))\n        self.ax_num_features = self.ax_encoder.num_features\n        self.ax_head = LSTMMIL(self.ax_num_features, 512, 3)\n        self.out = nn.Linear(self.ax_num_features, 3)\n        self.dropout = nn.Dropout(0.0)\n    def forward(self, ax):\n        ax_shape = ax.shape\n        #print(ax.shape, sagt2.shape)\n        ax = ax.reshape(ax_shape[0]*ax_shape[1]*ax_shape[2], 1, ax_shape[-2], ax_shape[-1])\n        ax = self.ax_encoder.forward_features(ax)\n        ax = self.ax_flatten(ax)\n        ax = ax.reshape(ax_shape[0]*ax_shape[1], ax_shape[2], -1)\n        ax_weighted_sum, ax_attn = self.ax_head(ax)\n        ax_attn = ax_attn.reshape(ax_shape[0], ax_shape[1], -1)\n\n        #out = torch.cat([ax_weighted_sum, sagt2_weighted_sum], dim=1)\n        out = ax_weighted_sum\n        out = self.out(out)\n        #print(sagt2_attn.shape, ax_attn.shape)\n        ax_attn = {'left_L1/L2': ax_attn[:, 0, :],'left_L2/L3': ax_attn[:, 1, :],'left_L3/L4': ax_attn[:, 2, :], 'left_L4/L5': ax_attn[:, 3, :], 'left_L5/S1': ax_attn[:, 4, :],\n                'right_L1/L2': ax_attn[:, 5, :], 'right_L2/L3': ax_attn[:, 6, :], 'right_L3/L4': ax_attn[:, 7, :], 'right_L4/L5': ax_attn[:, 8, :], 'right_L5/S1': ax_attn[:, 9, :]}\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:14:40.353486Z","iopub.execute_input":"2025-06-07T01:14:40.353755Z","iopub.status.idle":"2025-06-07T01:14:40.382841Z","shell.execute_reply.started":"2025-06-07T01:14:40.353731Z","shell.execute_reply":"2025-06-07T01:14:40.382249Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Severity prediction lightning modules","metadata":{}},{"cell_type":"code","source":"class ClassModule(pl.LightningModule):\n    def __init__(self, condition, model_name):\n        super().__init__()\n        self.condition = condition\n        if condition == 'scs':\n            if 'mil' in model_name: \n                self.model = SCSMIL(model_name)\n            else: \n                self.model = ConConvnextSCS('sagt2')\n        elif condition == 'nfn':\n            if 'mil' in model_name: \n                self.model = NFNMIL(model_name)\n            else: \n                self.model = ConConvnextNFN()\n        elif condition == 'ss':\n            if 'mil' in model_name: \n                self.model = SSMIL(model_name)\n            else: \n                self.model = ConvnextSS('ax')\n\n    def forward(self, batch):\n        preds = self.model(**batch)\n        return preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:14:40.383552Z","iopub.execute_input":"2025-06-07T01:14:40.383778Z","iopub.status.idle":"2025-06-07T01:14:40.396966Z","shell.execute_reply.started":"2025-06-07T01:14:40.383763Z","shell.execute_reply":"2025-06-07T01:14:40.396439Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Severity inference","metadata":{}},{"cell_type":"code","source":"%%time\nprefix = ''\nimport warnings\nwarnings.filterwarnings(\"ignore\")\nseverity_predict = {'scs': {\n                     'L1/L2':[], \n                     'L2/L3': [], \n                     'L3/L4': [], \n                     'L4/L5': [], \n                     'L5/S1': []\n                     }, \n                 'nfn': {\n                     'left_L1/L2': [], \n                     'left_L2/L3': [], \n                     'left_L3/L4': [], \n                     'left_L4/L5': [], \n                     'left_L5/S1': [], \n                     'right_L1/L2': [], \n                     'right_L2/L3': [], \n                     'right_L3/L4': [], \n                     'right_L4/L5': [], \n                     'right_L5/S1': [], \n                     }, \n                'ss': {\n                     'left_L1/L2': [], \n                     'left_L2/L3': [], \n                     'left_L3/L4': [], \n                     'left_L4/L5': [], \n                     'left_L5/S1': [], \n                     'right_L1/L2': [], \n                     'right_L2/L3': [], \n                     'right_L3/L4': [], \n                     'right_L4/L5': [], \n                     'right_L5/S1': [], \n                     }\n                    }\nmodel_path_dict = {\n    'scs': [\n           '/kaggle/input/rsna-spine-final-models/_scs_classify_5ch_axsagt2-lstm-mil_auxloss_auxdepth_convnext-s_for_exp0.ckpt', \n           '/kaggle/input/rsna-spine-final-models/_scs_classify_5ch_axsagt2-lstm-mil_auxloss_auxdepth_convnext-s_for_exp1.ckpt', \n           '/kaggle/input/rsna-spine-final-models/_scs_classify_5ch_axsagt2-lstm-mil_auxloss_auxdepth_convnext-s_for_exp2.ckpt', \n           '/kaggle/input/rsna-spine-final-models/_scs_classify_5ch_axsagt2-lstm-mil_auxloss_auxdepth_convnext-s_for_exp3.ckpt', \n           '/kaggle/input/rsna-spine-final-models/_scs_classify_5ch_axsagt2-lstm-mil_auxloss_auxdepth_convnext-s_for_exp4.ckpt', \n           '/kaggle/input/rsna-spine-final-models/_scs_classify_5ch_axsagt2-lstm-mil_auxloss_auxdepth_effv2s_for_exp0.ckpt', \n           '/kaggle/input/rsna-spine-final-models/_scs_classify_5ch_axsagt2-lstm-mil_auxloss_auxdepth_effv2s_for_exp1.ckpt', \n           '/kaggle/input/rsna-spine-final-models/_scs_classify_5ch_axsagt2-lstm-mil_auxloss_auxdepth_effv2s_for_exp2.ckpt', \n           '/kaggle/input/rsna-spine-final-models/_scs_classify_5ch_axsagt2-lstm-mil_auxloss_auxdepth_effv2s_for_exp3.ckpt', \n           '/kaggle/input/rsna-spine-final-models/_scs_classify_5ch_axsagt2-lstm-mil_auxloss_auxdepth_effv2s_for_exp4.ckpt', \n           ], \n    'nfn': [\n          '/kaggle/input/rsna-spine-final-models/nfn_classify_5ch_axsagt1-lstm-mil_auxloss_auxdepth_2shift_convnext-s_0.ckpt', \n          '/kaggle/input/rsna-spine-final-models/nfn_classify_5ch_axsagt1-lstm-mil_auxloss_auxdepth_2shift_convnext-s_1.ckpt', \n          '/kaggle/input/rsna-spine-final-models/nfn_classify_5ch_axsagt1-lstm-mil_auxloss_auxdepth_2shift_convnext-s_2.ckpt', \n          '/kaggle/input/rsna-spine-final-models/nfn_classify_5ch_axsagt1-lstm-mil_auxloss_auxdepth_2shift_convnext-s_3.ckpt',\n          '/kaggle/input/rsna-spine-final-models/nfn_classify_5ch_axsagt1-lstm-mil_auxloss_auxdepth_2shift_convnext-s_4.ckpt',\n          '/kaggle/input/rsna-spine-final-models/nfn_classify_5ch_axsagt1-lstm-mil_auxloss_auxdepth_2shift_effv2s_0.ckpt', \n          '/kaggle/input/rsna-spine-final-models/nfn_classify_5ch_axsagt1-lstm-mil_auxloss_auxdepth_2shift_effv2s_1.ckpt', \n          '/kaggle/input/rsna-spine-final-models/nfn_classify_5ch_axsagt1-lstm-mil_auxloss_auxdepth_2shift_effv2s_2.ckpt', \n          '/kaggle/input/rsna-spine-final-models/nfn_classify_5ch_axsagt1-lstm-mil_auxloss_auxdepth_2shift_effv2s_3.ckpt',\n          '/kaggle/input/rsna-spine-final-models/nfn_classify_5ch_axsagt1-lstm-mil_auxloss_auxdepth_2shift_effv2s_4.ckpt',\n           ], \n    'ss': [\n          '/kaggle/input/rsna-spine-final-models/ss_classify_5ch_ax-lstm-mil_auxloss_auxdepth_effv2s_0.ckpt', \n          '/kaggle/input/rsna-spine-final-models/ss_classify_5ch_ax-lstm-mil_auxloss_auxdepth_effv2s_1.ckpt', \n          '/kaggle/input/rsna-spine-final-models/ss_classify_5ch_ax-lstm-mil_auxloss_auxdepth_effv2s_2.ckpt', \n          '/kaggle/input/rsna-spine-final-models/ss_classify_5ch_ax-lstm-mil_auxloss_auxdepth_effv2s_3.ckpt',\n          '/kaggle/input/rsna-spine-final-models/ss_classify_5ch_ax-lstm-mil_auxloss_auxdepth_effv2s_4.ckpt',\n          '/kaggle/input/rsna-spine-final-models/ss_classify_5ch_ax-lstm-mil_auxloss_auxdepth_convnext-s_0.ckpt', \n          '/kaggle/input/rsna-spine-final-models/ss_classify_5ch_ax-lstm-mil_auxloss_auxdepth_convnext-s_1.ckpt', \n          '/kaggle/input/rsna-spine-final-models/ss_classify_5ch_ax-lstm-mil_auxloss_auxdepth_convnext-s_2.ckpt', \n          '/kaggle/input/rsna-spine-final-models/ss_classify_5ch_ax-lstm-mil_auxloss_auxdepth_convnext-s_3.ckpt',\n          '/kaggle/input/rsna-spine-final-models/ss_classify_5ch_ax-lstm-mil_auxloss_auxdepth_convnext-s_4.ckpt',\n          ]\n}\n##############SEVERITY PREDICT#########################\nfor condition in ['nfn', 'scs', 'ss']:\n    print(condition)\n    for path in model_path_dict[condition]:\n        _meta_df = meta_df.copy()\n        _coor_df = pred_coor_stage3.copy()\n        model_name = path.split('/')[-1]\n        if '5ch' in model_name: \n            dataset_channel = 5\n        else: \n            dataset_channel = 3\n        dataset_test = ClassDataset(_coor_df, _meta_df,  condition, dataset_channel, 'sub')\n        data_loader_test = DataLoader(\n            dataset_test,\n            batch_size=4,\n            shuffle=False,\n            num_workers=4,\n            pin_memory=False\n        )\n\n        model = ClassModule.load_from_checkpoint(path, condition=condition, model_name=model_name, strict=False)\n        model.eval()\n        model.zero_grad()\n        model.to(device)\n        pred_temp = {}\n        for k in severity_predict[condition].keys(): \n            pred_temp[k] = []\n        study_id_list = []\n        with torch.no_grad():\n            for data in tqdm(data_loader_test, total=len(data_loader_test)):\n                images, study_id = data\n                for k, v in images.items(): \n                    images[k] = v.to(device)\n                    bs = v.shape[0]\n                preds = model.forward(images)\n                preds = nn.functional.softmax(preds, dim=1)\n                preds = preds.reshape((bs, -1, 3))\n                preds = preds.to('cpu').detach().numpy()\n                #print(preds)\n                if condition == 'scs':  \n                    pred_temp['L1/L2'].append(preds[:, 0, :])\n                    pred_temp['L2/L3'].append(preds[:, 1, :])\n                    pred_temp['L3/L4'].append(preds[:, 2, :])\n                    pred_temp['L4/L5'].append(preds[:, 3, :])\n                    pred_temp['L5/S1'].append(preds[:, 4, :])\n                else: \n                    pred_temp['left_L1/L2'].append(preds[:, 0, :])\n                    pred_temp['left_L2/L3'].append(preds[:, 1, :])\n                    pred_temp['left_L3/L4'].append(preds[:, 2, :])\n                    pred_temp['left_L4/L5'].append(preds[:, 3, :])\n                    pred_temp['left_L5/S1'].append(preds[:, 4, :])\n                    pred_temp['right_L1/L2'].append(preds[:, 5, :])\n                    pred_temp['right_L2/L3'].append(preds[:, 6, :])\n                    pred_temp['right_L3/L4'].append(preds[:, 7, :])\n                    pred_temp['right_L4/L5'].append(preds[:, 8, :])\n                    pred_temp['right_L5/S1'].append(preds[:, 9, :])\n                study_id_list.append(study_id.to('cpu').reshape(-1).detach().numpy())\n                del images, preds\n                gc.collect()\n        for k, v in pred_temp.items(): \n            severity_predict[condition][k].append(np.concatenate(v))\n        all_study_id = np.concatenate(study_id_list)\n        del pred_temp, study_id_list\n        gc.collect()\n        \n    for k, v in severity_predict[condition].items(): \n        severity_predict[condition][k] = np.mean(np.array(severity_predict[condition][k]), axis=0)\n    severity_predict[condition]['study_id'] = all_study_id\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:14:40.397741Z","iopub.execute_input":"2025-06-07T01:14:40.398013Z","iopub.status.idle":"2025-06-07T01:20:35.988609Z","shell.execute_reply.started":"2025-06-07T01:14:40.397987Z","shell.execute_reply":"2025-06-07T01:20:35.986295Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Stage Final: Submission","metadata":{}},{"cell_type":"code","source":"condition_mapper = {'scs': 'spinal_canal_stenosis', 'nfn': 'neural_foraminal_narrowing', 'ss': 'subarticular_stenosis'}\npredict_list = []\nfor k, v in severity_predict.items():\n    condition = condition_mapper[k]\n    study_id = severity_predict[k]['study_id']\n    for kk, vv in v.items(): \n        if kk == 'study_id': \n            continue\n        level = kk.split('_')[-1].lower().replace('/', '_')\n        loc = kk.split('_')[0]\n        if loc == 'left':\n            loc = 'left_'\n        elif loc == 'right':\n            loc = 'right_'\n        else: \n            loc = ''\n        row_id = [f'{str(si)}_' + loc + condition + '_' + level for si in study_id]\n        df = pd.DataFrame({'row_id': row_id, 'normal_mild': vv[:, 0], 'moderate': vv[:, 1], 'severe': vv[:, 2]})\n        predict_list.append(df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:32:50.878934Z","iopub.execute_input":"2025-06-07T01:32:50.879685Z","iopub.status.idle":"2025-06-07T01:32:50.890969Z","shell.execute_reply.started":"2025-06-07T01:32:50.879658Z","shell.execute_reply":"2025-06-07T01:32:50.890312Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predict_df = pd.concat(predict_list)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:32:56.061027Z","iopub.execute_input":"2025-06-07T01:32:56.061583Z","iopub.status.idle":"2025-06-07T01:32:56.066783Z","shell.execute_reply.started":"2025-06-07T01:32:56.061560Z","shell.execute_reply":"2025-06-07T01:32:56.066233Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predict_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:33:52.197518Z","iopub.execute_input":"2025-06-07T01:33:52.197982Z","iopub.status.idle":"2025-06-07T01:33:52.209840Z","shell.execute_reply.started":"2025-06-07T01:33:52.197960Z","shell.execute_reply":"2025-06-07T01:33:52.209052Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/sample_submission.csv')\nsub = sub.drop(['normal_mild', 'moderate', 'severe'], axis=1)\nsub = sub.merge(predict_df, on='row_id', how='left')\nsub = sub.fillna(1/3)\nsub[['normal_mild', 'moderate', 'severe']] = sub[['normal_mild', 'moderate', 'severe']]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:33:02.152772Z","iopub.execute_input":"2025-06-07T01:33:02.153321Z","iopub.status.idle":"2025-06-07T01:33:02.169100Z","shell.execute_reply.started":"2025-06-07T01:33:02.153299Z","shell.execute_reply":"2025-06-07T01:33:02.168426Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub.to_csv('submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:33:28.309979Z","iopub.execute_input":"2025-06-07T01:33:28.310516Z","iopub.status.idle":"2025-06-07T01:33:28.315660Z","shell.execute_reply.started":"2025-06-07T01:33:28.310474Z","shell.execute_reply":"2025-06-07T01:33:28.314814Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.listdir('/kaggle/working/')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:54:36.178338Z","iopub.execute_input":"2025-06-07T01:54:36.179021Z","iopub.status.idle":"2025-06-07T01:54:36.184182Z","shell.execute_reply.started":"2025-06-07T01:54:36.178996Z","shell.execute_reply":"2025-06-07T01:54:36.183551Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_submiss = pd.read_csv('submission.csv')\ndata_submiss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-07T01:57:38.800206Z","iopub.execute_input":"2025-06-07T01:57:38.800479Z","iopub.status.idle":"2025-06-07T01:57:38.812985Z","shell.execute_reply.started":"2025-06-07T01:57:38.800459Z","shell.execute_reply":"2025-06-07T01:57:38.812324Z"}},"outputs":[],"execution_count":null}]}