{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"DEBUG = False","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-10-15T13:58:28.487983Z","iopub.execute_input":"2023-10-15T13:58:28.48874Z","iopub.status.idle":"2023-10-15T13:58:28.497283Z","shell.execute_reply.started":"2023-10-15T13:58:28.488708Z","shell.execute_reply":"2023-10-15T13:58:28.496414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install dicomsdl --no-index --find-links=file:///kaggle/input/read-dicom-set","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:58:28.49897Z","iopub.execute_input":"2023-10-15T13:58:28.499369Z","iopub.status.idle":"2023-10-15T13:58:40.036002Z","shell.execute_reply.started":"2023-10-15T13:58:28.49934Z","shell.execute_reply":"2023-10-15T13:58:40.034938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install timm==0.9.2 segmentation-models-pytorch\nimport sys\nsys.path = [\n    '../input/segment-pytorch-model-0-3-3/segmentation_models.pytorch/',\n    '../input/pretrained-model-pytorch/pretrained-models.pytorch',\n    '../input/efficientnet-pytorch/EfficientNet-PyTorch'\n    # '../input/tim-0-9-2/pytorch-image-models'\n] + sys.path\n\n!cp -r ../input/timm-0-9-2/pytorch-image-models/timm ./timm4smp","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:58:40.038379Z","iopub.execute_input":"2023-10-15T13:58:40.038714Z","iopub.status.idle":"2023-10-15T13:58:42.923246Z","shell.execute_reply.started":"2023-10-15T13:58:40.03868Z","shell.execute_reply":"2023-10-15T13:58:42.922062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm4smp\nprint(timm4smp.__version__)\nimport segmentation_models_pytorch as smp\nprint(smp.__version__)","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:58:42.925282Z","iopub.execute_input":"2023-10-15T13:58:42.925663Z","iopub.status.idle":"2023-10-15T13:58:49.953636Z","shell.execute_reply.started":"2023-10-15T13:58:42.925625Z","shell.execute_reply":"2023-10-15T13:58:49.952597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nimport os\nimport ast\nimport cv2\nimport time\n# import timm\nimport timm4smp\nimport pickle\nimport random\nimport pydicom\nimport argparse\nimport warnings\nimport threading\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom tqdm import tqdm\nfrom glob import glob\nimport albumentations\nimport matplotlib.pyplot as plt\nimport segmentation_models_pytorch as smp\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.cuda.amp as amp\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom pylab import rcParams\n\n%matplotlib inline\ndevice = torch.device('cuda')\ntorch.backends.cudnn.benchmark = True\n\n# timm.__version__","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:58:49.95616Z","iopub.execute_input":"2023-10-15T13:58:49.956722Z","iopub.status.idle":"2023-10-15T13:58:51.801584Z","shell.execute_reply.started":"2023-10-15T13:58:49.956688Z","shell.execute_reply":"2023-10-15T13:58:51.800236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    data_dir = '../input/rsna-2023-abdominal-trauma-detection/'\n    image_size_seg = (128, 128, 128)\n    msk_size = image_size_seg[0]\n    backbone_seg = 'resnet18'\n    drop_rate = 0.\n    drop_path_rate = 0.\n    n_blocks = 4\n    out_dim_seg = 5\n    \n    # cls model\n    in_chans = 6\n    backbone_cls = 'tf_efficientnetv2_s_in21ft1k'\n    image_size_cls = 224\n    n_slice_per_c = 15\n    n_ch = 5\n\n    batch_size_seg = 1\n    num_workers = 2\n    device = 'cuda'\n    seed = 42","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:58:51.803077Z","iopub.execute_input":"2023-10-15T13:58:51.803569Z","iopub.status.idle":"2023-10-15T13:58:51.810251Z","shell.execute_reply.started":"2023-10-15T13:58:51.803535Z","shell.execute_reply":"2023-10-15T13:58:51.808764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seeding(SEED):\n    np.random.seed(SEED)\n    random.seed(SEED)\n    os.environ['PYTHONHASHSEED'] = str(SEED)\n    torch.manual_seed(SEED)\n    torch.cuda.manual_seed(SEED)\n    torch.cuda.manual_seed_all(SEED)\n    print('seeding done!!!')\nseeding(CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:58:51.811527Z","iopub.execute_input":"2023-10-15T13:58:51.81222Z","iopub.status.idle":"2023-10-15T13:58:51.8263Z","shell.execute_reply.started":"2023-10-15T13:58:51.812191Z","shell.execute_reply":"2023-10-15T13:58:51.825378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom(path):\n    dicom = pydicom.read_file(path)\n    data = dicom.pixel_array\n    data = cv2.resize(data, (CFG.image_size_seg[0], CFG.image_size_seg[1]), interpolation = cv2.INTER_LINEAR)\n    return data\n\ndef load_dicom_line_par(path):\n\n    t_paths = sorted(glob(os.path.join(path, \"*\")), key=lambda x: int(x.split('/')[-1].split(\".\")[0]))\n    \n    n_scans = len(t_paths)\n    \n    indices = np.quantile(list(range(n_scans)), np.linspace(0., 1., CFG.image_size_seg[2])).round().astype(int)\n    t_paths = [t_paths[i] for i in indices]\n\n    images = []\n    for filename in t_paths:\n        images.append(load_dicom(filename))\n    images = np.stack(images, -1)\n    \n    images = images - np.min(images)\n    images = images / (np.max(images) + 1e-4)\n    images = (images * 255).astype(np.uint8)\n\n    return images\n\nclass SegTestDataset(Dataset):\n    def __init__(self, df):\n        self.df = df.reset_index()\n    \n    def __len__(self):\n        return self.df.shape[0]\n    \n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        image = load_dicom_line_par(row.dicom_folder)\n        if image.ndim < 4:\n            image = np.expand_dims(image, 0)\n        image = image.astype(np.float32).repeat(3, 0)\n        image = image / 255.\n        return torch.tensor(image).float()","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:58:51.827476Z","iopub.execute_input":"2023-10-15T13:58:51.828753Z","iopub.status.idle":"2023-10-15T13:58:51.838778Z","shell.execute_reply.started":"2023-10-15T13:58:51.828723Z","shell.execute_reply":"2023-10-15T13:58:51.837867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def is_folder_empty(folder_path):\n    return os.path.exists(folder_path) == 0\n\ndf = pd.read_csv(os.path.join(CFG.data_dir, 'test_series_meta.csv'))\n\ndf['dicom_folder'] = CFG.data_dir + 'test_images/' + df.patient_id.astype(str) + '/' + df.series_id.astype(str)\ndf['IsEmpty'] = df['dicom_folder'].apply(is_folder_empty)\ndf = df[~df['IsEmpty']]\ndf = df.drop(columns=['IsEmpty'])\ndf = df.reset_index()\nprint(df.head())\n","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:58:51.83997Z","iopub.execute_input":"2023-10-15T13:58:51.840805Z","iopub.status.idle":"2023-10-15T13:58:51.887682Z","shell.execute_reply.started":"2023-10-15T13:58:51.840777Z","shell.execute_reply":"2023-10-15T13:58:51.886766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_seg = SegTestDataset(df)\nloader_seg = torch.utils.data.DataLoader(dataset_seg, batch_size=CFG.batch_size_seg, shuffle=False, num_workers=CFG.num_workers)","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:58:51.88872Z","iopub.execute_input":"2023-10-15T13:58:51.889226Z","iopub.status.idle":"2023-10-15T13:58:51.895649Z","shell.execute_reply.started":"2023-10-15T13:58:51.889198Z","shell.execute_reply":"2023-10-15T13:58:51.894837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rcParams['figure.figsize'] = 20,8\nfor i in range(1):\n    f, axarr = plt.subplots(1, 2)\n    for p in range(1, 2):\n        idx = i*4+p\n        img = dataset_seg[idx]\n        img = img[:, :, :, 60]\n        axarr[p].imshow(img.transpose(0, 1).transpose(1,2).squeeze())","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:58:51.899236Z","iopub.execute_input":"2023-10-15T13:58:51.899676Z","iopub.status.idle":"2023-10-15T13:58:53.488475Z","shell.execute_reply.started":"2023-10-15T13:58:51.899646Z","shell.execute_reply":"2023-10-15T13:58:53.487575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nimport timm4smp\nimport segmentation_models_pytorch as smp\nfrom timm.layers import Conv2dSame\n# from conv3d_same import Conv3dSame\n\n\nclass TimmSegModel(nn.Module):\n    def __init__(self, backbone, segtype='unet', pretrained=False):\n        super(TimmSegModel, self).__init__()\n\n        self.encoder = timm4smp.create_model(\n            backbone,\n            in_chans=3,\n            features_only=True,\n            drop_rate=CFG.drop_rate,\n            drop_path_rate=CFG.drop_path_rate,\n            pretrained=pretrained\n        )\n        g = self.encoder(torch.rand(1, 3, 64, 64))\n        encoder_channels = [1] + [_.shape[1] for _ in g]\n        decoder_channels = [256, 128, 64, 32, 16]\n        n_blocks = CFG.n_blocks\n        if segtype == 'unet':\n            self.decoder = smp.decoders.unet.decoder.UnetDecoder(\n                encoder_channels=encoder_channels[:n_blocks+1],\n                decoder_channels=decoder_channels[:n_blocks],\n                n_blocks=n_blocks,\n            )\n\n        self.segmentation_head = nn.Conv2d(decoder_channels[n_blocks-1], CFG.out_dim_seg, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))\n\n    def forward(self,x):\n        global_features = [0] + self.encoder(x)[:CFG.n_blocks]\n        seg_features = self.decoder(*global_features)\n        seg_features = self.segmentation_head(seg_features)\n        return seg_features\n\n\ndef convert_3d(module):\n    module_output = module\n    if isinstance(module, torch.nn.BatchNorm2d):\n        module_output = torch.nn.BatchNorm3d(\n            module.num_features,\n            module.eps,\n            module.momentum,\n            module.affine,\n            module.track_running_stats,\n        )\n        if module.affine:\n            with torch.no_grad():\n                module_output.weight = module.weight\n                module_output.bias = module.bias\n        module_output.running_mean = module.running_mean\n        module_output.running_var = module.running_var\n        module_output.num_batches_tracked = module.num_batches_tracked\n        if hasattr(module, \"qconfig\"):\n            module_output.qconfig = module.qconfig\n            \n    # elif isinstance(module, Conv2dSame):\n    #     module_output = Conv3dSame(\n    #         in_channels=module.in_channels,\n    #         out_channels=module.out_channels,\n    #         kernel_size=module.kernel_size[0],\n    #         stride=module.stride[0],\n    #         padding=module.padding[0],\n    #         dilation=module.dilation[0],\n    #         groups=module.groups,\n    #         bias=module.bias is not None,\n    #     )\n    #     module_output.weight = torch.nn.Parameter(module.weight.unsqueeze(-1).repeat(1,1,1,1,module.kernel_size[0]))\n\n    elif isinstance(module, torch.nn.Conv2d):\n        module_output = torch.nn.Conv3d(\n            in_channels=module.in_channels,\n            out_channels=module.out_channels,\n            kernel_size=module.kernel_size[0],\n            stride=module.stride[0],\n            padding=module.padding[0],\n            dilation=module.dilation[0],\n            groups=module.groups,\n            bias=module.bias is not None,\n            padding_mode=module.padding_mode\n        )\n        module_output.weight = torch.nn.Parameter(module.weight.unsqueeze(-1).repeat(1,1,1,1,module.kernel_size[0]))\n\n    elif isinstance(module, torch.nn.MaxPool2d):\n        module_output = torch.nn.MaxPool3d(\n            kernel_size=module.kernel_size,\n            stride=module.stride,\n            padding=module.padding,\n            dilation=module.dilation,\n            ceil_mode=module.ceil_mode,\n        )\n    elif isinstance(module, torch.nn.AvgPool2d):\n        module_output = torch.nn.AvgPool3d(\n            kernel_size=module.kernel_size,\n            stride=module.stride,\n            padding=module.padding,\n            ceil_mode=module.ceil_mode,\n        )\n\n    for name, child in module.named_children():\n        module_output.add_module(\n            name, convert_3d(child)\n        )\n    del module\n\n    return module_output\n\n\nm = TimmSegModel(CFG.backbone_seg)\nm = convert_3d(m)\nm(torch.rand(1, 3, 128, 128, 128)).shape","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:58:53.489897Z","iopub.execute_input":"2023-10-15T13:58:53.490475Z","iopub.status.idle":"2023-10-15T13:59:05.248566Z","shell.execute_reply.started":"2023-10-15T13:58:53.490441Z","shell.execute_reply":"2023-10-15T13:59:05.247646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# load model\nmodel_seg_path = '../input/stage1-rsna'\nseg_model = TimmSegModel(CFG.backbone_seg, pretrained=False)\nseg_model = convert_3d(seg_model)\nseg_model = seg_model.to(CFG.device)\nload_model_file = os.path.join(model_seg_path, 'resnet18_fold0_best.pth')\nsd = torch.load(load_model_file)\nif 'model_state_dict' in sd.keys():\n    sd = sd['model_state_dict']\nsd = {k[7:] if k.startswith('module.') else k: sd[k] for k in sd.keys()}\nseg_model.load_state_dict(sd, strict=True)\nseg_model.eval()\nprint('load segmodel successfull!')","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:59:05.250183Z","iopub.execute_input":"2023-10-15T13:59:05.250797Z","iopub.status.idle":"2023-10-15T13:59:10.762682Z","shell.execute_reply.started":"2023-10-15T13:59:05.250766Z","shell.execute_reply":"2023-10-15T13:59:10.761677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class Timm1BoneModel(nn.Module):\n#     def __init__(self, backbone, image_size, pretrained=False):\n#         super(Timm1BoneModel, self).__init__()\n#         self.image_size = image_size\n\n#         self.encoder = timm4smp.create_model(\n#             backbone,\n#             in_chans=CFG.in_chans,\n#             num_classes=1,\n#             features_only=False,\n#             drop_rate=0,\n#             drop_path_rate=0,\n#             pretrained=pretrained\n#         )\n\n#         if 'efficient' in backbone:\n#             hdim = self.encoder.conv_head.out_channels\n#             self.encoder.classifier = nn.Identity()\n#         elif 'convnext' in backbone or 'nfnet' in backbone:\n#             hdim = self.encoder.head.fc.in_features\n#             self.encoder.head.fc = nn.Identity()\n\n#         self.lstm = nn.LSTM(hdim, 256, num_layers=2, dropout=0, bidirectional=True, batch_first=True)\n#         self.head = nn.Sequential(\n#             nn.Linear(512, 256),\n#             nn.BatchNorm1d(256),\n#             nn.Dropout(0),\n#             nn.LeakyReLU(0.1),\n#             nn.Linear(256, 1),\n#         )\n\n\n#     def forward(self, x):  # (bs, nslice, ch, sz, sz)\n#         bs = x.shape[0]\n#         x = x.view(bs * CFG.n_slice_per_c, CFG.in_chans, self.image_size, self.image_size)\n#         feat = self.encoder(x)\n#         feat = feat.view(bs, CFG.n_slice_per_c, -1)\n#         feat, _ = self.lstm(feat)\n#         feat = feat.contiguous().view(bs * CFG.n_slice_per_c, -1)\n#         feat = self.head(feat)\n#         feat = feat.view(bs, CFG.n_slice_per_c).contiguous()\n\n#         return feat","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:59:10.7642Z","iopub.execute_input":"2023-10-15T13:59:10.765199Z","iopub.status.idle":"2023-10-15T13:59:10.77122Z","shell.execute_reply.started":"2023-10-15T13:59:10.765163Z","shell.execute_reply":"2023-10-15T13:59:10.770268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from itertools import repeat\n\n\nclass SpatialDropout(nn.Module):\n    def __init__(self, drop=0.5):\n        super(SpatialDropout, self).__init__()\n        self.drop = drop\n        \n    def forward(self, inputs, noise_shape=None):\n        \"\"\"\n        @param: inputs, tensor\n        @param: noise_shape, tuple\n        \"\"\"\n        outputs = inputs.clone()\n        if noise_shape is None:\n            noise_shape = (inputs.shape[0], *repeat(1, inputs.dim()-2), inputs.shape[-1]) \n        \n        self.noise_shape = noise_shape\n        if not self.training or self.drop == 0:\n            return inputs\n        else:\n            noises = self._make_noises(inputs)\n            if self.drop == 1:\n                noises.fill_(0.0)\n            else:\n                noises.bernoulli_(1 - self.drop).div_(1 - self.drop)\n            noises = noises.expand_as(inputs)    \n            outputs.mul_(noises)\n            return outputs\n            \n    def _make_noises(self, inputs):\n        return inputs.new().resize_(self.noise_shape)\n\n\nclass MLPAttentionNetwork(nn.Module):\n \n    def __init__(self, hidden_dim, attention_dim=None):\n        super(MLPAttentionNetwork, self).__init__()\n \n        self.hidden_dim = hidden_dim\n        self.attention_dim = attention_dim\n        if self.attention_dim is None:\n            self.attention_dim = self.hidden_dim\n        # W * x + b\n        self.proj_w = nn.Linear(self.hidden_dim, self.attention_dim, bias=True)\n        # v.T\n        self.proj_v = nn.Linear(self.attention_dim, 1, bias=False)\n \n    def forward(self, x):\n        \"\"\"\n        :param x: seq_len, batch_size, hidden_dim\n        :return: batch_size * seq_len, batch_size * hidden_dim\n        \"\"\"\n        # print(f\"x shape:{x.shape}\")\n        batch_size, seq_len, _ = x.size()\n        # flat_inputs = x.reshape(-1, self.hidden_dim) # (batch_size*seq_len, hidden_dim)\n        # print(f\"flat_inputs shape:{flat_inputs.shape}\")\n        \n        H = torch.tanh(self.proj_w(x)) # (batch_size, seq_len, hidden_dim)\n        # print(f\"H shape:{H.shape}\")\n        \n        att_scores = torch.softmax(self.proj_v(H),axis=1) # (batch_size, seq_len)\n        # print(f\"att_scores shape:{att_scores.shape}\")\n        \n        attn_x = (x * att_scores).sum(1) # (batch_size, hidden_dim)\n        # print(f\"attn_x shape:{attn_x.shape}\")\n        return attn_x\n\nclass TimmModel2(nn.Module):\n    def __init__(self, backbone, pretrained=False):\n        super().__init__()\n        self.encoder = timm4smp.create_model(\n            backbone,\n            in_chans=CFG.in_chans,\n            num_classes=1,\n            features_only=False,\n            drop_rate=CFG.drop_rate,\n            drop_path_rate=CFG.drop_path_rate,\n            pretrained=pretrained\n        )\n        \n        hdim = self.encoder.conv_head.out_channels\n        self.encoder.classifier = nn.Identity()\n        \n        self.spatialdropout = SpatialDropout(0)\n        self.gru = nn.GRU(hdim, 256, 2, batch_first=True, bidirectional=True)\n        self.mlp_attention_layer = MLPAttentionNetwork(512)\n        self.head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.BatchNorm1d(256),\n            nn.Dropout(0),\n            nn.LeakyReLU(0.1),\n            nn.Linear(256, 1),\n        )\n\n    def forward(self, x):  # (bs, nslice, ch, sz, sz)\n        bs = x.shape[0]\n        x = x.view(bs * CFG.n_slice_per_c, CFG.in_chans, CFG.image_size_cls, CFG.image_size_cls)\n        feat = self.encoder(x)\n        feat = self.spatialdropout(feat)\n        feat = feat.view(bs, CFG.n_slice_per_c, -1)\n        feat, _ = self.gru(feat)\n        feat = self.mlp_attention_layer(feat) # [bs, 512]\n        # feat = feat.contiguous().view(bs * CFG.n_slice_per_c, -1)\n        feat = self.head(feat)\n        # feat = feat.view(bs, CFG.n_slice_per_c).contiguous()\n        # feat = self.mean_layer(feat)\n        feat = feat.view(bs, -1)\n        return feat","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:59:10.772927Z","iopub.execute_input":"2023-10-15T13:59:10.773807Z","iopub.status.idle":"2023-10-15T13:59:10.806192Z","shell.execute_reply.started":"2023-10-15T13:59:10.773772Z","shell.execute_reply":"2023-10-15T13:59:10.80514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# load cls model\nmodel_cls_path = '../input/stage2-rsna'\ncls_models = []\nfor idx_fold in range(5):\n    cls_model = TimmModel2(CFG.backbone_cls, pretrained=False)\n    load_cls_model_file = os.path.join(model_cls_path, f'epoch7_stage2_tf_efficientnetv2_s_in21ft1k_fold{idx_fold}_best.pth')\n    sd = torch.load(load_cls_model_file, map_location='cpu')\n    if 'model_state_dict' in sd.keys():\n        sd = sd['model_state_dict']\n    sd = {k[7:] if k.startswith('module.') else k: sd[k] for k in sd.keys()}\n    cls_model.load_state_dict(sd, strict=True)\n    cls_model = cls_model.to(device)\n    cls_model.eval()\n    cls_models.append(cls_model)\nprint(f'load cls_model succesfull for {len(cls_models)} models!')","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:59:10.807798Z","iopub.execute_input":"2023-10-15T13:59:10.808512Z","iopub.status.idle":"2023-10-15T13:59:15.980381Z","shell.execute_reply.started":"2023-10-15T13:59:10.808481Z","shell.execute_reply":"2023-10-15T13:59:15.975574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"output = []\nwith torch.no_grad():\n    for batch_id, images in enumerate(loader_seg):\n        print('start with index {}'.format(batch_id))\n        images = images.cuda()\n        # seg\n        pmask = seg_model(images).sigmoid()\n        pmask = pmask.squeeze(0)\n   \n        mask = pmask.detach().cpu().numpy()\n        paths = df.loc[batch_id, 'dicom_folder']\n        print('paths', paths)\n        study_id = paths.split('/')[-1]\n        patient_id = paths.split('/')[-2]\n        \n        t_paths = sorted(glob(os.path.join(paths, \"*\")), key=lambda x: int(x.split('/')[-1].split(\".\")[0]))\n        n_scans = len(t_paths)\n        print(f'path {paths} with total {n_scans} scan')\n        cls_inp = []\n        cropped_images = [None] * 5\n        \n        for cid in range(5):\n            bone = []\n            try:\n                msk_b = mask[cid] > 0.2\n                msk_c = mask[cid] > 0.05\n                x = np.where(msk_b.sum(1).sum(1) > 0)[0]\n                y = np.where(msk_b.sum(0).sum(1) > 0)[0]\n                z = np.where(msk_b.sum(0).sum(0) > 0)[0]\n\n                if len(x) == 0 or len(y) == 0 or len(z) == 0:\n                    x = np.where(msk_c.sum(1).sum(1) > 0)[0]\n                    y = np.where(msk_c.sum(0).sum(1) > 0)[0]\n                    z = np.where(msk_c.sum(0).sum(0) > 0)[0]\n\n                x1, x2 = max(0, x[0] - 1), min(mask.shape[1], x[-1] + 1)\n                y1, y2 = max(0, y[0] - 1), min(mask.shape[2], y[-1] + 1)\n                z1, z2 = max(0, z[0] - 1), min(mask.shape[3], z[-1] + 1)\n                zz1, zz2 = int(z1 / CFG.msk_size * n_scans), int(z2 / CFG.msk_size * n_scans)\n                inds = np.linspace(zz1 ,zz2-1, CFG.n_slice_per_c).astype(int)\n                inds_ = np.linspace(z1 ,z2-1, CFG.n_slice_per_c).astype(int)\n                for sid, (ind, ind_) in enumerate(zip(inds, inds_)):\n                    msk_this = mask[cid, :, :, ind_]\n                    images = []\n                    for i in range(-CFG.n_ch // 2 + 1, CFG.n_ch // 2 + 1):\n                        try:\n                            dicom = pydicom.read_file(t_paths[ind+i])\n                            images.append(dicom.pixel_array)\n                        except:\n                            images.append(np.zeros((512, 512)))\n                    data = np.stack(images, -1)\n                    data = data - np.min(data)\n                    data = data / (np.max(data) + 1e-4)\n                    data = (data * 255).astype(np.uint8)\n\n                    msk_this = msk_this[x1:x2, y1:y2]\n                    xx1 = int(x1 / CFG.msk_size * data.shape[0])\n                    xx2 = int(x2 / CFG.msk_size * data.shape[0])\n                    yy1 = int(y1 / CFG.msk_size * data.shape[1])\n                    yy2 = int(y2 / CFG.msk_size * data.shape[1])\n\n                    data = data[xx1:xx2, yy1:yy2]\n                    data = np.stack([cv2.resize(data[:, :, i], (CFG.image_size_cls, CFG.image_size_cls), interpolation = cv2.INTER_LINEAR) for i in range(CFG.n_ch)], -1)\n                    msk_this = (msk_this * 255).astype(np.uint8)\n                    msk_this = cv2.resize(msk_this, (CFG.image_size_cls, CFG.image_size_cls), interpolation = cv2.INTER_LINEAR)\n\n                    data = np.concatenate([data, msk_this[:, :, np.newaxis]], -1)\n                    bone.append(torch.tensor(data))\n            except:\n                bone = []\n                for sid in range(CFG.n_slice_per_c):\n                    bone.append(torch.ones((CFG.image_size_cls, CFG.image_size_cls, CFG.n_ch + 1)).int())\n            cropped_images[cid] = torch.stack(bone, 0)\n\n        cropped_images = torch.cat(cropped_images, 0)\n        print('crop', cropped_images.size())\n        cls_inp = cropped_images.permute(0, 3, 1, 2).float() / 255.\n        cls_inp = cls_inp.to(CFG.device)\n        # cls_inp = cls_inp.unsqueeze(0) # (1, 15*5, 6, 224, 224)\n    \n        pred_cls = []\n        cls_inp = cls_inp.view(5, 15, 6, CFG.image_size_cls, CFG.image_size_cls).contiguous()\n        pred_cls = []\n        for _, model in enumerate(cls_models):\n            logits = model(cls_inp)\n            logits = logits.sigmoid().view(-1, 5)\n            pred_cls.append(logits)\n        # logits = cls_model(cls_inp)\n        # print('logits', logits.size())\n        # logits = logits.sigmoid().view(-1, 5, CFG.n_slice_per_c)\n        pred_cls = torch.stack(pred_cls, 0).mean(0)\n        output.append(pred_cls.cpu())","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:59:15.981936Z","iopub.execute_input":"2023-10-15T13:59:15.98227Z","iopub.status.idle":"2023-10-15T13:59:48.761624Z","shell.execute_reply.started":"2023-10-15T13:59:15.982239Z","shell.execute_reply":"2023-10-15T13:59:48.760524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\nwith open('/kaggle/input/stage2-rsna/train_dict.json', 'r') as fr:\n    aohus = json.load(fr)\n    \n    \ndef norm_to_one(x):\n    s = sum(x)\n    x=[xx/s for xx in x]\n    return x\n\n\ndef post_process(preds, df):\n    all = []\n    # weight mean extra\n    train = pd.read_csv(f'{CFG.data_dir}/train.csv')\n    features_bowel = ['bowel_healthy', 'bowel_injury']\n    features = ['extravasation_healthy', 'extravasation_injury']\n    \n    extra = train[features].mean().tolist()\n    extra = norm_to_one([extra[0], extra[1] * 6])\n    extra = [0.625, 0.375]\n    # bowel = train[features_bowel].mean().tolist()\n    # bowel = norm_to_one([bowel[0], bowel[1] * 6])\n    print('extravasation value', extra)\n    # print('bowel value', bowel)\n    # extra[1] = extra[1] * 6\n    # extra[2] = extra[2] * 28\n    for idx, pred in enumerate(preds):\n#         if str(aohu) in data_aohu.keys():\n#             aohu_weight = data_aohu[str(aohu)]\n#         else:\n#             aohu_weight = 1\n        liver_pred, spleen_pred, kidney_left_pred, kidney_right_pred, bowel_pred = pred\n        liver_healthy, liver_low, liver_high = (1 - liver_pred), liver_pred * (2/3), liver_pred * (1/3)\n#         if liver_healthy < 0.2:\n#             liver_healthy, liver_low, liver_high = liver_healthy * 3, liver_low / 3, liver_high / 3\n        \n        spleen_healthy, spleen_low, spleen_high = (1 - spleen_pred), spleen_pred * (1/2), spleen_pred * (1/2)\n#         if spleen_healthy < 0.2:\n#             spleen_healthy, spleen_low, spleen_high = spleen_healthy * 3, spleen_low / 3, spleen_high / 3\n        \n        kidney_pred = (kidney_left_pred + kidney_right_pred) / 2\n        kidney_healthy, kidney_low, kidney_high = (1 - kidney_pred), kidney_pred * (1/3), kidney_pred * (2/3)\n#         if kidney_healthy < 0.2:\n#             kidney_healthy, kidney_low, kidney_high = kidney_healthy * 3, kidney_low / 3, kidney_high / 3\n        \n        bowel_healthy, bowel_injury = (1 - bowel_pred), bowel_pred\n        # bowel_healthy, bowel_injury = bowel[0], bowel[1]\n        all.append([df.loc[idx, 'patient_id'], bowel_healthy, bowel_injury, extra[0], extra[1], kidney_healthy,\n                    kidney_low, kidney_high, liver_healthy, liver_low, liver_high,\n                    spleen_healthy, spleen_low, spleen_high])\n    return all","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:59:48.763233Z","iopub.execute_input":"2023-10-15T13:59:48.763776Z","iopub.status.idle":"2023-10-15T13:59:48.782632Z","shell.execute_reply.started":"2023-10-15T13:59:48.763741Z","shell.execute_reply":"2023-10-15T13:59:48.781751Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mapping = {0: 'liver', 1: 'spleen', 2: 'kidney', 3: 'kidney', 4: 'bowel'}\noutputs = torch.cat(output)\npreds = outputs\n# preds = (outputs.mean(-1)).clamp(0.0001, 0.9999)\n# preds = preds.detach().cpu().numpy()\npreds = post_process(preds, df)","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:59:48.784154Z","iopub.execute_input":"2023-10-15T13:59:48.785105Z","iopub.status.idle":"2023-10-15T13:59:48.814404Z","shell.execute_reply.started":"2023-10-15T13:59:48.785074Z","shell.execute_reply":"2023-10-15T13:59:48.813476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.DataFrame(preds, columns=['patient_id', 'bowel_healthy', 'bowel_injury', 'extravasation_healthy', 'extravasation_injury',\n                                      'kidney_healthy', 'kidney_low', 'kidney_high', 'liver_healthy', 'liver_low', 'liver_high',\n                                      'spleen_healthy', 'spleen_low', 'spleen_high'])\nsub_df = sub_df.sort_values(by='patient_id')\nsub_df = sub_df.groupby('patient_id').mean().reset_index()\nsub_df.head(5)\n\n","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:59:48.815646Z","iopub.execute_input":"2023-10-15T13:59:48.816161Z","iopub.status.idle":"2023-10-15T13:59:48.857592Z","shell.execute_reply.started":"2023-10-15T13:59:48.816132Z","shell.execute_reply":"2023-10-15T13:59:48.856519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_df = pd.read_csv('/kaggle/input/rsna-2023-abdominal-trauma-detection/sample_submission.csv')\nfinal_df = final_df[['patient_id']]\nfinal_df = final_df.merge(sub_df, on='patient_id', how='left')\nfinal_df.to_csv('submission.csv', index=False)\nfinal_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-10-15T13:59:48.858926Z","iopub.execute_input":"2023-10-15T13:59:48.859808Z","iopub.status.idle":"2023-10-15T13:59:48.890389Z","shell.execute_reply.started":"2023-10-15T13:59:48.859769Z","shell.execute_reply":"2023-10-15T13:59:48.889445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}