{"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":"!pip install /kaggle/input/rsna2023atd-9th/munch-4.0.0-py2.py3-none-any.whl\n!pip install /kaggle/input/rsna2023atd-9th/efficientnet_pytorch-0.7.1-py3-none-any.whl\n!pip install /kaggle/input/rsna2023atd-9th/pretrainedmodels-0.7.4-py3-none-any.whl\n!pip install /kaggle/input/rsna2023atd-9th/timm-0.9.2-py3-none-any.whl\n!pip install /kaggle/input/rsna2023atd-9th/segmentation_models_pytorch-0.3.3-py3-none-any.whl\n!pip install /kaggle/input/rsna2023atd-9th/dicomsdl-0.109.2-cp310-cp310-manylinux_2_12_x86_64.manylinux2010_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2023-10-18T14:26:46.346985Z","iopub.execute_input":"2023-10-18T14:26:46.347317Z","iopub.status.idle":"2023-10-18T14:30:08.96019Z","shell.execute_reply.started":"2023-10-18T14:26:46.347293Z","shell.execute_reply":"2023-10-18T14:30:08.95905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nimport os\nimport sys\nimport time\nfrom pathlib import Path\nimport gc\nfrom typing import List, Optional, Tuple\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport dicomsdl as dicoml\nimport pytorch_lightning as pl\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom segmentation_models_pytorch.decoders.unet.decoder import UnetDecoder\nfrom timm.layers.conv2d_same import Conv2dSame\nfrom tqdm import tqdm\nimport gc\nfrom contextlib import contextmanager\nimport time","metadata":{"execution":{"iopub.status.busy":"2023-10-18T14:32:11.392931Z","iopub.execute_input":"2023-10-18T14:32:11.393816Z","iopub.status.idle":"2023-10-18T14:32:25.367497Z","shell.execute_reply.started":"2023-10-18T14:32:11.39378Z","shell.execute_reply":"2023-10-18T14:32:25.366619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Parameters","metadata":{}},{"cell_type":"code","source":"cols = ['liver_healthy', 'liver_low', 'liver_high', 'spleen_healthy', 'spleen_low', 'spleen_high', 'kidney_healthy', 'kidney_low',\n       'kidney_high', 'bowel_healthy', 'bowel_injury', 'extravasation_healthy', 'extravasation_injury']\nmean_values = [0.897998, 0.082301, 0.019701, 0.887512, 0.063235, 0.049253, 0.942167, 0.036543, 0.02129, 0.979663, 0.020337, 0.936447, 0.063553]\n\nseg_ckpt = \"/kaggle/input/rsna2023atd-9th/seg.ckpt\"\nthree_ckpts = [f\"/kaggle/input/rsna2023atd-9th/514_fold{fold}.ckpt\" for fold in range(4)]\nbowel_ckpts = [f\"/kaggle/input/rsna2023atd-9th/1102_fold{fold}.ckpt\" for fold in range(4)]\nbowel_ckpts2 = [f\"/kaggle/input/rsna2023atd-9th/606v3_fold{fold}.ckpt\" for fold in range(4)]\nev_high_ckpts = [f\"/kaggle/input/rsna2023atd-9th/609_fold{fold}.ckpt\" for fold in range(4)]\n\nimage_root = f\"/kaggle/input/rsna-2023-abdominal-trauma-detection/test_images\"\nseries_meta_path = f\"/kaggle/input/rsna-2023-abdominal-trauma-detection/test_series_meta.csv\"\ndicom_meta_path = f\"/kaggle/input/rsna-2023-abdominal-trauma-detection/test_dicom_tags.parquet\"\n\nsample_submission_df = pd.read_csv(\"/kaggle/input/rsna-2023-abdominal-trauma-detection/sample_submission.csv\")\n\n# segmentation params\nn_blocks = 4\ninit_lr = 3e-3\nbatch_size = 4\ndrop_rate = 0.0\ndrop_path_rate = 0.0\nloss_weights = [1, 1]\np_mixup = 0.1\nout_dim = 4\nimage_sizes_seg = [128, 128, 128]  # (h, w, d)\nbackbone_seg = \"resnet18d\"\n\n# three params\nD = 96\nC = 3\ndrop_rate_last = 0\nbackbone_three = \"seresnext26d_32x4d\"\nimage_size_three = 256\n\nD_BOWEL = 96\nC_BOWEL = 4\nC_BOWEL2 = 3\ndrop_rate_last_bowel = 0\nbackbone_bowel = \"seresnext26d_32x4d\"\nimage_size_bowel = 384\nimage_size_bowel2 = 256\n\n# ev params\nD_EV = 96\nC_EV = 5\ndrop_rate_last_ev = 0\nbackbone_ev = \"seresnext26d_32x4d\"\nimage_size_ev = 512\nuse_high_hu = True","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"@contextmanager\ndef timer(name):\n    t0 = time.time()\n    yield\n    print(f'[{name}] done in {time.time() - t0:.0f} s')\n\n    \ndef standardize_pixel_array(dcm) -> np.ndarray:\n    pixel_array = dcm.pixelData(storedvalue = True)\n    if dcm.PixelRepresentation == 1:\n        bit_shift = dcm.BitsAllocated - dcm.BitsStored\n        dtype = pixel_array.dtype\n        pixel_array = (pixel_array << bit_shift).astype(dtype) >> bit_shift\n\n    intercept = float(dcm.RescaleIntercept)\n    slope = float(dcm.RescaleSlope)\n    center = int(dcm.WindowCenter)\n    width = int(dcm.WindowWidth)\n    low = center - width / 2\n    high = center + width / 2\n\n    pixel_array = (pixel_array * slope) + intercept\n    pixel_array = np.clip(pixel_array, low, high)\n\n    return pixel_array\n\n\ndef process(dcm):\n    img = standardize_pixel_array(dcm)\n    img = (img - img.min()) / (img.max() - img.min() + 1e-6)\n\n    if dcm.PhotometricInterpretation == \"MONOCHROME1\":\n        img = 1 - img\n\n    return img\n\n\ndef get_images(series_id):\n    file_paths = list(Path(image_root).glob(f\"*/{series_id}/*.dcm\"))\n    indices = list(range(len(file_paths)))\n    if len(file_paths) > 800:\n        try:\n            file_paths = sorted(file_paths, key=lambda x: int(s.stem))\n            indices = np.quantile(indices, 800).round().astype(int)\n        except:\n            pass\n    pre = None\n    imgs = {}\n    for i in indices:\n        file_path = file_paths[i]\n        try:\n            crr = dicoml.open(str(file_path))\n        except:\n            crr = pre\n        pos_z = crr.ImagePositionPatient[-1]\n        img = process(crr)\n        imgs[pos_z] = img\n        pre = crr\n    imgs = {k: v for k, v in sorted(imgs.items())}\n    imgs = np.array(list(imgs.values()))\n    return imgs\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3D Segmentation","metadata":{}},{"cell_type":"code","source":"def get_padding(kernel_size: int, stride: int = 1, dilation: int = 1, **_) -> int:\n    padding = ((stride - 1) + dilation * (kernel_size - 1)) // 2\n    return padding\n\n\ndef get_same_padding(x: int, k: int, s: int, d: int):\n    return max((math.ceil(x / s) - 1) * s + (k - 1) * d + 1 - x, 0)\n\n\ndef is_static_pad(kernel_size: int, stride: int = 1, dilation: int = 1, **_):\n    return stride == 1 and (dilation * (kernel_size - 1)) % 2 == 0\n\n\ndef pad_same(x, k: List[int], s: List[int], d: List[int] = (1, 1, 1), value: float = 0):\n    ih, iw, iz = x.size()[-3:]\n    pad_h = get_same_padding(ih, k[0], s[0], d[0])\n    pad_w = get_same_padding(iw, k[1], s[1], d[1])\n    pad_z = get_same_padding(iz, k[2], s[2], d[2])\n    if pad_h > 0 or pad_w > 0 or pad_z > 0:\n        x = F.pad(\n            x,\n            [\n                pad_w // 2,\n                pad_w - pad_w // 2,\n                pad_h // 2,\n                pad_h - pad_h // 2,\n                pad_z // 2,\n                pad_z - pad_z // 2,\n            ],\n            value=value,\n        )\n    return x\n\n\ndef get_padding_value(padding, kernel_size, **kwargs) -> Tuple[Tuple, bool]:\n    dynamic = False\n    if isinstance(padding, str):\n        padding = padding.lower()\n        if padding == \"same\":\n            if is_static_pad(kernel_size, **kwargs):\n                padding = get_padding(kernel_size, **kwargs)\n            else:\n                padding = 0\n                dynamic = True\n        elif padding == \"valid\":\n            padding = 0\n        else:\n            padding = get_padding(kernel_size, **kwargs)\n    return padding, dynamic\n\n\ndef conv3d_same(\n    x,\n    weight: torch.Tensor,\n    bias: Optional[torch.Tensor] = None,\n    stride: Tuple[int, int, int] = (1, 1, 1),\n    padding: Tuple[int, int, int] = (0, 0, 0),\n    dilation: Tuple[int, int, int] = (1, 1, 1),\n    groups: int = 1,\n):\n    x = pad_same(x, weight.shape[-3:], stride, dilation)\n    return F.conv3d(x, weight, bias, stride, (0, 0, 0), dilation, groups)\n\n\nclass Conv3dSame(nn.Conv3d):\n    def __init__(\n        self,\n        in_channels,\n        out_channels,\n        kernel_size,\n        stride=1,\n        padding=0,\n        dilation=1,\n        groups=1,\n        bias=True,\n    ):\n        super(Conv3dSame, self).__init__(in_channels, out_channels, kernel_size, stride, 0, dilation, groups, bias)\n\n    def forward(self, x):\n        return conv3d_same(\n            x,\n            self.weight,\n            self.bias,\n            self.stride,\n            self.padding,\n            self.dilation,\n            self.groups,\n        )\n\n\ndef create_conv3d_pad(in_chs, out_chs, kernel_size, **kwargs):\n    padding = kwargs.pop(\"padding\", \"\")\n    kwargs.setdefault(\"bias\", False)\n    padding, is_dynamic = get_padding_value(padding, kernel_size, **kwargs)\n    if is_dynamic:\n        return Conv3dSame(in_chs, out_chs, kernel_size, **kwargs)\n    else:\n        return nn.Conv3d(in_chs, out_chs, kernel_size, padding=padding, **kwargs)\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(\n            module.weight.unsqueeze(-1).repeat(1, 1, 1, 1, module.kernel_size[0])\n        )\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(\n            module.weight.unsqueeze(-1).repeat(1, 1, 1, 1, module.kernel_size[0])\n        )\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(name, convert_3d(child))\n    del module\n\n    return module_output\n\n\nclass TimmSegModel(nn.Module):\n    def __init__(self, backbone, segtype=\"unet\", pretrained=False):\n        super(TimmSegModel, self).__init__()\n\n        self.encoder = timm.create_model(\n            backbone,\n            in_chans=3,\n            features_only=True,\n            drop_rate=drop_rate,\n            drop_path_rate=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        if segtype == \"unet\":\n            self.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(\n            decoder_channels[n_blocks - 1],\n            out_dim,\n            kernel_size=(3, 3),\n            stride=(1, 1),\n            padding=(1, 1),\n        )\n\n    def forward(self, x):\n        global_features = [0] + self.encoder(x)[:n_blocks]\n        seg_features = self.decoder(*global_features)\n        seg_features = self.segmentation_head(seg_features)\n        return seg_features\n\n\ndef prepare_seg_model(ckpt):\n    seg_model = TimmSegModel(backbone_seg, pretrained=False)\n    seg_model = convert_3d(seg_model)\n    seg_model = seg_model.cuda()\n    ckpt = torch.load(ckpt)\n    seg_model.load_state_dict(ckpt)\n    seg_model.eval()\n    return seg_model\n\n\ndef get_seg_input(images):\n    indices = np.quantile(list(range(images.shape[0])), np.linspace(0.0, 1.0, image_sizes_seg[2])).round().astype(int)\n    res = []\n    for index in indices:\n        res += [\n            cv2.resize(\n                images[index],\n                (image_sizes_seg[0], image_sizes_seg[1]),\n                interpolation=cv2.INTER_LINEAR,\n            )\n        ]\n    res = np.stack(res, -1)\n    if len(res.shape) < 4:\n        res = np.expand_dims(res, 0).repeat(3, 0)  # to 3ch\n    return res, indices  # (c,h,w,d)\n\n\ndef postprocess_seg_pred(seg_pred):\n    seg_pred = (seg_pred > 0.5).astype(np.uint8)\n    for cid in range(seg_pred.shape[0]):\n        seg_pred[cid] *= cid + 1\n    seg_pred = seg_pred.max(axis=0) # c,h,w,d\n    return seg_pred # h,w,d\n\ndef get_pad_voxel(seg_pred, threshold=10):\n    bbox = []\n    zz = []\n    for i in range(1, 5):\n        seg_bin_pred = seg_pred == i\n\n        try:\n            z = seg_bin_pred.sum(0).sum(0) > threshold\n            zmin = np.argmax(z)\n            zmax = len(z) - np.argmax(z[::-1])\n        except:\n            zmin = 0\n            zmax = 128\n\n        try:\n            x = seg_bin_pred.sum(0).sum(1) > threshold\n            xmin = np.argmax(x)\n            xmax = len(x) - np.argmax(x[::-1])\n        except:\n            xmin = 0\n            xmax = 128\n\n        try:\n            y = seg_bin_pred.sum(1).sum(1) > threshold\n            ymin = np.argmax(y)\n            ymax = len(y) - np.argmax(y[::-1])\n        except:\n            ymin = 0\n            ymax = 128\n\n        bbox += [[xmin, ymin, xmax, ymax]]\n        zz += [[zmin, zmax]]\n\n    bbox = np.stack(bbox).astype(np.float32)\n    bbox /= 128.0\n\n    bbox[:, :2] = (bbox[:, :2] - 0.01).clip(0.0, 1.0)\n    bbox[:, 2:] = (bbox[:, 2:] + 0.01).clip(0.0, 1.0)\n\n    zz = np.stack(zz)\n    return bbox, zz","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Liver/spleen/kidney Classification","metadata":{}},{"cell_type":"code","source":"class Attention(nn.Module):\n    def __init__(self, feature_dim, step_dim, bias=True, **kwargs):\n        super(Attention, self).__init__(**kwargs)\n\n        self.supports_masking = True\n\n        self.bias = bias\n        self.feature_dim = feature_dim\n        self.step_dim = step_dim\n        self.features_dim = 0\n\n        weight = torch.zeros(feature_dim, 1)\n        nn.init.xavier_uniform_(weight)\n        self.weight = nn.Parameter(weight)\n\n        if bias:\n            self.b = nn.Parameter(torch.zeros(step_dim))\n\n    def forward(self, x, mask=None):\n        feature_dim = self.feature_dim\n        step_dim = self.step_dim\n\n        eij = torch.mm(x.contiguous().view(-1, feature_dim), self.weight).view(-1, step_dim)\n\n        if self.bias:\n            eij = eij + self.b\n\n        eij = torch.tanh(eij)\n        a = torch.exp(eij)\n\n        if mask is not None:\n            a = a * mask\n\n        a = a / torch.sum(a, 1, keepdim=True) + 1e-10\n\n        weighted_input = x * torch.unsqueeze(a, -1)\n        return torch.sum(weighted_input, 1)\n\n\nclass ThreeModel(pl.LightningModule):\n    def __init__(self):\n        out_dim = 256\n        super(ThreeModel, self).__init__()\n\n        self.backbone = timm.create_model(backbone_three, pretrained=False, in_chans=C + 1)\n        if \"resnet\" in backbone_three or \"seresnext\" in backbone_three:\n            hdim = self.backbone.fc.in_features\n            self.backbone.fc = nn.Identity()\n        if \"convnext\" in backbone_three:\n            hdim = self.backbone.head.fc.in_features\n            self.backbone.head.fc = nn.Identity()\n\n        self.lstm = nn.LSTM(hdim, out_dim, num_layers=2, bidirectional=True, batch_first=True)\n        self.relu = nn.ReLU()\n\n        self.conv1d_kidney = nn.Conv1d(D, 1, 1)\n        self.attn_kidney = Attention(feature_dim=out_dim * 2, step_dim=D)\n        self.attn_bn_kidney = nn.BatchNorm1d(out_dim * 2)\n        self.head_kidney = nn.Sequential(\n            nn.Linear(out_dim * 2 * 2, out_dim),\n            nn.BatchNorm1d(out_dim),\n            nn.Dropout(drop_rate_last),\n            nn.LeakyReLU(0.1),\n            nn.Linear(out_dim, 3),\n        )\n        self.conv1d_liver = nn.Conv1d(D, 1, 1)\n        self.attn_liver = Attention(feature_dim=out_dim * 2, step_dim=D)\n        self.attn_bn_liver = nn.BatchNorm1d(out_dim * 2)\n        self.head_liver = nn.Sequential(\n            nn.Linear(out_dim * 2 * 2, out_dim),\n            nn.BatchNorm1d(out_dim),\n            nn.Dropout(drop_rate_last),\n            nn.LeakyReLU(0.1),\n            nn.Linear(out_dim, 3),\n        )\n        self.conv1d_spleen = nn.Conv1d(D, 1, 1)\n        self.attn_spleen = Attention(feature_dim=out_dim * 2, step_dim=D)\n        self.attn_bn_spleen = nn.BatchNorm1d(out_dim * 2)\n        self.head_spleen = nn.Sequential(\n            nn.Linear(out_dim * 2 * 2, out_dim),\n            nn.BatchNorm1d(out_dim),\n            nn.Dropout(drop_rate_last),\n            nn.LeakyReLU(0.1),\n            nn.Linear(out_dim, 3),\n        )\n\n    def forward(self, x):\n        b, d, c, h, w = x.shape\n        x = x.contiguous().view(b * d, c, h, w)\n        x = self.backbone(x)\n        x = x.view(b, d, -1)  # b,d,c\n        # lstm\n        x, _ = self.lstm(x)\n\n        # 1.liver\n        x_conv_liver = self.conv1d_liver(x)[:, 0]\n        x_attn_liver = self.attn_liver(x)\n        x_attn_liver = self.attn_bn_liver(x_attn_liver)\n        x_attn_liver = self.relu(x_attn_liver)\n        x_liver = torch.cat([x_conv_liver, x_attn_liver], dim=-1)\n        logit_liver = self.head_liver(x_liver)\n\n        # 2.spleen\n        x_conv_spleen = self.conv1d_spleen(x)[:, 0]\n        x_attn_spleen = self.attn_spleen(x)\n        x_attn_spleen = self.attn_bn_spleen(x_attn_spleen)\n        x_attn_spleen = self.relu(x_attn_spleen)\n        x_spleen = torch.cat([x_conv_spleen, x_attn_spleen], dim=-1)\n        logit_spleen = self.head_spleen(x_spleen)\n\n        # 3.kidney\n        x_conv_kidney = self.conv1d_kidney(x)[:, 0]\n        x_attn_kidney = self.attn_kidney(x)\n        x_attn_kidney = self.attn_bn_kidney(x_attn_kidney)\n        x_attn_kidney = self.relu(x_attn_kidney)\n        x_kidney = torch.cat([x_conv_kidney, x_attn_kidney], dim=-1)\n        logit_kidney = self.head_kidney(x_kidney)\n        return logit_liver, logit_spleen, logit_kidney  # b,n_organ,n_target\n\n\ndef prepare_three_model(ckpt):\n    three_model = ThreeModel.load_from_checkpoint(ckpt)\n    three_model = three_model.cuda()\n    three_model.eval()\n    return three_model\n\n\ndef get_three_crop_voxel_coord(bboxes, zz):\n    (xmin, ymin) = bboxes[:3, :2].min(axis=0)\n    (xmax, ymax) = bboxes[:3, 2:].max(axis=0)\n    (zmin, zmax) = zz[:3, 0].min(), zz[:3, 1].max()\n\n    return xmin, ymin, zmin, xmax, ymax, zmax\n\n\ndef pad_resize(image, h, w):\n    height, width = image.shape[:2]\n    aspect_ratio = width / height\n\n    if aspect_ratio > w / h:\n        new_width = w\n        new_height = int(new_width / aspect_ratio)\n    else:\n        new_height = h\n        new_width = int(new_height * aspect_ratio)\n\n    resized_image = cv2.resize(image, (new_width, new_height), interpolation=cv2.INTER_LINEAR)\n\n    top_pad = (h - new_height) // 2\n    bottom_pad = h - new_height - top_pad\n    left_pad = (w - new_width) // 2\n    right_pad = w - new_width - left_pad\n\n    padded_image = cv2.copyMakeBorder(\n        resized_image,\n        top_pad,\n        bottom_pad,\n        left_pad,\n        right_pad,\n        cv2.BORDER_CONSTANT,\n        value=0,\n    )\n    return padded_image\n\n\ndef get_three_input(images, seg_pred, seg_slice_indices, bbox, zmin, zmax):\n    # ======================\n    # segmentation mask\n    # ======================\n    spred = seg_pred * ((1 <= seg_pred) & (seg_pred <= 3)).astype(np.uint8)\n    # crop\n    xmin, ymin, xmax, ymax = (bbox * spred.shape[0]).astype(np.uint16)\n    spred = spred[ymin:ymax, xmin:xmax, zmin:zmax]\n    # resize\n    spred = pad_resize(spred, image_size_three, image_size_three)\n    spred = spred.transpose(2, 0, 1)  # d, h, w\n    spred = cv2.resize(spred, (image_size_three, D), interpolation=cv2.INTER_LINEAR)  # resize\n    spred = spred.transpose(1, 2, 0)\n    assert spred.shape == (image_size_three, image_size_three, D)\n\n    # ======================\n    # images\n    # ======================\n    use_z = (\n        np.quantile(\n            list(range(seg_slice_indices[zmin], seg_slice_indices[zmax - 1] + 1)),\n            np.linspace(0, 1, D),\n        )\n        .round()\n        .astype(int)\n    )\n    res = []\n    c = C // 2\n    for i, zi in enumerate(use_z):\n        tmp = []\n        for di in range(-c, c + 1):\n            j = min(max(zi + di, 0), images.shape[0] - 1)\n            img = images[j]\n            xmin, ymin, xmax, ymax = (bbox * img.shape[0]).astype(np.uint16)\n            img = img[ymin:ymax, xmin:xmax]  # crop\n            img = pad_resize(img, image_size_three, image_size_three)\n            tmp += [img]\n        tmp += [spred[:, :, i]]\n        tmp = np.stack(tmp, -1)  # h, w, c\n        res += [tmp]\n    res = np.stack(res)  # d, h, w, c\n\n    res = res.transpose(0, 3, 1, 2).astype(np.float64)\n    assert res.shape == (D, C + 1, image_size_three, image_size_three)\n    return res  # d, c, h, w\n\n\ndef get_three_pred_df(three_pred):\n    res = []\n    organ_names = [\"liver\", \"spleen\", \"kidney\"]\n    levels = [\"healthy\", \"low\", \"high\"]\n    for i in range(3):\n        organ = organ_names[i]\n        preds = torch.softmax(three_pred[i][0], dim=-1).view(1, 3).cpu().numpy()\n        cols = [f\"{organ}_{level}\" for level in levels]\n        pred_df = pd.DataFrame(preds, columns=cols)\n        res += [pred_df]\n    res = pd.concat(res, axis=1)\n    return res","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Bowel Classification","metadata":{}},{"cell_type":"code","source":"class BowelModel(pl.LightningModule):\n    def __init__(self):\n        out_dim = 256\n        super(BowelModel, self).__init__()\n\n        self.backbone = timm.create_model(backbone_bowel, pretrained=False, in_chans=C_BOWEL)\n        if \"resnet\" in backbone_bowel or \"seresnext\" in backbone_bowel:\n            hdim = self.backbone.fc.in_features\n            self.backbone.fc = nn.Identity()\n        if \"convnext\" in backbone_bowel:\n            hdim = self.backbone.head.fc.in_features\n            self.backbone.head.fc = nn.Identity()\n\n        self.lstm = nn.LSTM(hdim, out_dim, num_layers=2, bidirectional=True, batch_first=True)\n        self.relu = nn.ReLU()\n\n        self.conv1d = nn.Conv1d(D_BOWEL, 1, 1)\n        self.attn = Attention(feature_dim=out_dim * 2, step_dim=D)\n        self.attn_bn = nn.BatchNorm1d(out_dim * 2)\n        self.head = nn.Sequential(\n            nn.Linear(out_dim * 2 * 2, out_dim),\n            nn.BatchNorm1d(out_dim),\n            nn.Dropout(drop_rate_last),\n            nn.LeakyReLU(0.1),\n            nn.Linear(out_dim, 1),\n        )\n\n    def forward(self, x):\n        b, d, c, h, w = x.shape\n        x = x.contiguous().view(b * d, c, h, w)\n        x = self.backbone(x)\n        x = x.view(b, d, -1)  # b,d,c\n\n        # lstm\n        x, _ = self.lstm(x)\n\n        x_conv = self.conv1d(x)[:, 0]\n        x_attn = self.attn(x)\n        x_attn = self.attn_bn(x_attn)\n        x_attn = self.relu(x_attn)\n        x = torch.cat([x_conv, x_attn], dim=-1)\n        logit = self.head(x)\n\n        return logit  # b,n_organ,n_target\n\n\ndef prepare_bowel_model(ckpt):\n    bowel_model = BowelModel.load_from_checkpoint(ckpt)\n    bowel_model = bowel_model.cuda()\n    bowel_model.eval()\n    return bowel_model\n\n\ndef get_bowel_crop_voxel_coord(bboxes, zz):\n    (xmin, ymin) = bboxes[3, :2]\n    (xmax, ymax) = bboxes[3, 2:]\n    (zmin, zmax) = zz[3, 0], zz[3, 1]\n    xmin = min(max((xmin - 0.09), 0), 1)\n    ymin = min(max((ymin - 0.09), 0), 1)\n    xmax = min(max((xmax + 0.09), 0), 1)\n    ymax = min(max((ymax + 0.09), 0), 1)\n    return xmin, ymin, zmin, xmax, ymax, zmax\n\ndef resize_d(images):  # d, c, h, w\n    d, c, h, w = images.shape\n    res = [cv2.resize(images[:, ci, :, :], (h, D_BOWEL)) for ci in range(c)]\n    return np.stack(res, axis=1)\n\ndef get_bowel_input(images, seg_slice_indices, xmin, ymin, zmin, xmax, ymax, zmax, seg_pred):\n    d, h, w = images.shape\n    use_z = (\n        np.quantile(\n            range(images.shape[0]),\n            np.linspace(0, 1, 128),\n        )\n        .round()\n        .astype(int)\n    )\n    indices = np.stack([use_z-1, use_z, use_z+1], axis=-1).clip(0, len(images)-1)\n    images = images[indices] # d, c, h, w\n    mask = (seg_pred == 4).astype(float) # h, w, d\n    mask = cv2.resize(mask, (w, h)).transpose(2, 0, 1) # d, h, w\n    images = np.concatenate([images, mask[:, np.newaxis, :, :]], axis=1)\n    \n    # crop\n    xmin = int(round(xmin * w))\n    xmax = int(round(xmax * w))\n    ymin = int(round(ymin * h))\n    ymax = int(round(ymax * h))\n    images = images[zmin:zmax, :, ymin:ymax, xmin:xmax]\n    \n    d, c, h, w = images.shape\n    images = (\n        cv2.resize(images.transpose(2, 3, 0, 1).reshape(h, w, d * c), (384, 384))\n        .reshape(384, 384, d, c)\n        .transpose(2, 3, 0, 1)\n    )\n    images = resize_d(images)\n    return images  # (D, c, h, w)\n\n\ndef get_bowel_pred_df(bowel_pred):\n    pos_p = bowel_pred[0].sigmoid().cpu().numpy()\n    pred_df = pd.DataFrame(\n        np.concatenate([1 - pos_p, pos_p]).reshape(1, 2),\n        columns=[\"bowel_healthy\", \"bowel_injury\"],\n    )\n    return pred_df\n\n\nclass BowelModel2(pl.LightningModule):\n    def __init__(self):\n        out_dim = 256\n        super(BowelModel2, self).__init__()\n\n        self.backbone = timm.create_model(backbone_bowel, pretrained=False, in_chans=C_BOWEL2)\n        if \"resnet\" in backbone_bowel or \"seresnext\" in backbone_bowel:\n            hdim = self.backbone.fc.in_features\n            self.backbone.fc = nn.Identity()\n        if \"convnext\" in backbone_bowel:\n            hdim = self.backbone.head.fc.in_features\n            self.backbone.head.fc = nn.Identity()\n\n        self.lstm = nn.LSTM(hdim, out_dim, num_layers=2, bidirectional=True, batch_first=True)\n        self.relu = nn.ReLU()\n\n        self.conv1d = nn.Conv1d(D_BOWEL, 1, 1)\n        self.attn = Attention(feature_dim=out_dim * 2, step_dim=D)\n        self.attn_bn = nn.BatchNorm1d(out_dim * 2)\n        self.head = nn.Sequential(\n            nn.Linear(out_dim * 2 * 2, out_dim),\n            nn.BatchNorm1d(out_dim),\n            nn.Dropout(drop_rate_last),\n            nn.LeakyReLU(0.1),\n            nn.Linear(out_dim, 1),\n        )\n\n    def forward(self, x):\n        b, d, c, h, w = x.shape\n        x = x.contiguous().view(b * d, c, h, w)\n        x = self.backbone(x)\n        x = x.view(b, d, -1)  # b,d,c\n\n        # lstm\n        x, _ = self.lstm(x)\n\n        x_conv = self.conv1d(x)[:, 0]\n        x_attn = self.attn(x)\n        x_attn = self.attn_bn(x_attn)\n        x_attn = self.relu(x_attn)\n        x = torch.cat([x_conv, x_attn], dim=-1)\n        logit = self.head(x)\n\n        return logit  # b,n_organ,n_target\n\n\ndef prepare_bowel_model2(ckpt):\n    bowel_model = BowelModel2.load_from_checkpoint(ckpt)\n    bowel_model = bowel_model.cuda()\n    bowel_model.eval()\n    return bowel_model\n\n\ndef get_bowel_crop_voxel_coord2(bboxes, zz):\n    (xmin, ymin) = bboxes[3, :2]\n    (xmax, ymax) = bboxes[3, 2:]\n    (zmin, zmax) = zz[3, 0], zz[3, 1]\n    return xmin, ymin, zmin, xmax, ymax, zmax\n\n\ndef get_bowel_input2(images, seg_indices, bbox, zmin, zmax):\n    use_z = (\n        np.quantile(\n            list(range(seg_indices[zmin], seg_indices[zmax - 1] + 1)),\n            np.linspace(0, 1, D_BOWEL),\n        )\n        .round()\n        .astype(int)\n    )\n    indices = np.concatenate([use_z + di for di in range(-(C_BOWEL // 2), (C_BOWEL // 2) + 1)], axis=-1).clip(\n        0, images.shape[0] - 1\n    )\n    xmin, ymin, xmax, ymax = (bbox * images.shape[1]).astype(np.uint16)\n    res = images[indices][:, ymin:ymax, xmin:xmax].transpose(1, 2, 0)\n    res = pad_resize(res, image_size_bowel2, image_size_bowel2)\n    res = res.reshape(image_size_bowel2, image_size_bowel2, C_BOWEL2, D_BOWEL).transpose(3, 2, 0, 1)\n    return res  # (D, c, h, w)\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Extravasation Classification","metadata":{}},{"cell_type":"code","source":"class EVModel(pl.LightningModule):\n    def __init__(self):\n        out_dim = 256\n        super(EVModel, self).__init__()\n\n        self.backbone = timm.create_model(backbone_ev, pretrained=False, in_chans=C_EV)\n        if \"resnet\" in backbone_ev or \"seresnext\" in backbone_ev:\n            hdim = self.backbone.fc.in_features\n            self.backbone.fc = nn.Identity()\n        if \"convnext\" in backbone_ev:\n            hdim = self.backbone.head.fc.in_features\n            self.backbone.head.fc = nn.Identity()\n        self.lstm = nn.LSTM(hdim, out_dim, num_layers=2, bidirectional=True, batch_first=True)\n        self.relu = nn.ReLU()\n\n        self.conv1d = nn.Conv1d(D_EV, 1, 1)\n        self.attn = Attention(feature_dim=out_dim * 2, step_dim=D)\n        self.attn_bn = nn.BatchNorm1d(out_dim * 2)\n        self.head = nn.Sequential(\n            nn.Linear(out_dim * 2 * 2, out_dim),\n            nn.BatchNorm1d(out_dim),\n            nn.Dropout(drop_rate_last),\n            nn.LeakyReLU(0.1),\n            nn.Linear(out_dim, 1),\n        )\n\n    def forward(self, x):\n        b, d, c, h, w = x.shape\n        x = x.contiguous().view(b * d, c, h, w)\n        x = self.backbone(x)\n        x = x.view(b, d, -1)  # b,d,c\n\n        # lstm\n        x, _ = self.lstm(x)\n\n        x_conv = self.conv1d(x)[:, 0]\n        x_attn = self.attn(x)\n        x_attn = self.attn_bn(x_attn)\n        x_attn = self.relu(x_attn)\n        x = torch.cat([x_conv, x_attn], dim=-1)\n        logit = self.head(x)\n\n        return logit  # b,n_organ,n_target\n\n\ndef prepare_ev_model(ckpt):\n    ev_model = EVModel.load_from_checkpoint(ckpt)\n    ev_model = ev_model.cuda()\n    ev_model.eval()\n    return ev_model\n\n\ndef get_ev_input(images):\n    use_z = (\n        np.quantile(\n            range(len(images)),\n            np.linspace(0, 1, D_EV),\n        )\n        .round()\n        .astype(int)\n    )\n    indices = np.concatenate([use_z + di for di in range(-(C_EV // 2), (C_EV // 2) + 1)], axis=-1).clip(\n        0, images.shape[0] - 1\n    )\n    res = images[indices].transpose(1, 2, 0)\n    res = pad_resize(res, image_size_ev, image_size_ev)\n    res = res.reshape(image_size_ev, image_size_ev, C_EV, D_EV).transpose(3, 2, 0, 1)\n    return res  # (D, c, h, w)\n\n\ndef get_ev_pred_df(ev_pred):\n    pos_p = ev_pred[0].sigmoid().cpu().numpy()\n    pred_df = pd.DataFrame(\n        np.concatenate([1 - pos_p, pos_p]).reshape(1, 2),\n        columns=[\"extravasation_healthy\", \"extravasation_injury\"],\n    )\n    return pred_df","metadata":{"execution":{"iopub.status.busy":"2023-10-15T14:28:17.589717Z","iopub.execute_input":"2023-10-15T14:28:17.590342Z","iopub.status.idle":"2023-10-15T14:28:17.941941Z","shell.execute_reply.started":"2023-10-15T14:28:17.59031Z","shell.execute_reply":"2023-10-15T14:28:17.940905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load models","metadata":{}},{"cell_type":"code","source":"seg_model = prepare_seg_model(seg_ckpt)\nprint(\"seg model is loaded.\")\n\nthree_models = [prepare_three_model(three_ckpt) for three_ckpt in three_ckpts]\nprint(f\"{len(three_models)} three models are loaded.\")\n\nbowel_models = [prepare_bowel_model(bowel_ckpt) for bowel_ckpt in bowel_ckpts]\nprint(f\"{len(bowel_models)} bowel models are loaded.\")\n\nbowel_models2 = [prepare_bowel_model2(bowel_ckpt) for bowel_ckpt in bowel_ckpts2]\nprint(f\"{len(bowel_models2)} bowel models2 are loaded.\")\n\nev_models = [prepare_ev_model(ev_high_ckpt) for ev_high_ckpt in ev_high_ckpts]\nprint(f\"{len(ev_models)} ev models are loaded.\")","metadata":{"execution":{"iopub.status.busy":"2023-10-15T14:28:17.965446Z","iopub.execute_input":"2023-10-15T14:28:17.967375Z","iopub.status.idle":"2023-10-15T14:28:47.498273Z","shell.execute_reply.started":"2023-10-15T14:28:17.967351Z","shell.execute_reply":"2023-10-15T14:28:47.497355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Main","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(series_meta_path)\ndf = df.sort_values(\"aortic_hu\").reset_index(drop=True)\ndf_high_hu = df.groupby(\"patient_id\").tail(1).set_index(\"patient_id\")\nprint(df_high_hu.sort_values(\"patient_id\").head(5))\ndf = df.groupby(\"patient_id\").head(1).reset_index(drop=True)\nprint(df.sort_values(\"patient_id\").head(5))\nprint(\"pred iteration:\", len(df))\n\npred_dfs = []\nfor i, row in enumerate(tqdm(df.itertuples(), total=len(df))):\n    with torch.no_grad():\n        seg_input = None\n        three_input = None\n        bowel_input = None\n        ev_input = None\n\n        try:\n            with timer(\"get_images\"):\n                images = get_images(row.series_id)\n            \n            with timer(\"3D segmentation\"):\n                seg_input, seg_slice_indices = get_seg_input(images)\n\n                # preprocess\n                seg_input = torch.tensor(seg_input).float()\n                seg_input = seg_input.view(1, *seg_input.shape)\n\n                seg_pred = seg_model(seg_input.cuda())[0]\n                seg_pred = seg_pred.sigmoid().cpu().numpy()\n\n                seg_pred = postprocess_seg_pred(seg_pred)\n\n                bboxes, zz = get_pad_voxel(seg_pred)\n                \n            # ================\n            # kidney, spleen, liver\n            # ================\n            with timer(\"3 heads\"):\n                try:\n                    (xmin, ymin, zmin, xmax, ymax, zmax) = get_three_crop_voxel_coord(bboxes, zz)\n                    bbox = np.array([xmin, ymin, xmax, ymax])\n                    three_input = get_three_input(images, seg_pred, seg_slice_indices, bbox, zmin, zmax)\n                    # preprocess\n                    three_input = torch.tensor(three_input).float()  # D,c,h,w\n                    three_input = three_input.view(1, *three_input.shape)\n                    # forward\n                    dfs = []\n                    for three_model in three_models:\n                        three_pred = three_model(three_input.cuda())\n                        three_pred_df = get_three_pred_df(three_pred)\n                        dfs += [three_pred_df]\n                    three_pred_df = pd.concat(dfs, axis=0).mean(axis=0).to_frame().transpose()\n                except Exception as e:\n                    print(e)\n                    three_pred_df = pd.DataFrame([mean_values[:9]], columns=cols[:9])\n                    \n            # ================\n            # bowel\n            # ================\n            with timer(\"bowel\"):\n                try:\n                    (xmin, ymin, zmin, xmax, ymax, zmax) = get_bowel_crop_voxel_coord(bboxes, zz)\n                    bowel_input = get_bowel_input(images, seg_slice_indices, xmin, ymin, zmin, xmax, ymax, zmax, seg_pred)\n                    # preprocess\n                    bowel_input = torch.tensor(bowel_input).float()  # D,c,h,w\n                    bowel_input = bowel_input.view(1, *bowel_input.shape)\n                    # forward\n                    dfs = []\n                    for bowel_model in bowel_models:\n                        bowel_pred = bowel_model(bowel_input.cuda())\n                        bowel_pred_df = get_bowel_pred_df(bowel_pred)\n                        dfs += [bowel_pred_df]\n                    bowel_pred_df = pd.concat(dfs, axis=0).mean(axis=0).to_frame().transpose()\n                except Exception as e:\n                    print(e)\n                    bowel_pred_df = pd.DataFrame([mean_values[9:11]], columns=cols[9:11])\n            \n            with timer(\"bowel\"):\n                try:\n                    (xmin, ymin, zmin, xmax, ymax, zmax) = get_bowel_crop_voxel_coord2(bboxes, zz)\n                    bbox = np.array([xmin, ymin, xmax, ymax])\n                    bowel_input = get_bowel_input2(images, seg_slice_indices, bbox, zmin, zmax)\n                    # preprocess\n                    bowel_input = torch.tensor(bowel_input).float()  # D,c,h,w\n                    bowel_input = bowel_input.view(1, *bowel_input.shape)\n                    # forward\n                    dfs = []\n                    for bowel_model in bowel_models2:\n                        bowel_pred = bowel_model(bowel_input.cuda())\n                        a = get_bowel_pred_df(bowel_pred)\n                        dfs += [a]\n                    bowel_pred_df2 = pd.concat(dfs, axis=0).mean(axis=0).to_frame().transpose()\n                    bowel_pred_df = pd.concat([bowel_pred_df, bowel_pred_df2], axis=0).mean(axis=0)\n                except Exception as e:\n                    pass\n                \n            # ================\n            # extravasation\n            # ================\n            with timer(\"extravasation\"):\n                try:\n                    if use_high_hu: # update images\n                        ev_row = df_high_hu.loc[row.patient_id]\n                        high_series_id = int(ev_row.series_id)\n                        if high_series_id != row.series_id:\n                            with timer(\"get_images ev high\"):\n                                images = get_images(high_series_id)\n                    ev_input = get_ev_input(images)\n                    # preprocess\n                    ev_input = torch.tensor(ev_input).float()  # D,c,h,w\n                    ev_input = ev_input.view(1, *ev_input.shape)\n                    # forward\n                    dfs = []\n                    for ev_model in ev_models:\n                        ev_pred = ev_model(ev_input.cuda())\n                        ev_pred_df = get_ev_pred_df(ev_pred)\n                        dfs += [ev_pred_df]\n                    ev_pred_df = pd.concat(dfs, axis=0).mean(axis=0).to_frame().transpose()\n                except Exception as e:\n                    print(e)\n                    ev_pred_df = pd.DataFrame([mean_values[11:]], columns=cols[11:])\n                    \n            # ================\n            # merge\n            # ================\n            pred_dfs += [pd.concat([three_pred_df, bowel_pred_df, ev_pred_df], axis=1)]\n            \n        except Exception as e:\n            print(e)\n            del seg_input, three_input, bowel_input, ev_input\n            gc.collect()\n            pred_dfs += [pd.DataFrame([mean_values], columns=cols)]\n            \n    if i % 50 == 0:\n        torch.cuda.empty_cache()\n            \nsub_df = pd.concat(pred_dfs, axis=0).reset_index(drop=True)\nsub_df[\"patient_id\"] = df[\"patient_id\"]","metadata":{"execution":{"iopub.status.busy":"2023-10-15T14:28:47.49991Z","iopub.execute_input":"2023-10-15T14:28:47.500158Z","iopub.status.idle":"2023-10-15T14:30:41.76114Z","shell.execute_reply.started":"2023-10-15T14:28:47.500137Z","shell.execute_reply":"2023-10-15T14:30:41.759876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Post-Processing","metadata":{}},{"cell_type":"code","source":"low_coef = 1.5\nhigh_coef = 1.5\nev_coef = 1.5\nbowel_coef = 1.0\n\nlow_cols = [col for col in sub_df.columns if \"_low\" in col]\nhigh_cols = [col for col in sub_df.columns if \"_high\" in col]\nev = \"extravasation_injury\"\nbowel = \"bowel_injury\"\n\nfor lc in low_cols:\n    sub_df[lc] = sub_df[lc] * low_coef\n    \nfor hc in high_cols:\n    sub_df[hc] = sub_df[hc] * high_coef\n    \nsub_df[ev] = sub_df[ev] * ev_coef\nsub_df[bowel] = sub_df[bowel] * bowel_coef","metadata":{"execution":{"iopub.status.busy":"2023-10-15T14:30:41.764295Z","iopub.status.idle":"2023-10-15T14:30:41.764947Z","shell.execute_reply.started":"2023-10-15T14:30:41.764725Z","shell.execute_reply":"2023-10-15T14:30:41.764747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"sub_df = sub_df.set_index(\"patient_id\")\nsub_df = sub_df.loc[sample_submission_df[\"patient_id\"]]\nsub_df = sub_df.reset_index(drop=False)\nsub_df = sub_df[sample_submission_df.columns]\nsub_df.to_csv(\"submission.csv\",index=False)","metadata":{},"execution_count":null,"outputs":[]}]}