{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":699609,"sourceType":"datasetVersion","datasetId":255887},{"sourceId":6536400,"sourceType":"datasetVersion","datasetId":3778855},{"sourceId":6536441,"sourceType":"datasetVersion","datasetId":3778880},{"sourceId":9160420,"sourceType":"datasetVersion","datasetId":5528077},{"sourceId":9518875,"sourceType":"datasetVersion","datasetId":5795333},{"sourceId":9529116,"sourceType":"datasetVersion","datasetId":5521343},{"sourceId":9570179,"sourceType":"datasetVersion","datasetId":5831032},{"sourceId":9575214,"sourceType":"datasetVersion","datasetId":5521562}],"dockerImageVersionId":30747,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## library","metadata":{}},{"cell_type":"code","source":"import sys\n\nsys.path.append('/kaggle/input/segmentation-library/segmentation_models.pytorch')\nsys.path.append('/kaggle/input/pretrainedmodels/pretrainedmodels-0.7.4')\nsys.path.append('/kaggle/input/efficientnet-library/EfficientNet-PyTorch')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:33:02.438994Z","iopub.execute_input":"2024-10-08T15:33:02.439687Z","iopub.status.idle":"2024-10-08T15:33:02.450565Z","shell.execute_reply.started":"2024-10-08T15:33:02.439658Z","shell.execute_reply":"2024-10-08T15:33:02.44972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\n\nimport numpy as np\n\nimport glob\n\nimport gc\n\nimport matplotlib.pyplot as plt\n\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport albumentations as A\n\nimport pydicom\n\nimport timm\n\nfrom transformers import RobertaPreLayerNormConfig, RobertaPreLayerNormModel\n\nfrom segmentation_models_pytorch.decoders.unet.model import (\n    UnetDecoder,\n    SegmentationHead,\n)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:33:02.451933Z","iopub.execute_input":"2024-10-08T15:33:02.452199Z","iopub.status.idle":"2024-10-08T15:33:13.944206Z","shell.execute_reply.started":"2024-10-08T15:33:02.452168Z","shell.execute_reply":"2024-10-08T15:33:13.943283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## config","metadata":{}},{"cell_type":"code","source":"class CustomConfig:\n    device = 'cuda'\n    root = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\n    \n    orientations = ['Sagittal T2/STIR', 'Sagittal T1', 'Axial T2']\n    \n    stage1_weights = [\n        [\n                    '/kaggle/input/stage1-rsna-effnet/sagittal_t2/fold1/epoch030-trainloss0.164532-testloss0.183283-testscore0.92555.bin',\n                    '/kaggle/input/stage1-rsna-effnet/sagittal_t2/fold2/epoch035-trainloss0.163079-testloss0.187091-testscore0.863966.bin',\n                    '/kaggle/input/stage1-rsna-effnet/sagittal_t2/fold3/epoch040-trainloss0.162728-testloss0.184239-testscore0.881218.bin',\n                    '/kaggle/input/stage1-rsna-effnet/sagittal_t2/fold4/epoch030-trainloss0.164352-testloss0.18379-testscore0.868186.bin',\n                    '/kaggle/input/stage1-rsna-effnet/sagittal_t2/fold5/epoch025-trainloss0.16676-testloss0.183907-testscore0.856371.bin'\n        ],\n        [\n                    '/kaggle/input/stage1-rsna-effnet/sagittal_t1/fold1/epoch040-trainloss0.177427-testloss0.193425-testscore0.986041.bin',\n                    '/kaggle/input/stage1-rsna-effnet/sagittal_t1/fold2/epoch040-trainloss0.177861-testloss0.190003-testscore1.046954.bin',\n                    '/kaggle/input/stage1-rsna-effnet/sagittal_t1/fold3/epoch040-trainloss0.176816-testloss0.196593-testscore1.07519.bin',\n                    '/kaggle/input/stage1-rsna-effnet/sagittal_t1/fold4/epoch040-trainloss0.175746-testloss0.208547-testscore1.044726.bin',\n                    '/kaggle/input/stage1-rsna-effnet/sagittal_t1/fold5/epoch040-trainloss0.176303-testloss0.194872-testscore1.036294.bin'\n        ],\n        [\n                    '/kaggle/input/stage1-rsna-effnet/axial_t2/fold1/epoch035-trainloss0.095594-testloss0.102066-testscore1.408456.bin'\n        ],\n    ]\n    \n    n_locations = {\n        'Sagittal T2/STIR' : 5,\n        'Sagittal T1' : 10,\n        'Axial T2' : 10,\n    }\n    \n    batch_size = 4\n    n_worker = 4\n    \n    n_class = 3\n    \n    label_columns = pd.read_csv(root + 'train.csv').columns[1:]\n    \n    stage2_weights = [\n         '/kaggle/input/rsna-stage2-weights/v1/2024/fold2/epoch003-trainloss0.50246-testloss0.384054-testscore0.384054.bin',\n         '/kaggle/input/rsna-stage2-weights/v1/2024/fold4/epoch012-trainloss0.510078-testloss0.435106-testscore0.400931.bin',\n         '/kaggle/input/rsna-stage2-weights/v1/2024/fold5/epoch017-trainloss0.479376-testloss0.422528-testscore0.382907.bin',\n         '/kaggle/input/rsna-stage2-weights/v1/42/fold1/epoch003-trainloss0.501287-testloss0.395612-testscore0.395612.bin',\n         '/kaggle/input/rsna-stage2-weights/v1/42/fold3/epoch005-trainloss0.480527-testloss0.361974-testscore0.361974.bin',\n\n        '/kaggle/input/rsna-stage2-weights/v2/2024/fold4/epoch014-trainloss0.489339-testloss0.430131-testscore0.401807.bin',\n        '/kaggle/input/rsna-stage2-weights/v2/42/fold1/epoch006-trainloss0.518346-testloss0.39402-testscore0.39402.bin',\n        '/kaggle/input/rsna-stage2-weights/v2/42/fold2/epoch003-trainloss0.512317-testloss0.387429-testscore0.387429.bin',\n        '/kaggle/input/rsna-stage2-weights/v2/42/fold3/epoch005-trainloss0.487169-testloss0.360147-testscore0.360147.bin',\n        '/kaggle/input/rsna-stage2-weights/v2/42/fold5/epoch003-trainloss0.495823-testloss0.379182-testscore0.379182.bin',\n\n\n        '/kaggle/input/rsna-stage2-weights/v3/2024/fold2/epoch003-trainloss0.521357-testloss0.389746-testscore0.389746.bin',\n        '/kaggle/input/rsna-stage2-weights/v3/2024/fold3/epoch005-trainloss0.491033-testloss0.366064-testscore0.366064.bin',\n        '/kaggle/input/rsna-stage2-weights/v3/2024/fold4/epoch004-trainloss0.506416-testloss0.401614-testscore0.401614.bin',\n        '/kaggle/input/rsna-stage2-weights/v3/2024/fold5/epoch006-trainloss0.50526-testloss0.381305-testscore0.381305.bin',\n        '/kaggle/input/rsna-stage2-weights/v3/42/fold1/epoch004-trainloss0.506298-testloss0.392619-testscore0.392619.bin',\n\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch004-trainloss0.474929-testloss0.394903-testscore0.394903.bin',\n       '/kaggle/input/rsna-lsdc-stage2-weights/epoch006-trainloss0.485093-testloss0.38113-testscore0.38113.bin',\n       '/kaggle/input/rsna-lsdc-stage2-weights/epoch006-trainloss0.487418-testloss0.367523-testscore0.367523.bin',\n       '/kaggle/input/rsna-lsdc-stage2-weights/epoch003-trainloss0.490171-testloss0.387255-testscore0.387255.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch006-trainloss0.491637-testloss0.378409-testscore0.378409.bin',\n        \n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch010-trainloss0.488406-testloss0.400828-testscore0.400829.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch005-trainloss0.478115-testloss0.382623-testscore0.382623.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch010-trainloss0.490493-testloss0.368579-testscore0.368579.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch004-trainloss0.481913-testloss0.391444-testscore0.391444.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch006-trainloss0.479494-testloss0.381912-testscore0.381912.bin',\n        \n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch010-trainloss0.472179-testloss0.398393-testscore0.398393.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch006-trainloss0.481968-testloss0.384406-testscore0.384406.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch010-trainloss0.482629-testloss0.372525-testscore0.372525.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch005-trainloss0.484465-testloss0.389893-testscore0.389893.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch006-trainloss0.483708-testloss0.381362-testscore0.381362.bin',\n    \n    ]\n    \n\n    model_names = ['convnext_small.in12k_ft_in1k']*15 + \\\n                  ['convnext_tiny.in12k_ft_in1k'] * 5 + \\\n                  ['caformer_s18.sail_in22k_ft_in1k'] * 5 + \\\n                  ['pvt_v2_b3.in1k'] * 5\n    \n    versions = [1]*5 + [2]*5 + [3]*5 + [1]*15\n    \n    hidden_sizes = [768] * 20 + [512] * 10\n\n    patch_size = 32\n    image_size = 128\n    depth_size = 1 + 2*2\n    \n    mode = 'test'\n    \nif __name__ == \"__main__\":\n    args = CustomConfig()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:33:13.945855Z","iopub.execute_input":"2024-10-08T15:33:13.946374Z","iopub.status.idle":"2024-10-08T15:33:13.991885Z","shell.execute_reply.started":"2024-10-08T15:33:13.946311Z","shell.execute_reply":"2024-10-08T15:33:13.991143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# stage1","metadata":{}},{"cell_type":"markdown","source":"## utils","metadata":{}},{"cell_type":"code","source":"def volume2point(volumes):\n    volumes = volumes.cpu()\n\n    batch_size, n_class, _, h, w = volumes.shape\n\n    points = []\n    for i in range(batch_size):\n        volume = volumes[i]\n\n        point = []\n        for j in range(n_class):\n            flat_volume = volume[j].reshape(-1)\n            max_idx = torch.argmax(flat_volume)\n\n            z = max_idx // (h * w)\n            y = (max_idx % (h * w)) // w\n            x = (max_idx % (h * w)) % w\n            point.append(torch.stack([z, y, x]))\n\n        point = torch.stack(point, dim = 0)\n        points.append(point)\n\n    points = torch.stack(points, dim = 0)\n    return points.numpy()\n\n\ndef volume2point2(volumes):\n    volumes = volumes.cpu()\n\n    batch_size, n_class, _, h, w = volumes.shape\n\n    points = []\n    for i in range(batch_size):\n        volume = volumes[i]\n\n        point = []\n        for j in range(n_class):\n            flat_volume = volume[j].reshape(-1)\n            max_idx = torch.argmax(flat_volume)\n\n            max_idxs = []\n            for k in [-5, -3, -1, 1, 3]:\n                idx_start = max_idx + k*(h*w)//2\n                idx_start = max(0, idx_start)\n                idx_start = min(idx_start, flat_volume.shape[0]-h*w)\n                max_idxs.append(idx_start + torch.argmax(flat_volume[idx_start:idx_start + h*w]))\n\n            z = max_idx // (h * w)\n            y = [(idx % (h * w)) // w for idx in max_idxs]\n            x = [(idx % (h * w)) % w for idx in max_idxs]\n            point.append(torch.stack([z] +  y + x))\n\n        point = torch.stack(point, dim = 0)\n        points.append(point)\n\n    points = torch.stack(points, dim = 0)\n    return points.numpy()\n\nif __name__ == \"__main__\":\n    volumes = torch.zeros([32, 512, 512], dtype = torch.float)\n    volumes[16, 256, 256] = 1.0\n    \n    points = volume2point(volumes[None, None])\n    print('points : ', points)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:33:13.992922Z","iopub.execute_input":"2024-10-08T15:33:13.99319Z","iopub.status.idle":"2024-10-08T15:33:14.06359Z","shell.execute_reply.started":"2024-10-08T15:33:13.993167Z","shell.execute_reply":"2024-10-08T15:33:14.062685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## preprocess","metadata":{}},{"cell_type":"code","source":"def preprocess(args):\n    test_series_descriptions = pd.read_csv(args.root + f'{args.mode}_series_descriptions.csv')\n    \n    test = pd.DataFrame()\n    test['study_id'] = list(test_series_descriptions['study_id'].unique())\n    \n    return test, test_series_descriptions\n\n\nif __name__ == \"__main__\":\n    test, test_series_descriptions = preprocess(args)    ","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:33:14.06555Z","iopub.execute_input":"2024-10-08T15:33:14.065845Z","iopub.status.idle":"2024-10-08T15:33:14.084284Z","shell.execute_reply.started":"2024-10-08T15:33:14.06582Z","shell.execute_reply":"2024-10-08T15:33:14.083363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## dataset","metadata":{}},{"cell_type":"code","source":"def convert_to_8bit(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 * 255).astype(\"uint8\")\n\ndef get_imgs(dcms):\n    imgs = []\n    for dcm in dcms:\n        img = convert_to_8bit(dcm.pixel_array)\n        imgs.append(img)\n\n    try:\n        return np.stack(imgs, axis = 0)\n    except:\n        h, w = imgs[0].shape\n        imgs = [A.Resize(h, w)(image = x)['image'] for x in imgs]\n        return np.stack(imgs, axis = 0)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:33:14.085486Z","iopub.execute_input":"2024-10-08T15:33:14.085771Z","iopub.status.idle":"2024-10-08T15:33:14.09396Z","shell.execute_reply.started":"2024-10-08T15:33:14.085747Z","shell.execute_reply":"2024-10-08T15:33:14.093046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomDataset(torch.utils.data.Dataset):\n    def __init__(self, args, test, test_series_descriptions, orientation):\n        self.args = args\n        \n        self.test = test\n        self.test_series_descriptions = test_series_descriptions\n        \n        self.orientation = orientation\n        \n        self.volume_size = {\n            'Sagittal T2/STIR' : [29, 256, 256],\n            'Sagittal T1' : [38, 256, 256],\n            'Axial T2' : [192, 256, 256]\n        }[orientation]\n    \n    def __len__(self):\n        return len(self.test)\n\n    def get_subvolume(self, study_id, series_id):\n        dicoms = sorted(glob.glob(self.args.root + f'/{self.args.mode}_images/{study_id}/{series_id}/*.dcm'), key = lambda x: int(x.split('/')[-1].split('.')[0]))     \n        dicoms = [pydicom.dcmread(x) for x in dicoms]\n        \n        if 'Sagittal' in self.orientation:\n            pos = np.asarray([dicom.ImagePositionPatient for dicom in dicoms])[:, 0]\n        else:\n            pos = np.asarray([dicom.ImagePositionPatient for dicom in dicoms])[:, -1]\n\n        inputs = get_imgs(dicoms)\n        inputs = A.Resize(self.volume_size[1], self.volume_size[2])(image = inputs.transpose(1, 2, 0))['image'].transpose(2, 0, 1)\n        inputs = torch.tensor(inputs, dtype = torch.float)\n        inputs = inputs / 255.0\n\n        return inputs, pos\n    \n    def get_inputs(self, row):\n        row_series_descriptions = self.test_series_descriptions[self.test_series_descriptions['study_id'] == row['study_id']]\n        row_series_descriptions = row_series_descriptions[row_series_descriptions['series_description'] == self.orientation].reset_index(drop = True)\n        \n        if len(row_series_descriptions) > 0:\n            inputs, pos = [], []\n            for i in range(len(row_series_descriptions)):\n                study_id, series_id, _ = row_series_descriptions.loc[i]\n                \n                _inputs, _pos = self.get_subvolume(study_id, series_id)\n                inputs.append(_inputs)\n                pos.append(_pos)\n\n            inputs = torch.cat(inputs, dim = 0)\n            pos = np.concatenate(pos, axis = 0)\n        else:\n            inputs = torch.zeros(self.volume_size, dtype = torch.float32)\n            pos = np.zeros([self.volume_size[0]], dtype = np.float32)\n        \n        pos = np.argsort(pos)\n        inputs = inputs[pos]\n        \n        if inputs.shape[0] > self.volume_size[0]:\n            inputs = F.interpolate(inputs.unsqueeze(0).unsqueeze(0), size = self.volume_size, mode = 'trilinear').squeeze(0).squeeze(0)\n        \n        inputs = torch.cat([inputs, torch.zeros([self.volume_size[0] - inputs.shape[0]] + list(inputs.shape[1:]), dtype = torch.float)], dim = 0)\n        return inputs\n    \n    \n    def __getitem__(self, index):\n        row = self.test.loc[index]\n        inputs = self.get_inputs(row)\n        return inputs\n        \nif __name__ == \"__main__\":\n    test, test_series_descriptions = preprocess(args) \n    \n    orientation = 'Axial T2'\n    dataset = CustomDataset(args, test, test_series_descriptions, orientation)\n    \n    \n    inputs = dataset[0]\n    print('inputs : ', inputs.shape)\n    \n    fig, axes = plt.subplots(1, inputs.shape[0], figsize = (inputs.shape[0], 1))\n    for i in range(inputs.shape[0]):\n        axes[i].imshow(inputs[i], cmap = 'gray')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:33:14.095228Z","iopub.execute_input":"2024-10-08T15:33:14.095679Z","iopub.status.idle":"2024-10-08T15:33:31.929347Z","shell.execute_reply.started":"2024-10-08T15:33:14.095654Z","shell.execute_reply":"2024-10-08T15:33:31.928286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## model","metadata":{}},{"cell_type":"code","source":"class Segmenter(nn.Module):\n    def __init__(self,\n                 n_class,\n                 n_channel,\n                 volume_size,\n                 backbone = 'tf_efficientnet_b5_ns',\n                 n_blocks = 4,\n                 ):\n        super(Segmenter, self).__init__()\n\n        \n        self.n_class = n_class\n        self.n_channel = n_channel\n        self.volume_size = volume_size\n\n        self.extracter = timm.create_model(backbone, \n                              pretrained=False,\n                              features_only = True,\n                              out_indices=range(n_blocks),\n                                      )\n\n        encoder_channels = [n_channel] + [self.extracter.feature_info[i][\"num_chs\"] for i in range(n_blocks)]\n        decoder_channels = [256, 128, 64, 32, 16, 8][:n_blocks]\n\n        self.decoder = UnetDecoder(\n            encoder_channels = encoder_channels,\n            decoder_channels = decoder_channels,\n            n_blocks = n_blocks,\n            use_batchnorm = True,\n            center = False,\n            attention_type = None,\n        )\n\n        self.head = SegmentationHead(\n            in_channels = decoder_channels[-1],\n            out_channels = self.n_class,\n            activation = None,\n            kernel_size= 3,\n        )\n\n    def forward(self, x):\n        x = x.reshape(-1, self.n_channel, self.volume_size[1], self.volume_size[2])\n        _x = self.extracter(x)\n\n        x = self.decoder(*[x] + _x)\n        x = self.head(x)\n\n        x = x.reshape(-1, self.volume_size[0], self.n_class, self.volume_size[1], self.volume_size[2])\n        x = x.permute(0, 2, 1, 3, 4)\n        return x\n    \nclass Pointer(nn.Module):\n    def __init__(self,\n                 n_class,\n                 n_channel,\n                 volume_size,\n                 hidden_size,\n                 drop_rate,\n                 is_lstm = True,\n                 ):\n        super(Pointer, self).__init__()\n\n        self.n_channel = n_channel\n        self.volume_size = volume_size\n        \n\n\n        self.is_lstm = is_lstm\n\n\n        if self.is_lstm:\n    \n            self.extractor = timm.create_model('regnety_002', \n                                  pretrained=False,\n                                  features_only = True,\n                                  in_chans = n_channel,\n                                              )\n    \n            self.dense_size = self.extractor.feature_info[-1][\"num_chs\"]\n            self.rnn = nn.LSTM(\n                input_size = self.dense_size,\n                hidden_size = hidden_size,\n                batch_first = True,\n                bidirectional = True\n                )\n            self.out = nn.Sequential(\n                    nn.Dropout(p = drop_rate),\n                    nn.Linear(2 * hidden_size, n_class)\n                )\n        else:\n            \n            self.extracter = timm.create_model('regnety_002', \n                                  pretrained=False,\n                                  features_only = True,\n                                  in_chans = n_channel,\n                                              )\n    \n            self.dense_size = self.extracter.feature_info[-1][\"num_chs\"]\n\n            self.dense = nn.Linear(self.dense_size, hidden_size)\n    \n            self.transformer = RobertaPreLayerNormModel(\n                RobertaPreLayerNormConfig(\n                    hidden_size = hidden_size,\n                    num_hidden_layers = 1,\n                    num_attention_heads = 4,\n                    intermediate_size = 4 * hidden_size,\n                    hidden_act = 'gelu',\n                    )\n                )\n            del self.transformer.embeddings.word_embeddings\n\n            self.out = nn.Sequential(\n                nn.Dropout(p = drop_rate),\n                nn.Linear(hidden_size, n_class)\n            )\n\n    def forward(self, x, mask):\n        x = x.reshape(-1, self.n_channel, self.volume_size[1], self.volume_size[2])\n        if self.is_lstm:\n            x = self.extractor(x)\n            x = x[-1].mean(dim = [2, 3])\n            x = x.reshape(-1, self.volume_size[0], self.dense_size)\n            x, _ =  self.rnn(x)\n        else:\n            x = self.extracter(x)\n            x = x[-1].mean(dim = [2, 3])\n            x = x.reshape(-1, self.volume_size[0], self.dense_size)\n            x = self.dense(x)\n            x = self.transformer(inputs_embeds = x, attention_mask = mask).last_hidden_state\n        x = self.out(x)\n\n        x = x.permute(0, 2, 1)\n        return x\n    \nclass CustomModel(nn.Module):\n    def __init__(self,\n                 args,\n                 orientation,\n                 n_channel = 3,\n                 hidden_size = 256,\n                 drop_rate = 0.3,                 \n                 ):\n\n        super(CustomModel, self).__init__()\n        self.args = args\n\n        self.orientation = orientation\n\n        n_class = {\n            'Sagittal T2/STIR' : 5,\n            'Sagittal T1' : 10,\n            'Axial T2' : 10,\n        }[orientation]\n\n        volume_size = {\n            'Sagittal T2/STIR' : [29, 256, 256],\n            'Sagittal T1' : [38, 256, 256],\n            'Axial T2' : [192, 256, 256]\n        }[orientation]\n\n        segmenter_backbone = {\n            'Sagittal T2/STIR' : 'tf_efficientnet_b5_ns',\n            'Sagittal T1' :'tf_efficientnet_b5_ns',\n            'Axial T2' : 'regnety_002'\n        }[orientation]\n\n        n_blocks = {\n            'Sagittal T2/STIR' : 4,\n            'Sagittal T1' :4,\n            'Axial T2' : 5\n        }[orientation]\n\n        is_lstm = {\n            'Sagittal T2/STIR' : True,\n            'Sagittal T1' :True,\n            'Axial T2' : False\n        }[orientation]\n\n        self.n_channel = n_channel\n        self.volume_size = volume_size\n\n        self.segmenter = Segmenter(\n            n_class = n_class,\n            n_channel = n_channel,\n            volume_size = volume_size,\n            backbone = segmenter_backbone,\n            n_blocks = n_blocks,\n            )\n\n        self.pointer = Pointer(\n            n_class = n_class,\n            n_channel = n_channel,\n            volume_size = volume_size,\n            hidden_size = hidden_size,\n            drop_rate = drop_rate,\n            is_lstm = is_lstm,\n        )\n\n\n    def get_inputs(self, x):\n        x = F.pad(x, (0, 0, 0, 0, (self.n_channel-1)//2, (self.n_channel-1)//2), \"constant\", 0)\n        x = [x[:, i:i+self.n_channel] for i in range(self.volume_size[0])]\n        x = torch.stack(x, dim = 1)\n        return x\n\n    def get_masks(self, x):\n        x = (x.sum(dim = [2, 3]) != 0).float()\n        return x\n\n    def get_outputs(self, x1, x2, mask):\n        x1 = torch.sigmoid(x1)\n        x2 = torch.sigmoid(x2)\n\n        x = x1 * x2[:, :, :, None, None] * mask[:, None, :, None, None]\n        return x\n\n    def forward(self, x):\n        mask = self.get_masks(x)\n\n        x = self.get_inputs(x)\n\n        x1 = self.segmenter(x)\n\n        x2 = self.pointer(x, mask)\n\n        return self.get_outputs(x1, x2, mask)\n\nif __name__ == \"__main__\":\n    test, test_series_descriptions = preprocess(args) \n    \n    orientation = 'Sagittal T2/STIR'\n    dataset = CustomDataset(args, test, test_series_descriptions, orientation)\n    \n    loader = torch.utils.data.DataLoader(dataset, \n                                         batch_size = args.batch_size, \n                                         num_workers = args.n_worker,\n                                         shuffle = False,\n                                         drop_last = False)\n    inputs = next(iter(loader))\n    inputs = inputs.to(args.device)\n\n    model = CustomModel(args, orientation)\n    model = model.to(args.device)\n\n    with torch.no_grad():\n        outputs = model(inputs)\n        print(outputs.shape)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:33:31.930829Z","iopub.execute_input":"2024-10-08T15:33:31.931162Z","iopub.status.idle":"2024-10-08T15:33:35.436716Z","shell.execute_reply.started":"2024-10-08T15:33:31.931133Z","shell.execute_reply":"2024-10-08T15:33:35.435567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## inference","metadata":{}},{"cell_type":"code","source":"def inference(args, models, loader, batch_size):\n\n    preds = np.zeros([len(loader.dataset), args.n_locations[orientation], 11], dtype = int)\n    for bi, inputs in enumerate(tqdm(loader)):\n        inputs = inputs.to(args.device)\n        with torch.no_grad():\n            outputs = torch.stack([model(inputs) for model in models], dim = 0).mean(0)\n\n        preds[(bi) * batch_size:(bi + 1) * batch_size, :, :] = volume2point2(outputs)\n    return preds\n\nif __name__ == \"__main__\":\n    args = CustomConfig()\n\n    for orientation in args.orientations:\n        print('orientation : ', orientation)\n        \n        models_weight = {\n            'Sagittal T2/STIR' : args.stage1_weights[0],\n            'Sagittal T1' : args.stage1_weights[1],\n            'Axial T2' : args.stage1_weights[2],\n        }[orientation]\n\n        save_name = {\n            'Sagittal T2/STIR' : f'sagittal_t2',\n            'Sagittal T1' : f'sagittal_t1',\n            'Axial T2' : f'axial_t2',\n        }[orientation]\n        \n        models = []\n        for model_weight in models_weight:\n\n            model = CustomModel(args, orientation)\n            model = model.to(args.device)\n            model.load_state_dict(torch.load(model_weight))\n            model.eval()\n            models.append(model)\n\n        test, test_series_descriptions = preprocess(args) \n        dataset = CustomDataset(args, test, test_series_descriptions, orientation)\n        \n        batch_size = args.batch_size\n        n_worker = args.n_worker\n        if orientation == 'Axial T2':\n            batch_size = 1\n            n_worker = 2\n        loader = torch.utils.data.DataLoader(dataset,\n                                             batch_size = batch_size,\n                                             num_workers = n_worker,\n                                             shuffle = False,\n                                             drop_last = False)\n\n        preds = inference(args, models, loader, batch_size)\n        np.save('/kaggle/working/' + save_name + f'_preds.npy', preds)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:33:35.438384Z","iopub.execute_input":"2024-10-08T15:33:35.438788Z","iopub.status.idle":"2024-10-08T15:34:07.502332Z","shell.execute_reply.started":"2024-10-08T15:33:35.438751Z","shell.execute_reply":"2024-10-08T15:34:07.501148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## stage2","metadata":{}},{"cell_type":"markdown","source":"## dataset","metadata":{}},{"cell_type":"code","source":"class CustomTransform(nn.Module):\n    def __init__(self):\n        super(CustomTransform, self).__init__()\n        self.transform = A.Compose([\n            A.Resize(128, 128, p = 1.0),\n        ])\n\n    def forward(self, x):\n        x = x.transpose(1, 2, 0)\n        x = self.transform(image = x)['image']\n        x = x.transpose(2, 0, 1)\n        x = torch.tensor(x, dtype = torch.float)\n        return x\n    \ndef convert_to_8bit(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 * 255).astype(\"uint8\")\n\ndef get_imgs(dcms):\n    imgs = []\n    for dcm in dcms:\n        img = convert_to_8bit(dcm.pixel_array)\n        imgs.append(img)\n\n    try:\n        return np.stack(imgs, axis = 0)\n    except:\n        h, w = imgs[0].shape\n        imgs = [A.Resize(h, w)(image = x)['image'] for x in imgs]\n        return np.stack(imgs, axis = 0)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:34:07.503737Z","iopub.execute_input":"2024-10-08T15:34:07.504049Z","iopub.status.idle":"2024-10-08T15:34:07.514978Z","shell.execute_reply.started":"2024-10-08T15:34:07.504019Z","shell.execute_reply":"2024-10-08T15:34:07.514073Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomDataset(torch.utils.data.Dataset):\n    def __init__(self, args, test, test_series_descriptions, detects):\n        self.args = args\n        \n        self.test = test\n        self.test_series_descriptions = test_series_descriptions\n        \n        self.detects = {\n            'Sagittal T2/STIR' : detects[0],\n            'Sagittal T1' : detects[1], \n            'Axial T2' : detects[2],  \n        }\n        \n        self.transform = CustomTransform()\n        \n        self.volume_sizes = {\n            'Sagittal T2/STIR' : [29, 512, 512],\n            'Sagittal T1' : [38, 512, 512],\n            'Axial T2' : [192, 512, 512]\n        }\n        \n\n    def __len__(self):\n        return len(self.test)\n\n\n    def get_subvolume(self, study_id, series_id, orientation):\n        dicoms = sorted(glob.glob(self.args.root + f'/{self.args.mode}_images/{study_id}/{series_id}/*.dcm'), key = lambda x: int(x.split('/')[-1].split('.')[0]))     \n        dicoms = [pydicom.dcmread(x) for x in dicoms]\n        \n        if 'Sagittal' in orientation:\n            pos = np.asarray([dicom.ImagePositionPatient for dicom in dicoms])[:, 0]\n        else:\n            pos = np.asarray([dicom.ImagePositionPatient for dicom in dicoms])[:, -1]\n\n        inputs = get_imgs(dicoms)\n        inputs = A.Resize(self.volume_sizes[orientation][1], self.volume_sizes[orientation][2])(image = inputs.transpose(1, 2, 0))['image'].transpose(2, 0, 1)\n        inputs = torch.tensor(inputs, dtype = torch.float)\n        inputs = inputs / 255.0\n\n        return inputs, pos\n    \n    def get_volume(self, row, orientation):\n        row_series_descriptions = self.test_series_descriptions[self.test_series_descriptions['study_id'] == row['study_id']]\n        row_series_descriptions = row_series_descriptions[row_series_descriptions['series_description'] == orientation].reset_index(drop = True)\n        \n        if len(row_series_descriptions) > 0:\n            inputs, pos = [], []\n            for i in range(len(row_series_descriptions)):\n                study_id, series_id, _ = row_series_descriptions.loc[i]\n                \n                _inputs, _pos = self.get_subvolume(study_id, series_id, orientation)\n                inputs.append(_inputs)\n                pos.append(_pos)\n\n            inputs = torch.cat(inputs, dim = 0)\n            pos = np.concatenate(pos, axis = 0)\n        else:\n            inputs = torch.zeros(self.volume_sizes[orientation], dtype = torch.float32)\n            pos = np.zeros([self.volume_sizes[orientation][0]], dtype = np.float32)\n\n\n        pos = np.argsort(pos)\n        inputs = inputs[pos]\n\n        if inputs.shape[0] > self.volume_sizes[orientation][0]:\n            inputs = F.interpolate(inputs.unsqueeze(0).unsqueeze(0), size = self.volume_sizes[orientation], mode = 'trilinear').squeeze(0).squeeze(0)\n        return inputs\n    \n    def get_inputs(self, row, index):\n        row_series_descriptions = self.test_series_descriptions[self.test_series_descriptions['study_id'] == row['study_id']]\n        \n        inputs = []\n        for orientation in args.orientations:\n            if orientation in list(row_series_descriptions['series_description']):\n                volume = self.get_volume(row, orientation)\n                detect = self.detects[orientation][index]\n\n        \n                crops = []\n                for i in range(detect.shape[0]):\n                    z = detect[i][0] * 1\n                    \n                    z_all = range(z-2, z+3)\n                    y_all = detect[i][1:6] * 2\n                    x_all = detect[i][6:11] * 2\n                    crop = []\n                    for z, y, x in zip(z_all, y_all, x_all):\n\n                        z = min(max(z, 0), volume.shape[0]-1)\n                        y_start, y_end = y - self.args.patch_size, y + self.args.patch_size\n                        if y_start<0:\n                            y_start, y_end = 0, 2*self.args.patch_size\n                        if y_end>=self.volume_sizes[orientation][1]:\n                            y_start, y_end = self.volume_sizes[orientation][1]-1 - 2*self.args.patch_size, self.volume_sizes[orientation][1]-1\n\n                        x_start, x_end = x - self.args.patch_size, x + self.args.patch_size\n                        if x_start<0:\n                            x_start, x_end = 0, 2*self.args.patch_size\n                        if x_end>=self.volume_sizes[orientation][2]:\n                            x_start, x_end = self.volume_sizes[orientation][2]-1 - 2*self.args.patch_size, self.volume_sizes[orientation][2]-1\n\n                        crop.append(volume[z, y_start:y_end, x_start:x_end])\n                \n                    crop = torch.stack(crop, dim = 0)\n                    crop = crop.numpy()\n                    crops.append(crop)\n                crops = np.concatenate(crops, axis = 0)\n            else:\n                crops = np.zeros([self.args.depth_size * self.args.n_locations[orientation], 2*self.args.patch_size, 2*self.args.patch_size])\n\n            crops = crops.astype(np.float32)\n            crops = self.transform(crops)\n            crops = crops.reshape(-1, self.args.depth_size, self.args.image_size, self.args.image_size)\n            inputs.append(crops)\n\n        inputs = np.concatenate(inputs, axis = 0)\n        inputs = torch.tensor(inputs, dtype = torch.float)\n        return inputs\n            \n                    \n    def __getitem__(self, index):\n        row = self.test.loc[index]\n        inputs = self.get_inputs(row, index)\n        return inputs\n        \nif __name__ == \"__main__\":\n    test, test_series_descriptions = preprocess(args) \n    \n    detects = [\n        np.load(f'sagittal_t2_preds.npy'),\n        np.load(f'sagittal_t1_preds.npy'),\n        np.load(f'axial_t2_preds.npy'),\n    ]\n    \n    dataset = CustomDataset(args, test, test_series_descriptions, detects)\n    inputs = dataset[0]\n    print('inputs : ', inputs.shape)\n    \n    inputs = [inputs[0:5], inputs[5:15], inputs[15:25]]\n    print(f'{args.orientations[0]} : ', inputs[0].shape)\n    print(f'{args.orientations[1]} : ', inputs[1].shape)\n    print(f'{args.orientations[2]} : ', inputs[2].shape)\n    \n    for x in inputs:\n        fig, axes = plt.subplots(1, x.shape[0], figsize = (x.shape[0], 1))\n        for i in range(x.shape[0]):\n            axes[i].imshow(x[i][x.shape[1]//2], cmap = 'gray')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:34:07.519149Z","iopub.execute_input":"2024-10-08T15:34:07.519477Z","iopub.status.idle":"2024-10-08T15:34:12.918505Z","shell.execute_reply.started":"2024-10-08T15:34:07.519443Z","shell.execute_reply":"2024-10-08T15:34:12.917668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## model","metadata":{}},{"cell_type":"code","source":"class FeatureExtractor(nn.Module):\n    def __init__(self, n_channel, hidden_size, model_name = 'convnext_small.in12k_ft_in1k'):\n        super(FeatureExtractor, self).__init__()\n        self.backbone = timm.create_model(\n            model_name = model_name,\n            pretrained = False,\n            num_classes = 0,\n            in_chans = n_channel\n            )\n\n    def forward(self, x):\n        x = self.backbone(x)\n        return x\n\nclass CustomPooler(nn.Module):\n    def __init__(self, hidden_size):\n        super(CustomPooler, self).__init__()\n        self.fc = nn.Sequential(\n            nn.Linear(hidden_size, hidden_size, bias = True),\n            nn.Tanh(),\n            nn.Linear(hidden_size, 1, bias = False),\n            nn.Softmax(dim = 1),\n        )\n\n    def forward(self, x):\n        _x = self.fc(x)\n        x = torch.sum(x * _x, dim = 1)\n        return x\n\nclass CustomModel(nn.Module):\n    def __init__(self,\n                 args,\n                 version,\n                 n_channel = 1,\n                 hidden_size = 768,\n                 drop_rate = 0.2,\n                 model_name = 'convnext_small.in12k_ft_in1k'\n                 ):\n\n        super(CustomModel, self).__init__()\n        self.args = args\n        self.version = version\n\n        self.hidden_size = hidden_size\n\n        self.cnn = FeatureExtractor(\n            n_channel = n_channel,\n            hidden_size = hidden_size,\n            model_name = model_name,\n            )\n\n\n        self.rnn1 = nn.LSTM(\n            input_size = hidden_size,\n            hidden_size = hidden_size//2,\n            batch_first = True,\n            bidirectional = True\n            )\n        \n        self.pooler = CustomPooler(\n            hidden_size = hidden_size,\n            )\n\n\n        self.rnn2 = nn.LSTM(\n            input_size = hidden_size,\n            hidden_size = hidden_size//2,\n            batch_first = True,\n            bidirectional = True\n            )\n\n        \n        self.out1 = nn.Sequential(\n            nn.Dropout(drop_rate),\n            nn.Linear(hidden_size, args.n_class)\n        )\n        self.out2 = nn.Sequential(\n            nn.Dropout(drop_rate),\n            nn.Linear(hidden_size, args.n_class)\n        )\n        self.out3 = nn.Sequential(\n            nn.Dropout(drop_rate),\n            nn.Linear(hidden_size, args.n_class)\n        )\n\n    def forward(self, x):\n        _, _, d, h, w = x.shape\n\n        x = x.reshape(-1, h, w)\n        x = x.unsqueeze(1)\n        x = self.cnn(x)\n        \n        if self.version == 1:\n            x = x.reshape(-1, d, self.hidden_size)\n            x, _ = self.rnn1(x)\n            x = self.pooler(x)\n            x = x.reshape(-1, 25, self.hidden_size)\n            x = x.reshape(-1, 5, 5, self.hidden_size)\n            x = x.permute(0, 2, 1, 3)\n            x = x.reshape(-1, 5, self.hidden_size)\n            x, _ = self.rnn2(x)\n\n        elif self.version == 2:\n            x = x.reshape(-1, d, self.hidden_size)\n            x, _ = self.rnn1(x)\n            x = self.pooler(x)\n            x = x.reshape(-1, 25, self.hidden_size)\n            x = x.reshape(-1, 5, 5, self.hidden_size)\n            x = x.permute(0, 2, 1, 3)\n            x = x.reshape(-1, 5, self.hidden_size)\n            x_, _ = self.rnn2(x)\n            x = x + x_\n        \n        elif self.version == 3:\n\n            x = x.reshape(-1, 5, 5, d, self.hidden_size)\n            x = x.permute(0, 2, 1, 3, 4)\n            x = x.reshape(-1, 5 * d, self.hidden_size)\n            x, _ = self.rnn1(x)\n            x = x.reshape(-1, d, self.hidden_size)\n            x = self.pooler(x)\n        \n        x = x.reshape(-1, 5, 5, self.hidden_size)\n        x = x.permute(0, 2, 1, 3)\n        x = x.reshape(-1, 25, self.hidden_size)\n\n        x1 = x[:, 0:5]\n        x2 = x[:, 5:15]\n        x3 = x[:, 15:]\n\n        x1 = self.out1(x1)\n        x2 = self.out2(x2)\n        x3 = self.out3(x3)\n\n        x = torch.cat([x1, x2, x3], dim = 1)\n        return x\n\nclass EnsembleModel(nn.Module):\n    def __init__(self, args):\n        super(EnsembleModel, self).__init__()\n        self.models = []\n        for i, path in enumerate(args.stage2_weights):\n            model = CustomModel(args, \n                                version = args.versions[i], \n                                model_name = args.model_names[i],\n                                hidden_size = args.hidden_sizes[i])\n            model = model.to(args.device)\n            model.load_state_dict(torch.load(path))\n            model.eval()\n            self.models.append(model)\n    \n\n    def forward(self, x):\n        x = [model(x) for model in self.models]\n        x = torch.stack(x, dim = 0)\n        \n        x = 0.75 * torch.mean(x[:15], dim = 0) + 0.25 * torch.mean(x[15:], dim = 0)\n        \n        return x\n\nif __name__ == \"__main__\":\n    test, test_series_descriptions = preprocess(args) \n    \n    detects = [\n        np.load(f'/kaggle/working/sagittal_t2_preds.npy'),\n        np.load(f'/kaggle/working/sagittal_t1_preds.npy'),\n        np.load(f'/kaggle/working/axial_t2_preds.npy'),\n    ]\n    \n    dataset = CustomDataset(args, test, test_series_descriptions, detects)\n    \n    loader = torch.utils.data.DataLoader(dataset, batch_size = args.batch_size, num_workers = args.n_worker)\n    inputs = next(iter(loader))\n    inputs = inputs.to(args.device)\n\n    model = EnsembleModel(args)\n    model = model.to(args.device)\n\n    with torch.no_grad():\n        outputs = model(inputs)\n        print(outputs.shape)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:34:12.919595Z","iopub.execute_input":"2024-10-08T15:34:12.91989Z","iopub.status.idle":"2024-10-08T15:35:42.669968Z","shell.execute_reply.started":"2024-10-08T15:34:12.919864Z","shell.execute_reply":"2024-10-08T15:35:42.668955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## inference","metadata":{}},{"cell_type":"code","source":"def get_submission(args, test, preds):\n    row_id, normal_mild, moderate, severe = [], [], [], []\n    for i in range(len(test)):\n        x = test.loc[i]\n        x = x.fillna('Nan')\n\n        for j in range(len(args.label_columns)):\n            pred = preds[i, j, :]\n\n            row_id.append(f'{x.study_id}_{args.label_columns[j]}')\n            normal_mild.append(pred[0])\n            moderate.append(pred[1])\n            severe.append(pred[2])\n\n    submission = pd.DataFrame()\n    submission['row_id'] = row_id\n    submission['normal_mild'] = normal_mild\n    submission['moderate'] = moderate\n    submission['severe'] = severe\n    return submission\n\ndef inference(args, model, loader):\n    model.eval()\n    \n    preds = torch.zeros((len(loader.dataset), len(args.label_columns), args.n_class), dtype = torch.float)\n    for bi, inputs in enumerate(tqdm(loader)):\n        inputs = inputs.to(args.device)\n\n        with torch.no_grad():\n            outputs = model(inputs)\n            \n        preds[args.batch_size * bi:args.batch_size * (bi + 1)] = outputs.detach().cpu()\n        \n    preds = nn.Softmax(dim = -1)(preds)\n    preds = preds.numpy()\n    \n    submission = get_submission(args, loader.dataset.test, preds)\n    return submission\n\nif __name__ == \"__main__\":\n    args = CustomConfig()\n    \n    test, test_series_descriptions = preprocess(args) \n    \n    detects = [\n        np.load(f'/kaggle/working/sagittal_t2_preds.npy'),\n        np.load(f'/kaggle/working/sagittal_t1_preds.npy'),\n        np.load(f'/kaggle/working/axial_t2_preds.npy'),\n    ]\n    \n    dataset = CustomDataset(args, test, test_series_descriptions, detects)\n    loader = torch.utils.data.DataLoader(dataset, batch_size = args.batch_size, num_workers = args.n_worker)\n    \n    model = EnsembleModel(args)\n    \n    submission = inference(args, model, loader)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:35:42.671341Z","iopub.execute_input":"2024-10-08T15:35:42.671643Z","iopub.status.idle":"2024-10-08T15:36:21.861407Z","shell.execute_reply.started":"2024-10-08T15:35:42.671617Z","shell.execute_reply":"2024-10-08T15:36:21.860223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# submission","metadata":{}},{"cell_type":"code","source":"submission.to_csv('submission.csv', index = False)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:36:21.862886Z","iopub.execute_input":"2024-10-08T15:36:21.863205Z","iopub.status.idle":"2024-10-08T15:36:21.871583Z","shell.execute_reply.started":"2024-10-08T15:36:21.863175Z","shell.execute_reply":"2024-10-08T15:36:21.870651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:36:21.872723Z","iopub.execute_input":"2024-10-08T15:36:21.872976Z","iopub.status.idle":"2024-10-08T15:36:21.894807Z","shell.execute_reply.started":"2024-10-08T15:36:21.872955Z","shell.execute_reply":"2024-10-08T15:36:21.893989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm *.npy\ndel model, dataset, loader\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:36:21.895824Z","iopub.execute_input":"2024-10-08T15:36:21.896077Z","iopub.status.idle":"2024-10-08T15:36:22.595389Z","shell.execute_reply.started":"2024-10-08T15:36:21.896055Z","shell.execute_reply":"2024-10-08T15:36:22.594372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## config","metadata":{}},{"cell_type":"code","source":"class CustomConfig:\n    device = 'cuda'\n    root = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\n    \n    orientations = ['Sagittal T2/STIR', 'Sagittal T1', 'Axial T2']\n    \n    stage1_weights = {\n        'Sagittal T2/STIR' : [\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch020-trainloss0.16939-testloss0.181416-testscore1.332644.bin',\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch020-trainloss0.169008-testloss0.182547-testscore0.98592.bin',\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch020-trainloss0.168994-testloss0.181697-testscore1.278035.bin',\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch020-trainloss0.168799-testloss0.183231-testscore1.15847.bin',\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch020-trainloss0.169104-testloss0.181137-testscore1.341168.bin',\n        ],\n        'Sagittal T1' : [\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch020-trainloss0.183588-testloss0.189998-testscore1.199237.bin',\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch020-trainloss0.183153-testloss0.187536-testscore1.221008.bin',\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch020-trainloss0.182998-testloss0.191108-testscore1.184988.bin',\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch020-trainloss0.180799-testloss0.203025-testscore1.132657.bin',\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch015-trainloss0.184912-testloss0.192181-testscore1.2831.bin',\n        ],\n        'Axial T2' : [\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch020-trainloss0.101603-testloss0.104639-testscore1.152874.bin',\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch010-trainloss0.104989-testloss0.108265-testscore1.183024.bin',\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch020-trainloss0.100135-testloss0.107046-testscore1.109482.bin',\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch020-trainloss0.10015-testloss0.105416-testscore1.077373.bin',\n            '/kaggle/input/rsna-lsdc-stage1-weights/epoch020-trainloss0.100486-testloss0.106836-testscore1.154469.bin',\n        ]\n    }\n    \n    n_locations = {\n        'Sagittal T2/STIR' : 5,\n        'Sagittal T1' : 10,\n        'Axial T2' : 10,\n    }\n    \n    batch_size = 4\n    n_worker = 4\n    \n    n_class = 3\n    \n    label_columns = pd.read_csv(root + 'train.csv').columns[1:]\n    \n    stage2_weights = [\n        '/kaggle/input/rsna-stage2-weights/v1/2024/fold2/epoch003-trainloss0.50246-testloss0.384054-testscore0.384054.bin',\n        '/kaggle/input/rsna-stage2-weights/v1/2024/fold4/epoch012-trainloss0.510078-testloss0.435106-testscore0.400931.bin',\n        '/kaggle/input/rsna-stage2-weights/v1/2024/fold5/epoch017-trainloss0.479376-testloss0.422528-testscore0.382907.bin',\n        '/kaggle/input/rsna-stage2-weights/v1/42/fold1/epoch003-trainloss0.501287-testloss0.395612-testscore0.395612.bin',\n        '/kaggle/input/rsna-stage2-weights/v1/42/fold3/epoch005-trainloss0.480527-testloss0.361974-testscore0.361974.bin',\n\n        '/kaggle/input/rsna-stage2-weights/v2/2024/fold4/epoch014-trainloss0.489339-testloss0.430131-testscore0.401807.bin',\n        '/kaggle/input/rsna-stage2-weights/v2/42/fold1/epoch006-trainloss0.518346-testloss0.39402-testscore0.39402.bin',\n        '/kaggle/input/rsna-stage2-weights/v2/42/fold2/epoch003-trainloss0.512317-testloss0.387429-testscore0.387429.bin',\n        '/kaggle/input/rsna-stage2-weights/v2/42/fold3/epoch005-trainloss0.487169-testloss0.360147-testscore0.360147.bin',\n        '/kaggle/input/rsna-stage2-weights/v2/42/fold5/epoch003-trainloss0.495823-testloss0.379182-testscore0.379182.bin',\n\n        '/kaggle/input/rsna-stage2-weights/v3/2024/fold2/epoch003-trainloss0.521357-testloss0.389746-testscore0.389746.bin',\n        '/kaggle/input/rsna-stage2-weights/v3/2024/fold3/epoch005-trainloss0.491033-testloss0.366064-testscore0.366064.bin',\n        '/kaggle/input/rsna-stage2-weights/v3/2024/fold4/epoch004-trainloss0.506416-testloss0.401614-testscore0.401614.bin',\n        '/kaggle/input/rsna-stage2-weights/v3/2024/fold5/epoch006-trainloss0.50526-testloss0.381305-testscore0.381305.bin',\n        '/kaggle/input/rsna-stage2-weights/v3/42/fold1/epoch004-trainloss0.506298-testloss0.392619-testscore0.392619.bin',\n\n        \n       '/kaggle/input/rsna-lsdc-stage2-weights/epoch004-trainloss0.474929-testloss0.394903-testscore0.394903.bin',\n       '/kaggle/input/rsna-lsdc-stage2-weights/epoch006-trainloss0.485093-testloss0.38113-testscore0.38113.bin',\n       '/kaggle/input/rsna-lsdc-stage2-weights/epoch006-trainloss0.487418-testloss0.367523-testscore0.367523.bin',\n       '/kaggle/input/rsna-lsdc-stage2-weights/epoch003-trainloss0.490171-testloss0.387255-testscore0.387255.bin',\n       '/kaggle/input/rsna-lsdc-stage2-weights/epoch006-trainloss0.491637-testloss0.378409-testscore0.378409.bin',\n        \n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch010-trainloss0.488406-testloss0.400828-testscore0.400829.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch005-trainloss0.478115-testloss0.382623-testscore0.382623.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch010-trainloss0.490493-testloss0.368579-testscore0.368579.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch004-trainloss0.481913-testloss0.391444-testscore0.391444.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch006-trainloss0.479494-testloss0.381912-testscore0.381912.bin',\n        \n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch010-trainloss0.472179-testloss0.398393-testscore0.398393.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch006-trainloss0.481968-testloss0.384406-testscore0.384406.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch010-trainloss0.482629-testloss0.372525-testscore0.372525.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch005-trainloss0.484465-testloss0.389893-testscore0.389893.bin',\n        '/kaggle/input/rsna-lsdc-stage2-weights/epoch006-trainloss0.483708-testloss0.381362-testscore0.381362.bin',\n    \n    ]\n    \n\n    model_names = ['convnext_small.in12k_ft_in1k']*15 + \\\n                  ['convnext_tiny.in12k_ft_in1k'] * 5 + \\\n                  ['caformer_s18.sail_in22k_ft_in1k'] * 5 + \\\n                  ['pvt_v2_b3.in1k'] * 5\n    \n    versions = [1]*5 + [2]*5 + [3]*5 + [1]*15\n    \n    hidden_sizes = [768] * 20 + [512] * 10\n    \n    \n    \n    patch_size = 32\n    image_size = 128\n    depth_size = 1 + 2*2\n    \n    mode = 'test'\n    \nif __name__ == \"__main__\":\n    args = CustomConfig()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:36:56.065984Z","iopub.execute_input":"2024-10-08T15:36:56.066451Z","iopub.status.idle":"2024-10-08T15:36:56.098634Z","shell.execute_reply.started":"2024-10-08T15:36:56.066414Z","shell.execute_reply":"2024-10-08T15:36:56.097866Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# stage1","metadata":{}},{"cell_type":"markdown","source":"## utils","metadata":{}},{"cell_type":"code","source":"def volume2point(volumes):\n    volumes = volumes.cpu()\n\n    batch_size, n_class, _, h, w = volumes.shape\n\n    points = []\n    for i in range(batch_size):\n        volume = volumes[i]\n\n        point = []\n        for j in range(n_class):\n            flat_volume = volume[j].reshape(-1)\n            max_idx = torch.argmax(flat_volume)\n\n            z = max_idx // (h * w)\n            y = (max_idx % (h * w)) // w\n            x = (max_idx % (h * w)) % w\n            point.append(torch.stack([z, y, x]))\n\n        point = torch.stack(point, dim = 0)\n        points.append(point)\n\n    points = torch.stack(points, dim = 0)\n    return points.numpy()\n\nif __name__ == \"__main__\":\n    volumes = torch.zeros([32, 512, 512], dtype = torch.float)\n    volumes[16, 256, 256] = 1.0\n    \n    points = volume2point(volumes[None, None])\n    print('points : ', points)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:36:56.099782Z","iopub.execute_input":"2024-10-08T15:36:56.100118Z","iopub.status.idle":"2024-10-08T15:36:56.134052Z","shell.execute_reply.started":"2024-10-08T15:36:56.100088Z","shell.execute_reply":"2024-10-08T15:36:56.133176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## preprocess","metadata":{}},{"cell_type":"code","source":"def preprocess(args):\n    test_series_descriptions = pd.read_csv(args.root + f'{args.mode}_series_descriptions.csv')\n    \n    test = pd.DataFrame()\n    test['study_id'] = list(test_series_descriptions['study_id'].unique())\n    \n    return test, test_series_descriptions\n\n\nif __name__ == \"__main__\":\n    test, test_series_descriptions = preprocess(args)    ","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:36:56.135063Z","iopub.execute_input":"2024-10-08T15:36:56.135339Z","iopub.status.idle":"2024-10-08T15:36:56.143708Z","shell.execute_reply.started":"2024-10-08T15:36:56.135299Z","shell.execute_reply":"2024-10-08T15:36:56.142858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## dataset","metadata":{}},{"cell_type":"code","source":"def convert_to_8bit(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 * 255).astype(\"uint8\")\n\ndef get_imgs(dcms):\n    imgs = []\n    for dcm in dcms:\n        img = convert_to_8bit(dcm.pixel_array)\n        imgs.append(img)\n\n    try:\n        return np.stack(imgs, axis = 0)\n    except:\n        h, w = imgs[0].shape\n        imgs = [A.Resize(h, w)(image = x)['image'] for x in imgs]\n        return np.stack(imgs, axis = 0)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:36:56.144775Z","iopub.execute_input":"2024-10-08T15:36:56.145071Z","iopub.status.idle":"2024-10-08T15:36:56.153528Z","shell.execute_reply.started":"2024-10-08T15:36:56.145048Z","shell.execute_reply":"2024-10-08T15:36:56.152655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomDataset(torch.utils.data.Dataset):\n    def __init__(self, args, test, test_series_descriptions, orientation):\n        self.args = args\n        \n        self.test = test\n        self.test_series_descriptions = test_series_descriptions\n        \n        self.orientation = orientation\n        \n        self.volume_size = {\n            'Sagittal T2/STIR' : [32, 256, 256],\n            'Sagittal T1' : [32, 256, 256],\n            'Axial T2' : [64, 256, 256]\n        }[orientation]\n        \n    def __len__(self):\n        return len(self.test)\n    \n    def get_subvolume(self, study_id, series_id):\n        dicoms = sorted(glob.glob(self.args.root + f'{self.args.mode}_images/{study_id}/{series_id}/*.dcm'), key = lambda x: int(x.split('/')[-1].split('.')[0]))     \n        dicoms = [pydicom.dcmread(x) for x in dicoms]\n        \n        if 'Sagittal' in self.orientation:\n            pos = np.asarray([dicom.ImagePositionPatient for dicom in dicoms])[:, 0]\n        else:\n            pos = np.asarray([dicom.ImagePositionPatient for dicom in dicoms])[:, -1]\n\n        inputs = get_imgs(dicoms)\n        \n        inputs = A.Resize(self.volume_size[1], self.volume_size[2])(image = inputs.transpose(1, 2, 0))['image'].transpose(2, 0, 1)\n        inputs = torch.tensor(inputs, dtype = torch.float)\n        inputs = inputs / 255.0\n        return inputs, pos\n    \n    def get_inputs(self, row):\n        row_series_descriptions = self.test_series_descriptions[self.test_series_descriptions['study_id'] == row['study_id']]\n        row_series_descriptions = row_series_descriptions[row_series_descriptions['series_description'] == self.orientation].reset_index(drop = True)\n        \n        if len(row_series_descriptions) > 0:\n            inputs, pos = [], []\n            for i in range(len(row_series_descriptions)):\n                study_id, series_id, _ = row_series_descriptions.loc[i]\n                \n                _inputs, _pos = self.get_subvolume(study_id, series_id)\n                inputs.append(_inputs)\n                pos.append(_pos)\n                \n            inputs = torch.cat(inputs, dim = 0)\n            pos = np.concatenate(pos, axis = 0)\n            \n        else:\n            inputs = torch.zeros(self.volume_size, dtype = torch.float32)\n            pos = np.zeros([self.volume_size[0]], dtype = np.float32)\n            \n        pos = np.argsort(pos)\n        inputs = inputs[pos]\n        \n        if inputs.shape[0] > self.volume_size[0]:\n            inputs = F.interpolate(inputs.unsqueeze(0).unsqueeze(0), size = self.volume_size, mode = 'trilinear').squeeze(0).squeeze(0)\n        \n        inputs = torch.cat([inputs, torch.zeros([self.volume_size[0] - inputs.shape[0]] + list(inputs.shape[1:]), dtype = torch.float)], dim = 0)\n        return inputs\n    \n    def __getitem__(self, index):\n        row = self.test.loc[index]\n        inputs = self.get_inputs(row)\n        return inputs\n        \nif __name__ == \"__main__\":\n    test, test_series_descriptions = preprocess(args) \n    \n    orientation = 'Sagittal T2/STIR'\n    dataset = CustomDataset(args, test, test_series_descriptions, orientation)\n    \n    \n    inputs = dataset[0]\n    print('inputs : ', inputs.shape)\n    \n    fig, axes = plt.subplots(1, inputs.shape[0], figsize = (inputs.shape[0], 1))\n    for i in range(inputs.shape[0]):\n        axes[i].imshow(inputs[i], cmap = 'gray')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:36:56.154857Z","iopub.execute_input":"2024-10-08T15:36:56.155113Z","iopub.status.idle":"2024-10-08T15:37:00.22513Z","shell.execute_reply.started":"2024-10-08T15:36:56.155091Z","shell.execute_reply":"2024-10-08T15:37:00.22424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## model","metadata":{}},{"cell_type":"code","source":"class Segmenter(nn.Module):\n    def __init__(self,\n                 n_class,\n                 n_channel,\n                 volume_size,\n                 ):\n        super(Segmenter, self).__init__()\n        self.args = args\n\n        self.n_class = n_class\n        self.n_channel = n_channel\n        self.volume_size = volume_size\n\n        self.extracter = timm.create_model(\n            model_name = 'regnety_002',\n            pretrained = False,\n            features_only = True,\n            in_chans = n_channel,\n        )\n\n        encoder_channels = [n_channel] + [self.extracter.feature_info[i][\"num_chs\"] for i in range(len(self.extracter.feature_info))]\n        decoder_channels = [256, 128, 64, 32, 16]\n\n        self.decoder = UnetDecoder(\n            encoder_channels = encoder_channels,\n            decoder_channels = decoder_channels,\n            n_blocks = 5,\n            use_batchnorm = True,\n            center = False,\n            attention_type = None,\n        )\n\n        self.head = SegmentationHead(\n            in_channels = decoder_channels[-1],\n            out_channels = self.n_class,\n            activation = None,\n            kernel_size= 3,\n        )\n\n    def forward(self, x):\n        x = x.reshape(-1, self.n_channel, self.volume_size[1], self.volume_size[2])\n        _x = self.extracter(x)\n\n        x = self.decoder(*[x] + _x)\n        x = self.head(x)\n\n        x = x.reshape(-1, self.volume_size[0], self.n_class, self.volume_size[1], self.volume_size[2])\n        x = x.permute(0, 2, 1, 3, 4)\n        return x\n    \nclass Pointer(nn.Module):\n    def __init__(self,\n                 n_class,\n                 n_channel,\n                 volume_size,\n                 hidden_size,\n                 drop_rate,\n                 ):\n        super(Pointer, self).__init__()\n        self.args = args\n\n        self.n_channel = n_channel\n        self.volume_size = volume_size\n\n        self.extracter = timm.create_model(\n            model_name = 'regnety_002',\n            pretrained = False,\n            features_only = True,\n            in_chans = n_channel,\n        )\n\n        self.dense_size = self.extracter.feature_info[-1][\"num_chs\"]\n\n        self.dense = nn.Linear(self.dense_size, hidden_size)\n\n        self.rnn = nn.LSTM(\n            input_size = self.dense_size,\n            hidden_size = hidden_size,\n            batch_first = True,\n            bidirectional = True\n            )\n\n        self.out = nn.Sequential(\n            nn.Dropout(p = drop_rate),\n            nn.Linear(2 * hidden_size, n_class)\n        )\n\n    def forward(self, x, mask):\n        x = x.reshape(-1, self.n_channel, self.volume_size[1], self.volume_size[2])\n        x = self.extracter(x)\n        x = x[-1].mean(dim = [2, 3])\n\n        x = x.reshape(-1, self.volume_size[0], self.dense_size)\n        x, _ = self.rnn(x)\n        x = self.out(x)\n\n        x = x.permute(0, 2, 1)\n        return x\n    \nclass CustomModel(nn.Module):\n    def __init__(self,\n                 args,\n                 orientation,\n                 n_channel = 3,\n                 hidden_size = 256,\n                 drop_rate = 0.3,\n                 ):\n\n        super(CustomModel, self).__init__()\n        self.args = args\n\n        self.orientation = orientation\n\n        n_class = {\n            'Sagittal T2/STIR' : 5,\n            'Sagittal T1' : 10,\n            'Axial T2' : 10,\n        }[orientation]\n\n        volume_size = {\n            'Sagittal T2/STIR' : [32, 256, 256],\n            'Sagittal T1' : [32, 256, 256],\n            'Axial T2' : [64, 256, 256]\n        }[orientation]\n\n        self.n_channel = n_channel\n        self.volume_size = volume_size\n\n        self.segmenter = Segmenter(\n            n_class = n_class,\n            n_channel = n_channel,\n            volume_size = volume_size,\n            )\n\n        self.pointer = Pointer(\n            n_class = n_class,\n            n_channel = n_channel,\n            volume_size = volume_size,\n            hidden_size = hidden_size,\n            drop_rate = drop_rate,\n        )\n\n    def get_inputs(self, x):\n        x = F.pad(x, (0, 0, 0, 0, (self.n_channel-1)//2, (self.n_channel-1)//2), \"constant\", 0)\n        x = [x[:, i:i+self.n_channel] for i in range(self.volume_size[0])]\n        x = torch.stack(x, dim = 1)\n        return x\n\n    def get_masks(self, x):\n        x = (x.sum(dim = [2, 3]) != 0).float()\n        return x\n\n    def get_outputs(self, x1, x2, mask):\n        x1 = torch.sigmoid(x1)\n        x2 = torch.sigmoid(x2)\n\n        x = x1 * x2[:, :, :, None, None] * mask[:, None, :, None, None]\n        return x\n\n    def forward(self, x):\n        mask = self.get_masks(x)\n\n        x = self.get_inputs(x)\n\n        x1 = self.segmenter(x)\n\n        x2 = self.pointer(x, mask)\n\n        return self.get_outputs(x1, x2, mask)\n\nif __name__ == \"__main__\":\n    test, test_series_descriptions = preprocess(args) \n    \n    orientation = 'Sagittal T2/STIR'\n    dataset = CustomDataset(args, test, test_series_descriptions, orientation)\n    \n    loader = torch.utils.data.DataLoader(dataset, \n                                         batch_size = args.batch_size, \n                                         num_workers = args.n_worker,\n                                         shuffle = False,\n                                         drop_last = False)\n    inputs = next(iter(loader))\n    inputs = inputs.to(args.device)\n\n    model = CustomModel(args, orientation)\n    model = model.to(args.device)\n\n    with torch.no_grad():\n        outputs = model(inputs)\n        print(outputs.shape)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:37:00.226394Z","iopub.execute_input":"2024-10-08T15:37:00.226689Z","iopub.status.idle":"2024-10-08T15:37:02.26022Z","shell.execute_reply.started":"2024-10-08T15:37:00.226664Z","shell.execute_reply":"2024-10-08T15:37:02.259122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## inference","metadata":{}},{"cell_type":"code","source":"class EnsembleModel(nn.Module):\n    def __init__(self, args, orientation):\n        super(EnsembleModel, self).__init__()\n        self.models = []\n        for path in args.stage1_weights[orientation]:\n            model = CustomModel(args, orientation)\n            model = model.to(args.device)\n            model.load_state_dict(torch.load(path))\n            model.eval()\n            self.models.append(model)\n            \n\n    def forward(self, x):\n        x = [model(x) for model in self.models]\n        x = torch.stack(x, dim = 0)\n        return torch.mean(x, dim = 0)\n\ndef inference(args, model, loader):\n    model.eval()\n\n    preds = np.zeros([len(loader.dataset), args.n_locations[orientation], 3], dtype = int)\n    for bi, inputs in enumerate(tqdm(loader)):\n        inputs = inputs.to(args.device)\n        with torch.no_grad():\n            outputs = model(inputs)\n\n        preds[(bi) * args.batch_size:(bi + 1) * args.batch_size, :, :] = volume2point(outputs)\n    return preds\n\nif __name__ == \"__main__\":\n    args = CustomConfig()\n\n    for orientation in args.orientations:\n        print('orientation : ', orientation)\n\n        save_name = {\n            'Sagittal T2/STIR' : f'sagittal_t2',\n            'Sagittal T1' : f'sagittal_t1',\n            'Axial T2' : f'axial_t2',\n        }[orientation]\n\n        model = EnsembleModel(args, orientation)\n        model = model.to(args.device)\n\n        test, test_series_descriptions = preprocess(args) \n        dataset = CustomDataset(args, test, test_series_descriptions, orientation)\n        loader = torch.utils.data.DataLoader(dataset,\n                                             batch_size = args.batch_size,\n                                             num_workers = args.n_worker,\n                                             shuffle = False,\n                                             drop_last = False)\n\n        preds = inference(args, model, loader)\n        np.save('/kaggle/working/' + save_name + f'_preds.npy', preds)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:37:02.267011Z","iopub.execute_input":"2024-10-08T15:37:02.267419Z","iopub.status.idle":"2024-10-08T15:37:19.971915Z","shell.execute_reply.started":"2024-10-08T15:37:02.26739Z","shell.execute_reply":"2024-10-08T15:37:19.970702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## stage2","metadata":{}},{"cell_type":"markdown","source":"## dataset","metadata":{}},{"cell_type":"code","source":"class CustomTransform(nn.Module):\n    def __init__(self):\n        super(CustomTransform, self).__init__()\n        self.transform = A.Compose([\n            A.Resize(128, 128, p = 1.0),\n        ])\n\n    def forward(self, x):\n        x = x.transpose(1, 2, 0)\n        x = self.transform(image = x)['image']\n        x = x.transpose(2, 0, 1)\n        x = torch.tensor(x, dtype = torch.float)\n        return x\n    \ndef convert_to_8bit(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 * 255).astype(\"uint8\")\n\ndef get_imgs(dcms):\n    imgs = []\n    for dcm in dcms:\n        img = convert_to_8bit(dcm.pixel_array)\n        imgs.append(img)\n\n    try:\n        return np.stack(imgs, axis = 0)\n    except:\n        h, w = imgs[0].shape\n        imgs = [A.Resize(h, w)(image = x)['image'] for x in imgs]\n        return np.stack(imgs, axis = 0)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:37:19.973374Z","iopub.execute_input":"2024-10-08T15:37:19.973682Z","iopub.status.idle":"2024-10-08T15:37:19.984289Z","shell.execute_reply.started":"2024-10-08T15:37:19.973654Z","shell.execute_reply":"2024-10-08T15:37:19.983397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomDataset(torch.utils.data.Dataset):\n    def __init__(self, args, test, test_series_descriptions, detects):\n        self.args = args\n        \n        self.test = test\n        self.test_series_descriptions = test_series_descriptions\n        \n        self.detects = {\n            'Sagittal T2/STIR' : detects[0],\n            'Sagittal T1' : detects[1], \n            'Axial T2' : detects[2],  \n        }\n        \n        self.transform = CustomTransform()\n        \n        self.volume_sizes = {\n            'Sagittal T2/STIR' : [32, 512, 512],\n            'Sagittal T1' : [32, 512, 512],\n            'Axial T2' : [64, 512, 512]\n        }\n        \n\n    def __len__(self):\n        return len(self.test)\n    \n    def get_subvolume(self, study_id, series_id, orientation):\n        dicoms = sorted(glob.glob(self.args.root + f'{self.args.mode}_images/{study_id}/{series_id}/*.dcm'), key = lambda x: int(x.split('/')[-1].split('.')[0]))     \n        dicoms = [pydicom.dcmread(x) for x in dicoms]\n        \n        if 'Sagittal' in orientation:\n            pos = np.asarray([dicom.ImagePositionPatient for dicom in dicoms])[:, 0]\n        else:\n            pos = np.asarray([dicom.ImagePositionPatient for dicom in dicoms])[:, -1]\n\n        inputs = get_imgs(dicoms)\n        \n        inputs = A.Resize(self.volume_sizes[orientation][1], self.volume_sizes[orientation][2])(image = inputs.transpose(1, 2, 0))['image'].transpose(2, 0, 1)\n        inputs = torch.tensor(inputs, dtype = torch.float)\n        inputs = inputs / 255.0\n        return inputs, pos\n    \n    def get_volume(self, row, orientation):\n        row_series_descriptions = self.test_series_descriptions[self.test_series_descriptions['study_id'] == row['study_id']]\n        row_series_descriptions = row_series_descriptions[row_series_descriptions['series_description'] == orientation].reset_index(drop = True)\n\n        if len(row_series_descriptions) > 0:\n            inputs, pos = [], []\n            for i in range(len(row_series_descriptions)):\n                study_id, series_id, _ = row_series_descriptions.loc[i]\n                _inputs, _pos = self.get_subvolume(study_id, series_id, orientation)\n\n                inputs.append(_inputs)\n                pos.append(_pos)\n            inputs = torch.cat(inputs, dim = 0)\n            pos = np.concatenate(pos, axis = 0)\n        else:\n            inputs = torch.zeros(self.volume_sizes[orientation], dtype = torch.float32)\n            pos = np.zeros([self.volume_sizes[orientation][0]], dtype = np.float32)\n\n        pos = np.argsort(pos)\n        inputs = inputs[pos]\n\n        if inputs.shape[0] > self.volume_sizes[orientation][0]:\n            inputs = F.interpolate(inputs.unsqueeze(0).unsqueeze(0), size = self.volume_sizes[orientation], mode = 'trilinear').squeeze(0).squeeze(0)\n        return inputs\n    \n    def get_inputs(self, row, index):\n        row_series_descriptions = self.test_series_descriptions[self.test_series_descriptions['study_id'] == row['study_id']]\n        \n        inputs = []\n        for orientation in args.orientations:\n            if orientation in list(row_series_descriptions['series_description']):\n                volume = self.get_volume(row, orientation)\n                detect = self.detects[orientation][index]\n\n        \n                crops = []\n                for i in range(detect.shape[0]):\n                    z = detect[i][0] * 1\n                    y = detect[i][1] * 2\n                    x = detect[i][2] * 2\n                    \n                    z_start, z_end = max(z - (self.args.depth_size - 1)//2, 0), min(z + (self.args.depth_size - 1)//2 + 1, self.volume_sizes[orientation][0])\n                    y_start, y_end = max(y - self.args.patch_size, 0), min(y + self.args.patch_size, self.volume_sizes[orientation][1])\n                    x_start, x_end = max(x - self.args.patch_size, 0), min(x + self.args.patch_size, self.volume_sizes[orientation][2])\n  \n                    crop = volume[z_start:z_end, y_start:y_end, x_start:x_end]\n                    crop = F.interpolate(crop.unsqueeze(0).unsqueeze(0), size = [self.args.depth_size, 2*self.args.patch_size, 2*self.args.patch_size], mode = 'trilinear').squeeze(0).squeeze(0)\n                    crop = crop.numpy()\n                    crops.append(crop)\n                crops = np.concatenate(crops, axis = 0)\n            else:\n                crops = np.zeros([self.args.depth_size * self.args.n_locations[orientation], 2*self.args.patch_size, 2*self.args.patch_size])\n\n            crops = crops.astype(np.float32)\n            crops = self.transform(crops)\n            crops = crops.reshape(-1, self.args.depth_size, self.args.image_size, self.args.image_size)\n            inputs.append(crops)\n\n        inputs = np.concatenate(inputs, axis = 0)\n        inputs = torch.tensor(inputs, dtype = torch.float)\n        return inputs\n            \n                    \n    def __getitem__(self, index):\n        row = self.test.loc[index]\n        inputs = self.get_inputs(row, index)\n        return inputs\n        \nif __name__ == \"__main__\":\n    test, test_series_descriptions = preprocess(args) \n    \n    detects = [\n        np.load(f'/kaggle/working/sagittal_t2_preds.npy'),\n        np.load(f'/kaggle/working/sagittal_t1_preds.npy'),\n        np.load(f'/kaggle/working/axial_t2_preds.npy'),\n    ]\n    \n    dataset = CustomDataset(args, test, test_series_descriptions, detects)\n    inputs = dataset[0]\n    print('inputs : ', inputs.shape)\n    \n    inputs = [inputs[0:5], inputs[5:15], inputs[15:25]]\n    print(f'{args.orientations[0]} : ', inputs[0].shape)\n    print(f'{args.orientations[1]} : ', inputs[1].shape)\n    print(f'{args.orientations[2]} : ', inputs[2].shape)\n    \n    for x in inputs:\n        fig, axes = plt.subplots(1, x.shape[0], figsize = (x.shape[0], 1))\n        for i in range(x.shape[0]):\n            axes[i].imshow(x[i][x.shape[1]//2], cmap = 'gray')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:37:19.98607Z","iopub.execute_input":"2024-10-08T15:37:19.986483Z","iopub.status.idle":"2024-10-08T15:37:25.236534Z","shell.execute_reply.started":"2024-10-08T15:37:19.98645Z","shell.execute_reply":"2024-10-08T15:37:25.235534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## model","metadata":{}},{"cell_type":"code","source":"class FeatureExtractor(nn.Module):\n    def __init__(self, model_name, n_channel):\n        super(FeatureExtractor, self).__init__()\n        self.backbone = timm.create_model(\n            model_name = model_name,\n            pretrained = False,\n            num_classes = 0,\n            in_chans = n_channel\n            )\n\n    def forward(self, x):\n        x = self.backbone(x)\n        return x\n    \nclass CustomPooler(nn.Module):\n    def __init__(self, hidden_size):\n        super(CustomPooler, self).__init__()\n        self.fc = nn.Sequential(\n            nn.Linear(hidden_size, hidden_size, bias = True),\n            nn.Tanh(),\n            nn.Linear(hidden_size, 1, bias = False),\n            nn.Softmax(dim = 1),\n        )\n\n    def forward(self, x):\n        _x = self.fc(x)\n        x = torch.sum(x * _x, dim = 1)\n        return x\n\nclass CustomModel(nn.Module):\n    def __init__(self,\n                 args,\n                 version,\n                 n_channel = 1,\n                 hidden_size = 768,\n                 drop_rate = 0.2,\n                 model_name = 'convnext_small.in12k_ft_in1k'\n                 ):\n\n        super(CustomModel, self).__init__()\n        self.args = args\n        self.version = version\n\n        self.hidden_size = hidden_size\n\n        self.cnn = FeatureExtractor(\n            n_channel = n_channel,\n            model_name = model_name,\n            )\n\n\n        self.rnn1 = nn.LSTM(\n            input_size = hidden_size,\n            hidden_size = hidden_size//2,\n            batch_first = True,\n            bidirectional = True\n            )\n        \n        self.pooler = CustomPooler(\n            hidden_size = hidden_size,\n            )\n\n\n        self.rnn2 = nn.LSTM(\n            input_size = hidden_size,\n            hidden_size = hidden_size//2,\n            batch_first = True,\n            bidirectional = True\n            )\n\n        \n        self.out1 = nn.Sequential(\n            nn.Dropout(drop_rate),\n            nn.Linear(hidden_size, args.n_class)\n        )\n        self.out2 = nn.Sequential(\n            nn.Dropout(drop_rate),\n            nn.Linear(hidden_size, args.n_class)\n        )\n        self.out3 = nn.Sequential(\n            nn.Dropout(drop_rate),\n            nn.Linear(hidden_size, args.n_class)\n        )\n\n    def forward(self, x):\n        _, _, d, h, w = x.shape\n\n        x = x.reshape(-1, h, w)\n        x = x.unsqueeze(1)\n        x = self.cnn(x)\n        \n        if self.version == 1:\n            x = x.reshape(-1, d, self.hidden_size)\n            x, _ = self.rnn1(x)\n            x = self.pooler(x)\n            x = x.reshape(-1, 25, self.hidden_size)\n            x = x.reshape(-1, 5, 5, self.hidden_size)\n            x = x.permute(0, 2, 1, 3)\n            x = x.reshape(-1, 5, self.hidden_size)\n            x, _ = self.rnn2(x)\n\n        elif self.version == 2:\n            x = x.reshape(-1, d, self.hidden_size)\n            x, _ = self.rnn1(x)\n            x = self.pooler(x)\n            x = x.reshape(-1, 25, self.hidden_size)\n            x = x.reshape(-1, 5, 5, self.hidden_size)\n            x = x.permute(0, 2, 1, 3)\n            x = x.reshape(-1, 5, self.hidden_size)\n            x_, _ = self.rnn2(x)\n            x = x + x_\n        \n        elif self.version == 3:\n\n            x = x.reshape(-1, 5, 5, d, self.hidden_size)\n            x = x.permute(0, 2, 1, 3, 4)\n            x = x.reshape(-1, 5 * d, self.hidden_size)\n            x, _ = self.rnn1(x)\n            x = x.reshape(-1, d, self.hidden_size)\n            x = self.pooler(x)\n        \n        x = x.reshape(-1, 5, 5, self.hidden_size)\n        x = x.permute(0, 2, 1, 3)\n        x = x.reshape(-1, 25, self.hidden_size)\n\n        x1 = x[:, 0:5]\n        x2 = x[:, 5:15]\n        x3 = x[:, 15:]\n\n        x1 = self.out1(x1)\n        x2 = self.out2(x2)\n        x3 = self.out3(x3)\n\n        x = torch.cat([x1, x2, x3], dim = 1)\n        return x\n\nclass EnsembleModel(nn.Module):\n    def __init__(self, args):\n        super(EnsembleModel, self).__init__()\n        self.models = []\n        for i, path in enumerate(args.stage2_weights):\n            model = CustomModel(args, \n                                version = args.versions[i], \n                                model_name = args.model_names[i],\n                                hidden_size = args.hidden_sizes[i])\n            model = model.to(args.device)\n            model.load_state_dict(torch.load(path))\n            model.eval()\n            self.models.append(model)\n    \n\n    def forward(self, x):\n        x = [model(x) for model in self.models]\n        x = torch.stack(x, dim = 0)\n        \n        x = 0.25 * torch.mean(x[:15], dim = 0) + 0.75 * torch.mean(x[15:], dim = 0)\n        \n        return x\n\nif __name__ == \"__main__\":\n    test, test_series_descriptions = preprocess(args) \n    \n    detects = [\n        np.load(f'/kaggle/working/sagittal_t2_preds.npy'),\n        np.load(f'/kaggle/working/sagittal_t1_preds.npy'),\n        np.load(f'/kaggle/working/axial_t2_preds.npy'),\n    ]\n    \n    dataset = CustomDataset(args, test, test_series_descriptions, detects)\n    \n    loader = torch.utils.data.DataLoader(dataset, batch_size = args.batch_size, num_workers = args.n_worker)\n    inputs = next(iter(loader))\n    inputs = inputs.to(args.device)\n\n    model = EnsembleModel(args)\n    model = model.to(args.device)\n\n    with torch.no_grad():\n        outputs = model(inputs)\n        print(outputs.shape)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:39:42.92693Z","iopub.execute_input":"2024-10-08T15:39:42.927405Z","iopub.status.idle":"2024-10-08T15:40:21.015189Z","shell.execute_reply.started":"2024-10-08T15:39:42.927365Z","shell.execute_reply":"2024-10-08T15:40:21.014178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## inference","metadata":{}},{"cell_type":"code","source":"def get_submission(args, test, preds):\n    row_id, normal_mild, moderate, severe = [], [], [], []\n    for i in range(len(test)):\n        x = test.loc[i]\n        x = x.fillna('Nan')\n\n        for j in range(len(args.label_columns)):\n            pred = preds[i, j, :]\n\n            row_id.append(f'{x.study_id}_{args.label_columns[j]}')\n            normal_mild.append(pred[0])\n            moderate.append(pred[1])\n            severe.append(pred[2])\n\n    submission = pd.DataFrame()\n    submission['row_id'] = row_id\n    submission['normal_mild'] = normal_mild\n    submission['moderate'] = moderate\n    submission['severe'] = severe\n    return submission\n\ndef inference(args, model, loader):\n    model.eval()\n    \n    preds = torch.zeros((len(loader.dataset), len(args.label_columns), args.n_class), dtype = torch.float)\n    for bi, inputs in enumerate(tqdm(loader)):\n        inputs = inputs.to(args.device)\n\n        with torch.no_grad():\n            outputs = model(inputs)\n            \n        preds[args.batch_size * bi:args.batch_size * (bi + 1)] = outputs.detach().cpu()\n        \n    preds = nn.Softmax(dim = -1)(preds)\n    preds = preds.numpy()\n    \n    submission = get_submission(args, loader.dataset.test, preds)\n    return submission\n\nif __name__ == \"__main__\":\n    args = CustomConfig()\n    \n    test, test_series_descriptions = preprocess(args) \n    \n    detects = [\n        np.load(f'/kaggle/working/sagittal_t2_preds.npy'),\n        np.load(f'/kaggle/working/sagittal_t1_preds.npy'),\n        np.load(f'/kaggle/working/axial_t2_preds.npy'),\n    ]\n    \n    dataset = CustomDataset(args, test, test_series_descriptions, detects)\n    loader = torch.utils.data.DataLoader(dataset, batch_size = args.batch_size, num_workers = args.n_worker)\n    \n    model = EnsembleModel(args)\n    \n    submission = inference(args, model, loader)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:40:24.103397Z","iopub.execute_input":"2024-10-08T15:40:24.103803Z","iopub.status.idle":"2024-10-08T15:41:02.228417Z","shell.execute_reply.started":"2024-10-08T15:40:24.103769Z","shell.execute_reply":"2024-10-08T15:41:02.227404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Ensemble","metadata":{}},{"cell_type":"code","source":"!rm -r *.npy\nsubmission2 = pd.read_csv('submission.csv')\nsubmission = submission.sort_values('row_id')\nsubmission2 = submission2.sort_values('row_id')\ndisplay(submission)\ndisplay(submission2)\nsubmission[['normal_mild', 'moderate', 'severe']] = 0.4*submission[['normal_mild', 'moderate', 'severe']].values + 0.6*submission2[['normal_mild', 'moderate', 'severe']].values","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:41:02.230602Z","iopub.execute_input":"2024-10-08T15:41:02.230925Z","iopub.status.idle":"2024-10-08T15:41:02.262945Z","shell.execute_reply.started":"2024-10-08T15:41:02.230897Z","shell.execute_reply":"2024-10-08T15:41:02.262056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# submission","metadata":{}},{"cell_type":"code","source":"submission.to_csv('submission.csv', index = False)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:41:02.264167Z","iopub.execute_input":"2024-10-08T15:41:02.264666Z","iopub.status.idle":"2024-10-08T15:41:02.270637Z","shell.execute_reply.started":"2024-10-08T15:41:02.264633Z","shell.execute_reply":"2024-10-08T15:41:02.269742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:41:02.272655Z","iopub.execute_input":"2024-10-08T15:41:02.273119Z","iopub.status.idle":"2024-10-08T15:41:02.288509Z","shell.execute_reply.started":"2024-10-08T15:41:02.27307Z","shell.execute_reply":"2024-10-08T15:41:02.287524Z"},"trusted":true},"execution_count":null,"outputs":[]}]}