{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9127652,"sourceType":"datasetVersion","datasetId":5494026},{"sourceId":9729124,"sourceType":"datasetVersion","datasetId":5953414}],"dockerImageVersionId":30747,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# import & configure","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":{"execution":{"iopub.status.busy":"2024-10-26T16:02:51.3836Z","iopub.execute_input":"2024-10-26T16:02:51.384083Z","iopub.status.idle":"2024-10-26T16:02:58.019676Z","shell.execute_reply.started":"2024-10-26T16:02:51.384049Z","shell.execute_reply":"2024-10-26T16:02:58.018726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 126 # friend's birthday\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":{"execution":{"iopub.status.busy":"2024-10-26T16:02:58.021626Z","iopub.execute_input":"2024-10-26T16:02:58.022094Z","iopub.status.idle":"2024-10-26T16:02:58.064466Z","shell.execute_reply.started":"2024-10-26T16:02:58.022067Z","shell.execute_reply":"2024-10-26T16:02:58.063663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"デバッグモードとseriesの書き換え, 各inferenceのfold","metadata":{}},{"cell_type":"code","source":"%%writefile config.yaml\n\ndata_path: \"/content/\"\noutput_dir: \"/content/gdrive/MyDrive/RSNA_SPINE/models/\"\n\nseed: 1101\ndebug: False\ntrain_bs: 4\nvalid_bs: 4\ntest_bs: 8\nworkers: 1\n\nprogress_bar_refresh_rate: 1\n\npseudo_train: 0\n\nsave_topk: 1\nfold: 5\n\ntask:\n    kind: 'detect'\n    #kind: 'classify'\n    #kind: 'depth'\n    condition: 'nfn'\n    #condition: 'scs'\n    #condition: 'scs'\n    #condition: 'all'\n    #direction: 'sagt2'\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:\n    ","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:02:58.065675Z","iopub.execute_input":"2024-10-26T16:02:58.066326Z","iopub.status.idle":"2024-10-26T16:02:58.072445Z","shell.execute_reply.started":"2024-10-26T16:02:58.066292Z","shell.execute_reply":"2024-10-26T16:02:58.071617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(\"config.yaml\", \"r\") as file_obj:\n    config = yaml.safe_load(file_obj)","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:02:58.073467Z","iopub.execute_input":"2024-10-26T16:02:58.073746Z","iopub.status.idle":"2024-10-26T16:02:58.089833Z","shell.execute_reply.started":"2024-10-26T16:02:58.073723Z","shell.execute_reply":"2024-10-26T16:02:58.088936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# use train data when debug mode on\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":{"execution":{"iopub.status.busy":"2024-10-26T16:02:58.092209Z","iopub.execute_input":"2024-10-26T16:02:58.092727Z","iopub.status.idle":"2024-10-26T16:02:58.103015Z","shell.execute_reply.started":"2024-10-26T16:02:58.092703Z","shell.execute_reply":"2024-10-26T16:02:58.102169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"series.head()","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:02:58.104138Z","iopub.execute_input":"2024-10-26T16:02:58.104464Z","iopub.status.idle":"2024-10-26T16:02:58.117523Z","shell.execute_reply.started":"2024-10-26T16:02:58.104432Z","shell.execute_reply":"2024-10-26T16:02:58.116671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# initial stage (create meta file)\n\n- create meta file using dicom's meta data","metadata":{}},{"cell_type":"code","source":"def create_dcm_df(study_id, series_id, series_desc): \n    try: \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        dcm_list = []\n        for i in in_list: \n            dcm_list.append(dcm.dcmread(IMAGE_PATH + f'{study_id}/{series_id}/{i}.dcm'))\n        #dcm_list = [dcm.dcmread(IMAGE_PATH + f'{study_id}/{series_id}/{i}.dcm') for i in in_list]\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        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')\n        ps_x = ps[:, 0]\n        ps_y = ps[:, 1]\n        meta_dict = {'instance_number': in_list, 'ipp_x': ipp_x, 'ipp_y': ipp_y, 'ipp_z': ipp_z, 'sbs': sbs, 'ps_x': ps_x, 'ps_y': ps_y}\n        meta_df = pd.DataFrame(meta_dict)\n        meta_df['series_id'] = series_id\n        meta_df['study_id'] = study_id\n        meta_df['series_description'] = series_desc\n        meta_df['height'] = shape[:, 0]\n        meta_df['width'] = shape[:, 1]\n        meta_df['iop'] = pd.Series(iop)\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_desc)\n        return None","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:02:58.118765Z","iopub.execute_input":"2024-10-26T16:02:58.119105Z","iopub.status.idle":"2024-10-26T16:02:58.13388Z","shell.execute_reply.started":"2024-10-26T16:02:58.119076Z","shell.execute_reply":"2024-10-26T16:02:58.133034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nif config['debug']: \n    #meta_df = pd.read_parquet('/kaggle/input/rsna-newmeta/meta.parquet')\n    meta_df_list = []\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    #for _, row in tqdm(test_series.iterrows(), total=len(test_series)): \n    #    meta_df_list.append(create_dcm_df(row.study_id, row.series_id, row.series_description))\n    meta_df = pd.concat(meta_df_list)\n    del meta_df_list\n    gc.collect()\n    meta_df.to_parquet('meta.parquet')\nelse: \n    # spend about 20 min for train data\n    meta_df_list = []\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    #for _, row in tqdm(test_series.iterrows(), total=len(test_series)): \n    #    meta_df_list.append(create_dcm_df(row.study_id, row.series_id, row.series_description))\n    meta_df = pd.concat(meta_df_list)\n    del meta_df_list\n    gc.collect()\n    meta_df.to_parquet('meta.parquet')","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:02:58.135133Z","iopub.execute_input":"2024-10-26T16:02:58.135392Z","iopub.status.idle":"2024-10-26T16:03:00.67465Z","shell.execute_reply.started":"2024-10-26T16:02:58.13537Z","shell.execute_reply":"2024-10-26T16:03:00.673609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# first stage (depth inference)\n\n- infer depth of sagittal t1 & t2\n- align infered depth using rule base algorithms","metadata":{}},{"cell_type":"markdown","source":"## dataset","metadata":{}},{"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.read_file(path)\n        data = dicom.pixel_array\n        return data\n","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:03:00.67613Z","iopub.execute_input":"2024-10-26T16:03:00.676772Z","iopub.status.idle":"2024-10-26T16:03:00.696983Z","shell.execute_reply.started":"2024-10-26T16:03:00.676735Z","shell.execute_reply":"2024-10-26T16:03:00.696029Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## models\n\nreference: [Implementing ConvNext in PyTorch](https://github.com/FrancescoSaverioZuppichini/ConvNext)","metadata":{}},{"cell_type":"code","source":"from torchvision.ops import StochasticDepth\nfrom typing import List, Dict\nfrom torch import Tensor\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":{"execution":{"iopub.status.busy":"2024-10-26T16:03:00.69843Z","iopub.execute_input":"2024-10-26T16:03:00.69878Z","iopub.status.idle":"2024-10-26T16:03:00.781767Z","shell.execute_reply.started":"2024-10-26T16:03:00.698744Z","shell.execute_reply":"2024-10-26T16:03:00.780865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## lightning module","metadata":{}},{"cell_type":"code","source":"class 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":{"execution":{"iopub.status.busy":"2024-10-26T16:03:00.782943Z","iopub.execute_input":"2024-10-26T16:03:00.78325Z","iopub.status.idle":"2024-10-26T16:03:00.79497Z","shell.execute_reply.started":"2024-10-26T16:03:00.783226Z","shell.execute_reply":"2024-10-26T16:03:00.794154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## inference","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":{"execution":{"iopub.status.busy":"2024-10-26T16:03:00.796188Z","iopub.execute_input":"2024-10-26T16:03:00.796461Z","iopub.status.idle":"2024-10-26T16:07:40.142654Z","shell.execute_reply.started":"2024-10-26T16:03:00.796438Z","shell.execute_reply":"2024-10-26T16:07:40.141661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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'])\n","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:07:40.143951Z","iopub.execute_input":"2024-10-26T16:07:40.144254Z","iopub.status.idle":"2024-10-26T16:07:40.199439Z","shell.execute_reply.started":"2024-10-26T16:07:40.144228Z","shell.execute_reply":"2024-10-26T16:07:40.19862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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)","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:07:40.204797Z","iopub.execute_input":"2024-10-26T16:07:40.205126Z","iopub.status.idle":"2024-10-26T16:07:40.385031Z","shell.execute_reply.started":"2024-10-26T16:07:40.205102Z","shell.execute_reply":"2024-10-26T16:07:40.384241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Second Stage (xy inference)\n\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.read_file(path)\n        data = dicom.pixel_array\n        return data","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:07:40.386282Z","iopub.execute_input":"2024-10-26T16:07:40.386598Z","iopub.status.idle":"2024-10-26T16:07:40.426379Z","shell.execute_reply.started":"2024-10-26T16:07:40.386572Z","shell.execute_reply":"2024-10-26T16:07:40.425493Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-26T16:07:40.427654Z","iopub.execute_input":"2024-10-26T16:07:40.42814Z","iopub.status.idle":"2024-10-26T16:07:40.448342Z","shell.execute_reply.started":"2024-10-26T16:07:40.428109Z","shell.execute_reply":"2024-10-26T16:07:40.447515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-26T16:07:40.449471Z","iopub.execute_input":"2024-10-26T16:07:40.450018Z","iopub.status.idle":"2024-10-26T16:07:40.461409Z","shell.execute_reply.started":"2024-10-26T16:07:40.449993Z","shell.execute_reply":"2024-10-26T16:07:40.460631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-26T16:07:40.462632Z","iopub.execute_input":"2024-10-26T16:07:40.462923Z","iopub.status.idle":"2024-10-26T16:11:25.947602Z","shell.execute_reply.started":"2024-10-26T16:07:40.462881Z","shell.execute_reply":"2024-10-26T16:11:25.945782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_coor.head(1)","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:11:25.949181Z","iopub.execute_input":"2024-10-26T16:11:25.949949Z","iopub.status.idle":"2024-10-26T16:11:25.960218Z","shell.execute_reply.started":"2024-10-26T16:11:25.949888Z","shell.execute_reply":"2024-10-26T16:11:25.959432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-26T16:11:25.961522Z","iopub.execute_input":"2024-10-26T16:11:25.961797Z","iopub.status.idle":"2024-10-26T16:11:26.034603Z","shell.execute_reply.started":"2024-10-26T16:11:25.961775Z","shell.execute_reply":"2024-10-26T16:11:26.033746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_coor_stage2 = pd.merge(pred_coor, _pred_coor, on=['study_id', 'series_id', 'level', 'condition'], how='inner')","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:11:26.035522Z","iopub.execute_input":"2024-10-26T16:11:26.035777Z","iopub.status.idle":"2024-10-26T16:11:26.044588Z","shell.execute_reply.started":"2024-10-26T16:11:26.035755Z","shell.execute_reply":"2024-10-26T16:11:26.043749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(pred_coor_stage2.head())\npred_coor_stage2.to_csv('stage2_coor.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:11:26.045624Z","iopub.execute_input":"2024-10-26T16:11:26.04587Z","iopub.status.idle":"2024-10-26T16:11:26.060763Z","shell.execute_reply.started":"2024-10-26T16:11:26.045848Z","shell.execute_reply":"2024-10-26T16:11:26.059956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Third Stage (calc. location of axial t2)\n\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":"# thanks @hengck\n# customized for my pipeline\n# project 2d to 3d\ndef 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":{"execution":{"iopub.status.busy":"2024-10-26T16:11:26.061803Z","iopub.execute_input":"2024-10-26T16:11:26.062099Z","iopub.status.idle":"2024-10-26T16:11:26.075013Z","shell.execute_reply.started":"2024-10-26T16:11:26.062077Z","shell.execute_reply":"2024-10-26T16:11:26.074093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-26T16:11:26.07605Z","iopub.execute_input":"2024-10-26T16:11:26.076337Z","iopub.status.idle":"2024-10-26T16:11:26.110772Z","shell.execute_reply.started":"2024-10-26T16:11:26.076304Z","shell.execute_reply":"2024-10-26T16:11:26.110002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"closest_ax = pd.concat(closest_ax_list)\nclosest_ax.head()","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:11:26.111872Z","iopub.execute_input":"2024-10-26T16:11:26.112139Z","iopub.status.idle":"2024-10-26T16:11:26.121168Z","shell.execute_reply.started":"2024-10-26T16:11:26.112117Z","shell.execute_reply":"2024-10-26T16:11:26.120271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## subarticular stenosis 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.read_file(path)\n        data = dicom.pixel_array\n        return data","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:11:26.122285Z","iopub.execute_input":"2024-10-26T16:11:26.122586Z","iopub.status.idle":"2024-10-26T16:11:26.135Z","shell.execute_reply.started":"2024-10-26T16:11:26.122552Z","shell.execute_reply":"2024-10-26T16:11:26.134253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## subarticular stenosis coordinate detection 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":{"execution":{"iopub.status.busy":"2024-10-26T16:11:26.135989Z","iopub.execute_input":"2024-10-26T16:11:26.13622Z","iopub.status.idle":"2024-10-26T16:11:26.148521Z","shell.execute_reply.started":"2024-10-26T16:11:26.1362Z","shell.execute_reply":"2024-10-26T16:11:26.147813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-26T16:11:26.149695Z","iopub.execute_input":"2024-10-26T16:11:26.14998Z","iopub.status.idle":"2024-10-26T16:11:26.160988Z","shell.execute_reply.started":"2024-10-26T16:11:26.149958Z","shell.execute_reply":"2024-10-26T16:11:26.160158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## subarticular stenosis 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":{"execution":{"iopub.status.busy":"2024-10-26T16:11:26.162103Z","iopub.execute_input":"2024-10-26T16:11:26.162419Z","iopub.status.idle":"2024-10-26T16:12:17.431476Z","shell.execute_reply.started":"2024-10-26T16:11:26.162395Z","shell.execute_reply":"2024-10-26T16:12:17.430506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-26T16:12:17.433052Z","iopub.execute_input":"2024-10-26T16:12:17.433501Z","iopub.status.idle":"2024-10-26T16:12:17.4568Z","shell.execute_reply.started":"2024-10-26T16:12:17.433464Z","shell.execute_reply":"2024-10-26T16:12:17.455938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-26T16:12:17.457867Z","iopub.execute_input":"2024-10-26T16:12:17.458172Z","iopub.status.idle":"2024-10-26T16:12:17.474687Z","shell.execute_reply.started":"2024-10-26T16:12:17.458147Z","shell.execute_reply":"2024-10-26T16:12:17.47388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-26T16:12:17.475758Z","iopub.execute_input":"2024-10-26T16:12:17.476037Z","iopub.status.idle":"2024-10-26T16:12:17.49368Z","shell.execute_reply.started":"2024-10-26T16:12:17.476014Z","shell.execute_reply":"2024-10-26T16:12:17.492833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_coor_stage3","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:12:17.495025Z","iopub.execute_input":"2024-10-26T16:12:17.495598Z","iopub.status.idle":"2024-10-26T16:12:17.509451Z","shell.execute_reply.started":"2024-10-26T16:12:17.495558Z","shell.execute_reply":"2024-10-26T16:12:17.508551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del pred_coor_stage2, ax_coor_pred, pred_coor, closest_ax\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:12:17.51063Z","iopub.execute_input":"2024-10-26T16:12:17.510913Z","iopub.status.idle":"2024-10-26T16:12:17.69886Z","shell.execute_reply.started":"2024-10-26T16:12:17.51087Z","shell.execute_reply":"2024-10-26T16:12:17.69794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Fourth Stage (predict each severity)","metadata":{}},{"cell_type":"markdown","source":"## severity prediction datasets","metadata":{}},{"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.read_file(path)\n        data = dicom.pixel_array\n        return data","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:12:17.700247Z","iopub.execute_input":"2024-10-26T16:12:17.700562Z","iopub.status.idle":"2024-10-26T16:12:17.776701Z","shell.execute_reply.started":"2024-10-26T16:12:17.700538Z","shell.execute_reply":"2024-10-26T16:12:17.775848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\n    \n\n","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:12:17.778201Z","iopub.execute_input":"2024-10-26T16:12:17.77854Z","iopub.status.idle":"2024-10-26T16:12:17.805034Z","shell.execute_reply.started":"2024-10-26T16:12:17.778511Z","shell.execute_reply":"2024-10-26T16:12:17.804127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-26T16:12:17.806489Z","iopub.execute_input":"2024-10-26T16:12:17.806735Z","iopub.status.idle":"2024-10-26T16:12:17.857275Z","shell.execute_reply.started":"2024-10-26T16:12:17.806713Z","shell.execute_reply":"2024-10-26T16:12:17.856428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\n\n","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:12:17.858385Z","iopub.execute_input":"2024-10-26T16:12:17.858638Z","iopub.status.idle":"2024-10-26T16:12:17.87037Z","shell.execute_reply.started":"2024-10-26T16:12:17.858615Z","shell.execute_reply":"2024-10-26T16:12:17.869636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-26T16:12:17.871566Z","iopub.execute_input":"2024-10-26T16:12:17.871831Z","iopub.status.idle":"2024-10-26T16:19:02.696094Z","shell.execute_reply.started":"2024-10-26T16:12:17.871809Z","shell.execute_reply":"2024-10-26T16:19:02.693793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# fifth stage (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":{"execution":{"iopub.status.busy":"2024-10-26T16:19:02.697633Z","iopub.execute_input":"2024-10-26T16:19:02.698103Z","iopub.status.idle":"2024-10-26T16:19:02.714213Z","shell.execute_reply.started":"2024-10-26T16:19:02.698067Z","shell.execute_reply":"2024-10-26T16:19:02.713441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predict_df = pd.concat(predict_list)","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:19:02.720404Z","iopub.execute_input":"2024-10-26T16:19:02.72107Z","iopub.status.idle":"2024-10-26T16:19:02.727003Z","shell.execute_reply.started":"2024-10-26T16:19:02.721024Z","shell.execute_reply":"2024-10-26T16:19:02.726167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-10-26T16:19:02.728103Z","iopub.execute_input":"2024-10-26T16:19:02.728429Z","iopub.status.idle":"2024-10-26T16:19:02.752854Z","shell.execute_reply.started":"2024-10-26T16:19:02.728398Z","shell.execute_reply":"2024-10-26T16:19:02.752142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:19:02.753915Z","iopub.execute_input":"2024-10-26T16:19:02.754169Z","iopub.status.idle":"2024-10-26T16:19:02.760528Z","shell.execute_reply.started":"2024-10-26T16:19:02.754147Z","shell.execute_reply":"2024-10-26T16:19:02.759716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"sub = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/sample_submission.csv')\n\nif len(sub) == 25: \n    sub.to_csv('submission.csv', index=False)","metadata":{}},{"cell_type":"code","source":"!rm /kaggle/working/meta.parquet\n!rm /kaggle/working/config.yaml\n!rm /kaggle/working/stage3_coor.csv\n!rm /kaggle/working/stage1_coor.csv\n!rm /kaggle/working/stage2_coor.csv","metadata":{"execution":{"iopub.status.busy":"2024-10-26T16:19:02.761521Z","iopub.execute_input":"2024-10-26T16:19:02.761842Z","iopub.status.idle":"2024-10-26T16:19:07.965586Z","shell.execute_reply.started":"2024-10-26T16:19:02.761813Z","shell.execute_reply":"2024-10-26T16:19:07.964389Z"},"trusted":true},"execution_count":null,"outputs":[]}]}