{"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":"markdown","source":"# 🛠 Install Libraries","metadata":{}},{"cell_type":"code","source":"# !pip install -q ../input/mmlablibsv2/einops-0.4.1-py3-none-any.whl\n# !pip install -q ../input/mmlablibsv2/yapf-0.31.0-py2.py3-none-any.whl\n# !pip install -q ../input/mmlablibsv2/addict-2.4.0-py3-none-any.whl\n# !pip install -q ../input/mmlablibsv2/terminaltables-3.1.0-py3-none-any.whl\n# !pip install -q ../input/mmlablibsv2/mmcv_full-1.3.17-cp37-cp37m-linux_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2022-07-07T17:33:58.89075Z","iopub.execute_input":"2022-07-07T17:33:58.89132Z","iopub.status.idle":"2022-07-07T17:33:58.921083Z","shell.execute_reply.started":"2022-07-07T17:33:58.891206Z","shell.execute_reply":"2022-07-07T17:33:58.920468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip uninstall -qy transformers\n!pip uninstall -qy tokenizers\n!pip install -q ../input/pytorch-segmentation-models-lib/pretrainedmodels-0.7.4/pretrainedmodels-0.7.4\n!pip install -q ../input/pytorch-segmentation-models-lib/efficientnet_pytorch-0.6.3/efficientnet_pytorch-0.6.3\n!pip install -q ../input/segmentation-models-pytorch-030/timm-0.5.4-py3-none-any.whl\n!pip install -q ../input/segmentation-models-pytorch-030/segmentation_models_pytorch-0.3.0.dev0-py3-none-any.whl\n!pip uninstall -qy transformers\n!pip uninstall -qy tokenizers\n!pip install -q ../input/uwmgiseg/dependencies/tokenizers-0.12.1-cp37-cp37m-manylinux_2_12_x86_64.manylinux2010_x86_64.whl\n!pip install ../input/uwmgiseg/dependencies/huggingface_hub-0.5.1-py3-none-any.whl\n!pip install -q ../input/uwmgiseg/dependencies/transformers-4.18.0-py3-none-any.whl\n!pip install -q ../input/uwmgiseg/dependencies/monai-0.8.1-202202162213-py3-none-any.whl\n!python ../input/uwmgiseg/setup.py develop","metadata":{"execution":{"iopub.status.busy":"2022-07-11T04:41:17.164121Z","iopub.execute_input":"2022-07-11T04:41:17.164466Z","iopub.status.idle":"2022-07-11T04:45:14.850159Z","shell.execute_reply.started":"2022-07-11T04:41:17.164383Z","shell.execute_reply":"2022-07-11T04:45:14.849258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 📚 Import Libraries ","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\npd.options.plotting.backend = \"plotly\"\nimport random\nfrom glob import glob\nimport os, shutil\nfrom tqdm import tqdm\ntqdm.pandas()\nimport time\nimport copy\nimport joblib\nfrom collections import defaultdict\nimport gc\nfrom IPython import display as ipd\n\n# visualization\nimport cv2\nimport cupy as cp\nimport gc\nimport matplotlib.pyplot as plt\n\n# PyTorch \nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda import amp\nfrom torch.cuda.amp import autocast\nimport torch.nn.functional as F\n\nimport timm\n\n# Albumentations for augmentations\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# For colored terminal text\nfrom colorama import Fore, Back, Style\nc_  = Fore.GREEN\nsr_ = Style.RESET_ALL\n\n# Warnings\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nfrom pandarallel import pandarallel\npandarallel.initialize(progress_bar=True)\n\nfrom segmentation_models_pytorch.base.modules import Activation\nfrom segmentation_models_pytorch.base import modules as md\nfrom segmentation_models_pytorch.decoders.deeplabv3.decoder import ASPP, SeparableConv2d","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-07-11T05:50:01.018137Z","iopub.execute_input":"2022-07-11T05:50:01.019991Z","iopub.status.idle":"2022-07-11T05:50:01.031343Z","shell.execute_reply.started":"2022-07-11T05:50:01.019957Z","shell.execute_reply":"2022-07-11T05:50:01.030527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append(\"../input/uwmgiseg\")\nimport cvcore\nfrom cvcore.config import get_cfg\nfrom cvcore.modeling.meta_arch import build_model","metadata":{"execution":{"iopub.status.busy":"2022-07-11T05:50:01.03321Z","iopub.execute_input":"2022-07-11T05:50:01.034144Z","iopub.status.idle":"2022-07-11T05:50:02.251671Z","shell.execute_reply.started":"2022-07-11T05:50:01.034099Z","shell.execute_reply":"2022-07-11T05:50:02.250847Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🔨 Utility","metadata":{}},{"cell_type":"code","source":"def mask2rle(msk):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    msk    = cp.array(msk)\n    pixels = msk.flatten()\n    pad    = cp.array([0])\n    pixels = cp.concatenate([pad, pixels, pad])\n    runs   = cp.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\n\ndef read_image(name):\n    img = cv2.imread(name, cv2.IMREAD_ANYDEPTH) / 65535.0\n    img = 255 * ((img - img.min()) / (img.max() - img.min()))\n    img = img.astype(np.uint8)\n    return img\n\ndef df_preprocessing(df, globbed_file_list):\n    \"\"\"The preprocessing steps applied to get column information\"\"\"\n    # 1. Get Case-ID as a column (str and int)\n    df[\"case_id_str\"] = df[\"id\"].apply(lambda x: x.split(\"_\", 2)[0])\n    df[\"case_id\"] = df[\"id\"].apply(lambda x: int(x.split(\"_\", 2)[0].replace(\"case\", \"\")))\n\n    # 2. Get Day as a column\n    df[\"day_num_str\"] = df[\"id\"].apply(lambda x: x.split(\"_\", 2)[1])\n    df[\"day_num\"] = df[\"id\"].apply(lambda x: int(x.split(\"_\", 2)[1].replace(\"day\", \"\")))\n\n    # 3. Get Slice Identifier as a column\n    df[\"slice_id\"] = df[\"id\"].apply(lambda x: x.split(\"_\", 2)[2])\n\n    # 4. Get full file paths for the representative scans\n    df[\"_partial_ident\"] = (\n        globbed_file_list[0].rsplit(\"/\", 4)[0]\n        + \"/\"\n        + df[\"case_id_str\"]  # /kaggle/input/uw-madison-gi-tract-image-segmentation/train/\n        + \"/\"\n        + df[\"case_id_str\"]  # .../case###/\n        + \"_\"\n        + df[\"day_num_str\"]\n        + \"/scans/\"  # .../case###_day##/\n        + df[\"slice_id\"]\n    )  # .../slice_####\n    _tmp_merge_df = pd.DataFrame(\n        {\n            \"_partial_ident\": [x.rsplit(\"_\", 4)[0] for x in globbed_file_list],\n            \"f_path\": globbed_file_list,\n        }\n    )\n    df = df.merge(_tmp_merge_df, on=\"_partial_ident\").drop(columns=[\"_partial_ident\"])\n\n    # 5. Get slice dimensions from filepath (int in pixels)\n    df[\"slice_h\"] = df[\"f_path\"].apply(lambda x: int(x[:-4].rsplit(\"_\", 4)[1]))\n    df[\"slice_w\"] = df[\"f_path\"].apply(lambda x: int(x[:-4].rsplit(\"_\", 4)[2]))\n\n    # 6. Pixel spacing from filepath (float in mm)\n    df[\"px_spacing_h\"] = df[\"f_path\"].apply(lambda x: float(x[:-4].rsplit(\"_\", 4)[3]))\n    df[\"px_spacing_w\"] = df[\"f_path\"].apply(lambda x: float(x[:-4].rsplit(\"_\", 4)[4]))\n\n    # 7. Reorder columns to the a new ordering (drops class and segmentation as no longer necessary)\n    new_col_order = [\n        \"id\",\n        \"f_path\",\n        \"slice_h\",\n        \"slice_w\",\n        \"px_spacing_h\",\n        \"px_spacing_w\",\n        \"case_id_str\",\n        \"case_id\",\n        \"day_num_str\",\n        \"day_num\",\n        \"slice_id\",\n    ]\n    new_col_order = [_c for _c in new_col_order if _c in df.columns]\n    df = df[new_col_order]\n    return df\n\n\ndef get_nearby_slices(id_, case_id_str, day_num_str, case_length, num_slices=3, num_strides=1):\n    slice_idx = int(id_.split(\"_\")[-1])\n    get_idxs = np.arange(slice_idx - num_slices//2 * num_strides, \n                         slice_idx + num_slices//2 * num_strides + 1, num_strides) # -7 -5 -3 -1 1 3 5 7 9\n    \n    min_idx = 2 if slice_idx%2 == 0 else 1\n    if case_length % 2 == 0:\n        max_idx = case_length if slice_idx%2 == 0 else case_length - 1\n    else:\n        max_idx = case_length if slice_idx%2 != 0 else case_length - 1\n\n    get_idxs = np.clip(get_idxs, min_idx, max_idx)\n    get_ids = [f\"{case_id_str}_{day_num_str}_slice_{slice_idx:04d}\" for slice_idx in get_idxs]\n    return get_ids","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-07-11T05:50:02.254052Z","iopub.execute_input":"2022-07-11T05:50:02.254363Z","iopub.status.idle":"2022-07-11T05:50:02.278964Z","shell.execute_reply.started":"2022-07-11T05:50:02.254323Z","shell.execute_reply":"2022-07-11T05:50:02.27811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test","metadata":{}},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/uw-madison-gi-tract-image-segmentation/\"\nTEST_DIR = os.path.join(DATA_DIR, \"test\")\nTRAIN_DIR = os.path.join(DATA_DIR, \"train\")\nSUB_CSV = os.path.join(DATA_DIR, \"sample_submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-07-11T05:50:02.281087Z","iopub.execute_input":"2022-07-11T05:50:02.281612Z","iopub.status.idle":"2022-07-11T05:50:02.288088Z","shell.execute_reply.started":"2022-07-11T05:50:02.281499Z","shell.execute_reply":"2022-07-11T05:50:02.287357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.read_csv(SUB_CSV)\n\nif not len(sub_df):\n    # Infer on train cases\n    debug = True \n    sub_df = pd.read_csv(os.path.join(DATA_DIR, \"train.csv\"))\n    sub_df = sub_df.drop(columns=['class','segmentation']).drop_duplicates()\n    paths = glob(f'/kaggle/input/uw-madison-gi-tract-image-segmentation/train/**/*png', recursive=True)\n    sub_df = df_preprocessing(sub_df, paths)\n    cases = sub_df[\"case_id_str\"].unique()[:1]\n    sub_df = sub_df[sub_df[\"case_id_str\"].isin(cases)].reset_index(drop=True)\nelse:\n    debug = False\n    sub_df = sub_df.drop(columns=['class','predicted']).drop_duplicates()\n    paths = glob(f'/kaggle/input/uw-madison-gi-tract-image-segmentation/test/**/*png',recursive=True)\n    sub_df = df_preprocessing(sub_df, paths)\n    sub_df = sub_df.reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T06:30:52.642667Z","iopub.execute_input":"2022-07-11T06:30:52.643307Z","iopub.status.idle":"2022-07-11T06:30:54.197151Z","shell.execute_reply.started":"2022-07-11T06:30:52.643265Z","shell.execute_reply":"2022-07-11T06:30:54.196347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## TH Model","metadata":{}},{"cell_type":"code","source":"# sys.path.append(\"../input/uwsegv2\")\n# from segmentation_models_pytorch.base import ClassificationHead as TClassificationHead\n# from giseg.cvcore.modeling.backbone import TimmUniversalEncoder as TTimmUniversalEncoder\n# from giseg.cvcore.modeling.backbone import mit_PLD_b4\n# from giseg.cvcore.modeling.heads import UnetDecoder as TUnetDecoder\n# from giseg.cvcore.modeling.heads import ASPPHead\n# from giseg.cvcore.modeling.heads import SegmentationHead as TSegmentationHead\n\n# class TimmUNetASPP(nn.Module):\n#     def __init__(self, arch, img_size, num_slices, num_strides):\n#         super(TimmUNetASPP, self).__init__()\n#         num_slices = 31\n#         self.encoder = TTimmUniversalEncoder(\n#             arch,\n#             pretrained=False,\n#             in_channels=num_slices,\n#             drop_path_rate=None,\n#             img_size=img_size[0],\n#         )\n#         encoder_channels = list(self.encoder.out_channels)\n#         if len(encoder_channels) == 5 or  len(encoder_channels) == 6: \n#             common_stride = 2\n#         else:\n#             common_stride = 1\n#         decoder_channels = [2048, 1024, 512, 256]\n#         n_blocks = len(decoder_channels)\n#         num_classes = 3\n\n#         self.decoder = TUnetDecoder(\n#             encoder_channels,\n#             decoder_channels,\n#             n_blocks=n_blocks,\n#             center=False,\n#             attention_type='scse',\n#             norm=\"BN\",\n#             act=\"relu\",\n#         )\n#         self.segmentation_head = TSegmentationHead(\n#             in_channels=decoder_channels[-1],\n#             out_channels=num_classes,\n#             upsampling=common_stride,\n#         )\n#         self.aux_decoder = ASPPHead(\n#                             encoder_channels = encoder_channels,\n#         )\n        \n#         self.segmentation_head_aux = TSegmentationHead(\n#             in_channels=decoder_channels[-1],\n#             out_channels=num_classes,\n#             upsampling=4,\n#         )\n#         self.classification_head = TClassificationHead(\n#             in_channels=self.encoder.out_channels[-1], classes=num_classes\n#         )\n#         self._add_hausdorff = False\n\n#     @autocast()\n#     def forward(self, images, gt_masks=None, image_sizes=None):\n#         features = self.encoder(images)\n#         decoder_output = self.decoder(*features)\n#         masks = self.segmentation_head(decoder_output)\n#         ## No need to infer this auxilary branch \n# #         aux_out = self.aux_decoder(*features) \n# #         aux_mask = self.segmentation_head_aux(aux_out)\n#         if self.training:\n#             losses = seg_criterion(masks, gt_masks, self._add_hausdorff)\n#             losses.update(seg_criterion(aux_mask,gt_masks, self._add_hausdorff, weights=[1.,1.], aux=True))\n#             gt_classes = (gt_masks.sum((2, 3)) > 0).float()\n#             cls_logits = self.classification_head(features[-1])\n#             cls_loss = cls_criterion(cls_logits, gt_classes)\n#             losses.update({\"bce_cls\": cls_loss})\n#             return losses\n#         else:\n#             # masks = masks.view(masks.shape[0], -1, 3, masks.shape[2], masks.shape[3])\n#             # masks = masks[:, masks.shape[1] // 2, ...]\n# #             return (masks + aux_mask)/2\n#             return torch.sigmoid(masks)\n\n# class TimmssFormerASPP(nn.Module):\n#     def __init__(self, img_size, num_slices, num_strides):\n#         super(TimmssFormerASPP, self).__init__()\n\n#         if num_strides == 1:\n#             num_slices = num_slices // num_strides + 1\n#         else:\n#             num_slices = num_slices\n#         num_classes = 3\n\n#         self.encoderdecoder = mit_PLD_b4(class_num=num_classes)\n#         self.segmentation_head = TSegmentationHead(\n#             in_channels=128,\n#             out_channels=num_classes,\n#             upsampling=4, \n#         )\n#         self.classification_head = TClassificationHead(\n#             in_channels=512, classes=num_classes\n#         )\n#         self._add_hausdorff = False\n\n#     @autocast()\n#     def forward(self, images, gt_masks=None, image_sizes=None):\n#         features, masks = self.encoderdecoder(images)\n#         if self.training:\n#             losses = seg_criterion(masks, gt_masks, self._add_hausdorff)\n#             gt_classes = (gt_masks.sum((2, 3)) > 0).float()\n#             cls_logits = self.classification_head(features[-1])\n#             cls_loss = cls_criterion(cls_logits, gt_classes)\n#             losses.update({\"bce_cls\": cls_loss})\n#             return losses\n#         else:\n#             # masks = masks.view(masks.shape[0], -1, 3, masks.shape[2], masks.shape[3])\n#             # masks = masks[:, masks.shape[1] // 2, ...]\n#             # return masks \n#             return torch.sigmoid(masks)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T05:50:08.28323Z","iopub.execute_input":"2022-07-11T05:50:08.283482Z","iopub.status.idle":"2022-07-11T05:50:08.290429Z","shell.execute_reply.started":"2022-07-11T05:50:08.283445Z","shell.execute_reply":"2022-07-11T05:50:08.289637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TimmUniversalEncoder(nn.Module):\n    def __init__(\n        self,\n        name,\n        pretrained=False,\n        in_channels=3,\n        drop_path_rate=0.0,\n        depth=5,\n        output_stride=32,\n        img_size=224,\n    ):\n        super().__init__()\n        kwargs = dict(\n            in_chans=in_channels,\n            features_only=True,\n            pretrained=pretrained,\n            drop_path_rate=drop_path_rate,\n            img_size=img_size,\n        )\n        kwargs.pop(\"img_size\")\n        self.model = timm.create_model(name, **kwargs)\n        if name.startswith('convnext'):\n            old_conv = self.model.stages_3.downsample[1]\n            old_in, old_out = old_conv.in_channels, old_conv.out_channels\n            self.model.stages_3.downsample[1] = nn.Conv2d(\n                old_in, old_out, kernel_size=1, stride=1\n            )\n            self.model.stages_3.downsample[1].weight.data = old_conv.weight.mean(\n                dim=(2, 3), keepdim=True\n            )\n            self.model.stages_3.downsample[1].bias.data = old_conv.bias\n        self._out_channels = [\n            in_channels,\n        ] + self.model.feature_info.channels()\n        self._depth = depth\n        self._output_stride = output_stride\n\n    def forward(self, x):\n        features = self.model(x)\n        features = [\n            x,\n        ] + features\n        return features\n\n    @property\n    def out_channels(self):\n        return self._out_channels\n\n    @property\n    def output_stride(self):\n        return min(self._output_stride, 2 ** self._depth)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T05:50:08.292054Z","iopub.execute_input":"2022-07-11T05:50:08.292477Z","iopub.status.idle":"2022-07-11T05:50:08.305586Z","shell.execute_reply.started":"2022-07-11T05:50:08.292438Z","shell.execute_reply":"2022-07-11T05:50:08.304807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SegmentationHead(nn.Sequential):\n    def __init__(self, in_channels, out_channels, kernel_size=3, activation=None, upsampling=1):\n        conv2d = nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, padding=kernel_size // 2)\n        upsampling = nn.UpsamplingBilinear2d(scale_factor=upsampling) if upsampling > 1 else nn.Identity()\n        activation = Activation(activation)\n        super().__init__(conv2d, upsampling, activation)\n\nclass SegmentationHeadDouble(nn.Sequential):\n    def __init__(self, in_channels, out_channels, kernel_size=3, activation_func=None, upsampling_scale=1):\n        conv2d = nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, padding=kernel_size // 2)\n        upsampling = nn.UpsamplingBilinear2d(scale_factor=upsampling_scale) if upsampling_scale > 1 else nn.Identity()\n        activation = Activation(activation_func)\n        conv2d_2 = nn.Conv2d(out_channels, out_channels, kernel_size=kernel_size, padding=kernel_size // 2)\n        upsampling_2 = nn.UpsamplingBilinear2d(scale_factor=upsampling_scale) if upsampling_scale > 1 else nn.Identity()\n        activation_2 = Activation(activation_func)\n        super().__init__(conv2d, upsampling, activation, conv2d_2, upsampling_2, activation_2)\n\n\nclass ClassificationHead(nn.Sequential):\n    def __init__(self, in_channels, classes, pooling=\"avg\", dropout=0.2, activation=None):\n        if pooling not in (\"max\", \"avg\"):\n            raise ValueError(\"Pooling should be one of ('max', 'avg'), got {}.\".format(pooling))\n        pool = nn.AdaptiveAvgPool2d(1) if pooling == \"avg\" else nn.AdaptiveMaxPool2d(1)\n        flatten = nn.Flatten()\n        dropout = nn.Dropout(p=dropout, inplace=True) if dropout else nn.Identity()\n        linear = nn.Linear(in_channels, classes, bias=True)\n        activation = Activation(activation)\n        super().__init__(pool, flatten, dropout, linear, activation)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T05:50:08.307188Z","iopub.execute_input":"2022-07-11T05:50:08.307478Z","iopub.status.idle":"2022-07-11T05:50:08.323305Z","shell.execute_reply.started":"2022-07-11T05:50:08.307443Z","shell.execute_reply":"2022-07-11T05:50:08.322304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DeepLabV3PlusDecoder(nn.Module):\n    def __init__(\n        self,\n        encoder_channels,\n        out_channels=256,\n        atrous_rates=(12, 24, 36),\n        output_stride=16,\n    ):\n        super().__init__()\n        if output_stride not in {8, 16}:\n            raise ValueError(\"Output stride should be 8 or 16, got {}.\".format(output_stride))\n\n        self.out_channels = out_channels\n        self.output_stride = output_stride\n\n        self.aspp = nn.Sequential(\n            ASPP(encoder_channels[-1], out_channels, atrous_rates, separable=True),\n            SeparableConv2d(out_channels, out_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(),\n        )\n\n        scale_factor = 2 if output_stride == 8 else 4\n        self.up = nn.UpsamplingBilinear2d(scale_factor=scale_factor)\n\n        highres_in_channels = encoder_channels[-4]\n        highres_out_channels = 48  # proposed by authors of paper\n        self.block1 = nn.Sequential(\n            nn.Conv2d(highres_in_channels, highres_out_channels, kernel_size=1, bias=False),\n            nn.BatchNorm2d(highres_out_channels),\n            nn.ReLU(),\n        )\n        self.block2 = nn.Sequential(\n            SeparableConv2d(\n                highres_out_channels + out_channels,\n                out_channels,\n                kernel_size=3,\n                padding=1,\n                bias=False,\n            ),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(),\n        )\n\n    def forward(self, *features):\n        aspp_features = self.aspp(features[-1])\n        aspp_features = self.up(aspp_features)\n        high_res_features = self.block1(features[-4])\n        concat_features = torch.cat([aspp_features, high_res_features], dim=1)\n        fused_features = self.block2(concat_features)\n        return fused_features\n\n    \nclass DeepLabV3PlusDecoderFix(DeepLabV3PlusDecoder):\n    def __init__(self, encoder_channels):\n        super().__init__(encoder_channels,)\n        aspp_channels = 256\n        self.block3 = nn.Sequential(\n            nn.Conv2d(aspp_channels, aspp_channels, kernel_size=1, bias=False),\n            nn.BatchNorm2d(aspp_channels),\n            nn.ReLU(),\n        )\n        self.up_2 = nn.UpsamplingBilinear2d(scale_factor=2)\n\n    def forward(self, *features):\n        aspp_features = self.aspp(features[-1])\n        aspp_features = self.up(aspp_features)\n        aspp_features = self.block3(aspp_features)\n        aspp_features = self.up_2(aspp_features)\n        high_res_features = self.block1(features[-4])\n        concat_features = torch.cat([aspp_features, high_res_features], dim=1)\n        fused_features = self.block2(concat_features)\n        return fused_features\n    \n\nclass DecoderBlock(nn.Module):\n    def __init__(\n        self,\n        in_channels,\n        skip_channels,\n        out_channels,\n        use_batchnorm=True,\n        attention_type=None,\n    ):\n        super().__init__()\n        self.conv1 = md.Conv2dReLU(\n            in_channels + skip_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        self.attention1 = md.Attention(attention_type, in_channels=in_channels + skip_channels)\n        self.conv2 = md.Conv2dReLU(\n            out_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        self.attention2 = md.Attention(attention_type, in_channels=out_channels)\n\n    def forward(self, x, skip=None, scale=True):\n        if scale:\n            x = F.interpolate(x, scale_factor=2, mode=\"nearest\")\n        if skip is not None:\n            x = torch.cat([x, skip], dim=1)\n            x = self.attention1(x)\n        x = self.conv1(x)\n        x = self.conv2(x)\n        x = self.attention2(x)\n        return x\n\n\nclass CenterBlock(nn.Sequential):\n    def __init__(self, in_channels, out_channels, use_batchnorm=True):\n        conv1 = md.Conv2dReLU(\n            in_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        conv2 = md.Conv2dReLU(\n            out_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        super().__init__(conv1, conv2)\n        \n        \nclass UnetPlusPlusDecoder(nn.Module):\n    def __init__(\n        self,\n        encoder_channels,\n        decoder_channels=(256, 128, 64, 32, 16),\n        n_blocks=5,\n        use_batchnorm=True,\n        attention_type=None,\n        center=False,\n    ):\n        super().__init__()\n\n        if n_blocks != len(decoder_channels):\n            raise ValueError(\n                \"Model depth is {}, but you provide `decoder_channels` for {} blocks.\".format(\n                    n_blocks, len(decoder_channels)\n                )\n            )\n\n        # remove first skip with same spatial resolution\n        encoder_channels = encoder_channels[1:]\n        # reverse channels to start from head of encoder\n        encoder_channels = encoder_channels[::-1]\n\n        # computing blocks input and output channels\n        head_channels = encoder_channels[0]\n        self.in_channels = [head_channels] + list(decoder_channels[:-1])\n        self.skip_channels = list(encoder_channels[1:]) + [0]\n        self.out_channels = decoder_channels\n        if center:\n            self.center = CenterBlock(head_channels, head_channels, use_batchnorm=use_batchnorm)\n        else:\n            self.center = nn.Identity()\n\n        # combine decoder keyword arguments\n        kwargs = dict(use_batchnorm=use_batchnorm, attention_type=attention_type)\n\n        blocks = {}\n        for layer_idx in range(len(self.in_channels) - 1):\n            for depth_idx in range(layer_idx + 1):\n                if depth_idx == 0:\n                    in_ch = self.in_channels[layer_idx]\n                    skip_ch = self.skip_channels[layer_idx] * (layer_idx + 1)\n                    out_ch = self.out_channels[layer_idx]\n                else:\n                    out_ch = self.skip_channels[layer_idx]\n                    skip_ch = self.skip_channels[layer_idx] * (layer_idx + 1 - depth_idx)\n                    in_ch = self.skip_channels[layer_idx - 1]\n                blocks[f\"x_{depth_idx}_{layer_idx}\"] = DecoderBlock(in_ch, skip_ch, out_ch, **kwargs)\n        blocks[f\"x_{0}_{len(self.in_channels)-1}\"] = DecoderBlock(\n            self.in_channels[-1], 0, self.out_channels[-1], **kwargs\n        )\n        self.blocks = nn.ModuleDict(blocks)\n        self.depth = len(self.in_channels) - 1\n\n    def forward(self, *features):\n        features = features[1:]  # remove first skip with same spatial resolution\n        features = features[::-1]  # reverse channels to start from head of encoder\n        # start building dense connections\n        dense_x = {}\n        for layer_idx in range(len(self.in_channels) - 1):\n            for depth_idx in range(self.depth - layer_idx):\n                if layer_idx == 0:\n                    output = self.blocks[f\"x_{depth_idx}_{depth_idx}\"](features[depth_idx], features[depth_idx + 1])\n                    dense_x[f\"x_{depth_idx}_{depth_idx}\"] = output\n                else:\n                    dense_l_i = depth_idx + layer_idx\n                    cat_features = [dense_x[f\"x_{idx}_{dense_l_i}\"] for idx in range(depth_idx + 1, dense_l_i + 1)]\n                    cat_features = torch.cat(cat_features + [features[dense_l_i + 1]], dim=1)\n                    dense_x[f\"x_{depth_idx}_{dense_l_i}\"] = self.blocks[f\"x_{depth_idx}_{dense_l_i}\"](\n                        dense_x[f\"x_{depth_idx}_{dense_l_i-1}\"], cat_features\n                    )\n        dense_x[f\"x_{0}_{self.depth}\"] = self.blocks[f\"x_{0}_{self.depth}\"](dense_x[f\"x_{0}_{self.depth-1}\"])\n        return dense_x[f\"x_{0}_{self.depth}\"]\n\n\nclass UnetPlusPlusDecoderFix(nn.Module):\n    def __init__(\n        self,\n        encoder_channels,\n        decoder_channels=(256, 128, 64, 32),\n        n_blocks=4,\n        use_batchnorm=True,\n        attention_type=None,\n        center=False,\n    ):\n        super().__init__()\n\n        if n_blocks != len(decoder_channels):\n            raise ValueError(\n                \"Model depth is {}, but you provide `decoder_channels` for {} blocks.\".format(\n                    n_blocks, len(decoder_channels)\n                )\n            )\n\n        # remove first skip with same spatial resolution\n        encoder_channels = encoder_channels[1:]\n        # reverse channels to start from head of encoder\n        encoder_channels = encoder_channels[::-1]\n\n        # computing blocks input and output channels\n        head_channels = encoder_channels[0]\n        self.in_channels = [head_channels] + list(decoder_channels[:-1])\n        self.skip_channels = list(encoder_channels[1:]) + [0]\n        self.out_channels = decoder_channels\n        if center:\n            self.center = CenterBlock(head_channels, head_channels, use_batchnorm=use_batchnorm)\n        else:\n            self.center = nn.Identity()\n\n        # combine decoder keyword arguments\n        kwargs = dict(use_batchnorm=use_batchnorm, attention_type=attention_type)\n\n        blocks = {}\n        for layer_idx in range(len(self.in_channels) - 1):\n            for depth_idx in range(layer_idx + 1):\n                if depth_idx == 0:\n                    in_ch = self.in_channels[layer_idx]\n                    skip_ch = self.skip_channels[layer_idx] * (layer_idx + 1)\n                    out_ch = self.out_channels[layer_idx]\n                else:\n                    out_ch = self.skip_channels[layer_idx]\n                    skip_ch = self.skip_channels[layer_idx] * (layer_idx + 1 - depth_idx)\n                    in_ch = self.skip_channels[layer_idx - 1]\n                blocks[f\"x_{depth_idx}_{layer_idx}\"] = DecoderBlock(in_ch, skip_ch, out_ch, **kwargs)\n        blocks[f\"x_{0}_{len(self.in_channels)-1}\"] = DecoderBlock(\n            self.in_channels[-1], 0, self.out_channels[-1], **kwargs\n        )\n        self.blocks = nn.ModuleDict(blocks)\n        self.depth = len(self.in_channels) - 1\n\n    def forward(self, *features):\n        features = features[1:]  # remove first skip with same spatial resolution\n        features = features[::-1]  # reverse channels to start from head of encoder\n        # start building dense connections\n        dense_x = {}\n        for layer_idx in range(len(self.in_channels) - 1):\n            for depth_idx in range(self.depth - layer_idx):\n                if layer_idx == 0:\n                    output = self.blocks[f\"x_{depth_idx}_{depth_idx}\"](features[depth_idx], features[depth_idx + 1], depth_idx >= 1)\n                    dense_x[f\"x_{depth_idx}_{depth_idx}\"] = output\n                else:\n                    dense_l_i = depth_idx + layer_idx\n                    cat_features = [dense_x[f\"x_{idx}_{dense_l_i}\"] for idx in range(depth_idx + 1, dense_l_i + 1)]\n                    cat_features = torch.cat(cat_features + [features[dense_l_i + 1]], dim=1)\n                    dense_x[f\"x_{depth_idx}_{dense_l_i}\"] = self.blocks[f\"x_{depth_idx}_{dense_l_i}\"](\n                        dense_x[f\"x_{depth_idx}_{dense_l_i-1}\"], cat_features\n                    )\n        dense_x[f\"x_{0}_{self.depth}\"] = self.blocks[f\"x_{0}_{self.depth}\"](dense_x[f\"x_{0}_{self.depth-1}\"])\n        return dense_x[f\"x_{0}_{self.depth}\"]","metadata":{"execution":{"iopub.status.busy":"2022-07-11T05:50:08.326358Z","iopub.execute_input":"2022-07-11T05:50:08.326831Z","iopub.status.idle":"2022-07-11T05:50:08.376888Z","shell.execute_reply.started":"2022-07-11T05:50:08.326791Z","shell.execute_reply":"2022-07-11T05:50:08.37609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BaselineSegTimm(nn.Module):\n    def __init__(self, index, num_stride):\n        super(BaselineSegTimm, self).__init__()\n\n        self.encoder = TimmUniversalEncoder(\n            CFG.ENCODER[index],\n            in_channels=num_stride,\n            img_size=CFG.img_size[index][0],\n        )\n\n        with torch.no_grad():\n            dummy_inputs = torch.randn(2, num_stride, *CFG.img_size[index])\n            out = self.encoder(dummy_inputs)\n            common_stride = CFG.img_size[index][0] // out[1].shape[2]\n        encoder_channels = self.encoder.out_channels\n        num_classes = 3\n\n        if CFG.ARCH[index] == \"DeepLabV3Plus\":\n            self.decoder = DeepLabV3PlusDecoder(\n                encoder_channels=encoder_channels,\n            )\n        elif CFG.ARCH[index] == \"DeepLabV3PlusFix\":\n            self.decoder = DeepLabV3PlusDecoderFix(\n                encoder_channels=encoder_channels,\n            )\n        elif CFG.ARCH[index] == \"UnetPlusPlus\":\n            self.decoder = UnetPlusPlusDecoder(\n                encoder_channels=encoder_channels,\n            )\n        elif CFG.ARCH[index] == \"UnetPlusPlusFix\":\n            self.decoder = UnetPlusPlusDecoderFix(\n                encoder_channels=encoder_channels,\n            )\n\n        with torch.no_grad():\n            out = self.decoder(*out)\n        \n        if CFG.ARCH[index] == \"UnetPlusPlus\":\n            self.segmentation_head = SegmentationHead(\n                    in_channels=out.shape[1],\n                    out_channels=num_classes,\n                    activation=None,\n                    kernel_size=3,\n                    upsampling=1,\n                )\n        elif CFG.ARCH[index] == \"UnetPlusPlusFix\":\n            self.segmentation_head = SegmentationHead(\n                    in_channels=out.shape[1],\n                    out_channels=num_classes,\n                    activation=None,\n                    kernel_size=3,\n                    upsampling=common_stride//2,\n                )\n        else:\n            if CFG.ENCODER[index].startswith('tf_efficientnet') or CFG.ENCODER[index].startswith('ecaresnet'):\n                self.segmentation_head = SegmentationHeadDouble(\n                    in_channels=out.shape[1],\n                    out_channels=num_classes,\n                    activation_func=None,\n                    kernel_size=3,\n                    upsampling_scale=common_stride,\n                )\n            else:\n                self.segmentation_head = SegmentationHead(\n                    in_channels=out.shape[1],\n                    out_channels=num_classes,\n                    activation=None,\n                    kernel_size=3,\n                    upsampling=common_stride,\n                )\n\n    @autocast()\n    def forward(self, images, gt_masks=None):\n        features = self.encoder(images)\n        decoder_output = self.decoder(*features)\n        masks = self.segmentation_head(decoder_output)\n        return torch.sigmoid(masks)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T05:50:08.378391Z","iopub.execute_input":"2022-07-11T05:50:08.378828Z","iopub.status.idle":"2022-07-11T05:50:08.394947Z","shell.execute_reply.started":"2022-07-11T05:50:08.378788Z","shell.execute_reply":"2022-07-11T05:50:08.393859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    ARCH = ['DeepLabV3Plus', \n#             'DeepLabV3PlusFix',\n            'Unet3d', \n            'Unet3d',\n            'Unet',\n            'UnetPlusPlusFix', 'UnetPlusPlus',\n            'Unet'\n    ]\n    YAML = [None, \n#             None,\n            \"../input/uwmgiseg/configs/cxl-unetplus3d.yaml\",\n            \"../input/uwmgiseg/configs/cb-unetacs.yaml\",\n            \"../input/uwmgiseg/configs/cb-acs-unet.yaml\",\n            None, None, \n            \"../input/uwmgiseg/configs/cb-unet.yaml\"\n           ]\n    ENCODER = ['convnext_xlarge_in22ft1k', \n#                'tf_efficientnetv2_l',\n               'convnext_xlarge',\n               'convnext_base', \n               'convnext_base',\n               'convnext_xlarge_in22ft1k', 'tf_efficientnetv2_l',\n               'convnext_base'\n              ]\n    WEIGHTS = [\n       '../input/newgiseg/deeplabv3plus_convnext_xlarge_17_608_fold-1_e28.pth',\n#        '../input/uwmgiweights/deeplabv3plus_v2l_17_608_fold-1_e28.pth',\n       '../input/uwmgiseg/weights/convnext_xlarge_unetplus3d_e1280.pth',\n       '../input/uwmgiseg/weights/convnext_base_unet3d_2000.pth',\n       '../input/uwmgiseg/weights/convnext_base_acs_unet_epoch30.pth',\n       '../input/uwmgiweights/unetplus_convnext_xlarge_fold-1_e28.pth',\n       '../input/uwmgiweights/unetplus_v2l_fold-1_e28.pth',\n       '../input/uwmgiseg/weights/convnext_base_unet_epoch10.pth',\n              ]\n    NUM_CLASSES = 3\n\n    img_size = [[640, 640], \n#                 [640, 640], \n                [224, 224],\n                [224, 224], \n                [512, 512], [640, 640], [640, 640], [512, 512]]\n    slices = [17, \n#               17, \n              80,\n              80, \n              33, 9, 9, 5]\n    strides = [2, \n#                2,\n               1,\n               1, \n               1, 2, 3, 1]","metadata":{"execution":{"iopub.status.busy":"2022-07-11T05:50:08.396348Z","iopub.execute_input":"2022-07-11T05:50:08.396829Z","iopub.status.idle":"2022-07-11T05:50:08.40718Z","shell.execute_reply.started":"2022-07-11T05:50:08.396789Z","shell.execute_reply":"2022-07-11T05:50:08.406351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GISegDataset(Dataset):\n    def __init__(self, df):\n        super().__init__()\n        df[\"case_day_str\"] = df[\"case_id_str\"] + \"_\" + df[\"day_num_str\"]\n        self.cases_length = df[\"case_day_str\"].value_counts().to_dict()\n        self.images_dict = {id: f_path for id, f_path in zip(df[\"id\"].values, df[\"f_path\"].values)}\n        self.aug512 = A.Compose([A.Resize(512, 512), ToTensorV2()])\n        self.aug640 = A.Compose([A.Resize(640, 640), ToTensorV2()])\n        self.df = df\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        info = self.df.iloc[idx]\n        center_id = info[\"id\"]\n        case_length = self.cases_length.get(info[\"case_day_str\"])\n        case_id = info[\"case_id_str\"]\n        day = info[\"day_num_str\"]\n        \n        ids172 = get_nearby_slices(center_id, case_id, day, case_length,\n                                   num_slices=17, num_strides=2)\n        ids331 = get_nearby_slices(center_id, case_id, day, case_length,\n                                   num_slices=33, num_strides=1)\n        ids92 = get_nearby_slices(center_id, case_id, day, case_length,\n                                  num_slices=9, num_strides=2)\n        ids93 = get_nearby_slices(center_id, case_id, day, case_length,\n                                  num_slices=9, num_strides=3)\n        \n        img_ids = set(ids172 + ids331 + ids92 + ids93)\n        \n        imgs_dict = {id_: read_image(self.images_dict.get(id_)) for id_ in img_ids}\n        h, w = imgs_dict.get(ids172[0]).shape\n        \n        img172 = np.stack([imgs_dict.get(id_) for id_ in ids172], axis=-1)\n        img172 = self.aug640(image=img172)[\"image\"].float() / 255.\n        \n        img331 = np.stack([imgs_dict.get(id_) for id_ in ids331], axis=-1)\n        img331 = self.aug512(image=img331)[\"image\"].float() / 255.\n        img51 = img331[14:19]\n        \n        idx92 = []\n        for i in ids92:\n            idx92.append(ids172.index(i))\n        img92 = img172[idx92]\n            \n        img93 = np.stack([imgs_dict.get(id_) for id_ in ids93], axis=-1)\n        img93 = self.aug640(image=img93)[\"image\"].float() / 255.\n        \n        img = imgs_dict.get(center_id)\n        img = torch.from_numpy(img).float() / 255.\n    \n        return img, img172, img331, img92, img93, img51, center_id, h, w","metadata":{"execution":{"iopub.status.busy":"2022-07-11T05:50:08.408732Z","iopub.execute_input":"2022-07-11T05:50:08.409212Z","iopub.status.idle":"2022-07-11T05:50:08.427205Z","shell.execute_reply.started":"2022-07-11T05:50:08.409173Z","shell.execute_reply":"2022-07-11T05:50:08.426368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🔭 Inference","metadata":{}},{"cell_type":"code","source":"# all models\ndef load_weight(wt, model, cls=False):\n    ckpt = torch.load(wt, \"cpu\")\n    if not cls:\n        ckpt[\"model\"] = {k: v for k, v in ckpt[\"model\"].items() if \"classification_head\" not in k}\n    else:\n        ckpt[\"model\"] = {k: v for k, v in ckpt[\"model\"].items()}\n    model.load_state_dict(ckpt.pop(\"model\"))\n    print(wt, ckpt[\"best_metric\"])\n    del ckpt; gc.collect()\n    model.eval()\n    model = model.cuda()\n    return model\n\nall_models = []\nfor i in range(len(CFG.WEIGHTS)):\n    if CFG.YAML[i] is None:\n        model = BaselineSegTimm(i, CFG.slices[i])\n        model = load_weight(CFG.WEIGHTS[i], model)\n    else:\n        cfg = get_cfg()\n        cfg.merge_from_file(CFG.YAML[i])\n        cfg.MODEL.BACKBONE.PRETRAINED = False\n        model = build_model(cfg)\n        model = load_weight(CFG.WEIGHTS[i], model)\n    all_models.append(model)\n    del model","metadata":{"execution":{"iopub.status.busy":"2022-07-11T05:50:08.430805Z","iopub.execute_input":"2022-07-11T05:50:08.4318Z","iopub.status.idle":"2022-07-11T05:53:11.63532Z","shell.execute_reply.started":"2022-07-11T05:50:08.431762Z","shell.execute_reply":"2022-07-11T05:53:11.634508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from monai.inferers import sliding_window_inference\n\ndef infer3d(test_inputs, model, roi_size=(224, 224, 80), sw_batch_size=2):\n    # do 4 tta\n    pred_all = sliding_window_inference(test_inputs, roi_size, \n                                        sw_batch_size, model)\n    for dims in [[2], [3], [2, 3]]:\n        pred_all += torch.flip(\n            sliding_window_inference(torch.flip(test_inputs, dims=dims), roi_size, \n                                     sw_batch_size, model),\n            dims=dims)\n    pred_all = torch.sigmoid(pred_all / 4.)\n    return pred_all\n\n\ndef masks2rles(masks, ids):\n    pred_strings = []; pred_ids = []; pred_classes = []\n    masks[:, 0] = torch.where(masks[:, 0] >= 0.3, 1, 0)\n    masks[:, 1] = torch.where(masks[:, 1] >= 0.3, 1, 0)\n    masks[:, 2] = torch.where(masks[:, 2] >= 0.4, 1, 0)\n    masks = masks.to(torch.uint8).permute(0, 2, 3, 1).cpu().numpy()\n    for idx, mask in enumerate(masks):\n        rle = [None] * 3\n        for class_idx in [0, 1, 2]:\n            rle[class_idx] = mask2rle(mask[..., class_idx])\n        pred_strings.extend(rle)\n        pred_ids.extend([ids[idx]] * 3)\n        pred_classes.extend(['large_bowel', 'small_bowel', 'stomach'])\n    return pred_strings, pred_ids, pred_classes\n\n# IMAGES = []\n# MASKS = []\n\n@torch.no_grad()\ndef inference(models, test_loader):\n    pred_strings = []; pred_ids = []; pred_classes = []\n    case_img = []\n    case_mask = []\n    slice_ids = []\n    # 2.5d\n    for imgs, imgs172, imgs331, imgs92, imgs93, imgs51, ids, heights, widths in test_loader:\n        case_img.append(imgs)\n        slice_ids.extend(ids)\n        imgs172 = imgs172.half().cuda(non_blocking=True)\n        imgs331 = imgs331.half().cuda(non_blocking=True)\n        imgs92 = imgs92.half().cuda(non_blocking=True)\n        imgs93 = imgs93.half().cuda(non_blocking=True)\n        imgs51 = imgs51.half().cuda(non_blocking=True)\n        masks = 0\n        with autocast():\n            for model_index, model in enumerate(models):\n#                 if model_index in [2]:\n                if model_index in [1, 2]:\n                    continue\n                else:\n                    if model_index == 0:\n                        out = model(imgs172)\n#                     elif model_index == 1:\n#                         out = model(imgs172)\n                    elif model_index == 3:\n                        out = model(imgs331)\n                    elif model_index == 4:\n                        out = model(imgs92)\n                    elif model_index == 5:\n                        out = model(imgs93)\n                    elif model_index == 6:\n                        out = model(imgs51)\n                    out = F.interpolate(out, size=(heights[0].item(), widths[0].item()), \n                                        mode='bilinear', align_corners=False)\n                    masks += out / len(models)\n        case_mask.append(masks)\n    del imgs172, imgs331, imgs92, imgs93, imgs51\n    gc.collect()\n    torch.cuda.empty_cache()\n    # 3d\n    case_img = torch.cat(case_img).permute(1, 2, 0).unsqueeze(0).unsqueeze(0)\n    case_img = case_img.cuda(non_blocking=True)\n    case_mask = torch.cat(case_mask)\n#     for model_index in [2]:\n    for model_index in [1, 2]:\n        out = infer3d(case_img, models[model_index])[0]\n        out = torch.permute(out, [3, 0, 1, 2])\n        case_mask += out / len(models)\n#     print(case_mask.min(), case_mask.max())\n    result = masks2rles(case_mask, slice_ids)\n    pred_strings.extend(result[0])\n    pred_ids.extend(result[1])\n    pred_classes.extend(result[2])\n    del result\n    gc.collect()\n    return pred_strings, pred_ids, pred_classes","metadata":{"execution":{"iopub.status.busy":"2022-07-11T06:34:02.027586Z","iopub.execute_input":"2022-07-11T06:34:02.028039Z","iopub.status.idle":"2022-07-11T06:34:02.050847Z","shell.execute_reply.started":"2022-07-11T06:34:02.027996Z","shell.execute_reply":"2022-07-11T06:34:02.048647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_strings = []; pred_ids = []; pred_classes = []\nfor group in tqdm(sub_df.groupby([\"case_id_str\", \"day_num_str\"])):\n    case_id_str, day_num_str = group[0]\n    group_id = case_id_str + \"_\" + day_num_str\n    group_df = group[1].sort_values(\"slice_id\", ascending=True).reset_index(drop=True)\n    test_dataset = GISegDataset(group_df)\n    test_loader = DataLoader(test_dataset, batch_size=4,\n                             num_workers=2, shuffle=False, pin_memory=True)\n    results = inference(all_models, test_loader)\n    pred_strings.extend(results[0])\n    pred_ids.extend(results[1])\n    pred_classes.extend(results[2])","metadata":{"execution":{"iopub.status.busy":"2022-07-11T06:34:05.338987Z","iopub.execute_input":"2022-07-11T06:34:05.340023Z","iopub.status.idle":"2022-07-11T06:42:58.218767Z","shell.execute_reply.started":"2022-07-11T06:34:05.339974Z","shell.execute_reply":"2022-07-11T06:42:58.217964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# del IMAGES, MASKS\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-07-11T06:44:49.16913Z","iopub.execute_input":"2022-07-11T06:44:49.169422Z","iopub.status.idle":"2022-07-11T06:44:49.438261Z","shell.execute_reply.started":"2022-07-11T06:44:49.169383Z","shell.execute_reply":"2022-07-11T06:44:49.437512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 📝 Submission","metadata":{}},{"cell_type":"code","source":"pred_df = pd.DataFrame({\n    \"id\":pred_ids,\n    \"class\":pred_classes,\n    \"predicted\":pred_strings\n})","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-07-11T06:44:52.665127Z","iopub.execute_input":"2022-07-11T06:44:52.665742Z","iopub.status.idle":"2022-07-11T06:44:52.67152Z","shell.execute_reply.started":"2022-07-11T06:44:52.665698Z","shell.execute_reply":"2022-07-11T06:44:52.670783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not debug:\n    sub_df = pd.read_csv('../input/uw-madison-gi-tract-image-segmentation/sample_submission.csv')\n    del sub_df['predicted']\nelse:\n    sub_df = pd.read_csv('../input/uw-madison-gi-tract-image-segmentation/train.csv')\n    del sub_df['segmentation']\n    sub_df = sub_df[sub_df['id'].apply(lambda x: any([x.startswith(case_id) for case_id in cases]))]\n\nassert len(sub_df) == len(pred_df)\nsub_df = sub_df.merge(pred_df, on=['id','class'])","metadata":{"execution":{"iopub.status.busy":"2022-07-11T06:45:37.793015Z","iopub.execute_input":"2022-07-11T06:45:37.793296Z","iopub.status.idle":"2022-07-11T06:45:38.258026Z","shell.execute_reply.started":"2022-07-11T06:45:37.793265Z","shell.execute_reply":"2022-07-11T06:45:38.257229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df.to_csv('submission.csv',index=False)\nprint(sub_df.head(10))","metadata":{"execution":{"iopub.status.busy":"2022-07-11T06:45:39.837986Z","iopub.execute_input":"2022-07-11T06:45:39.838665Z","iopub.status.idle":"2022-07-11T06:45:39.864016Z","shell.execute_reply.started":"2022-07-11T06:45:39.838623Z","shell.execute_reply":"2022-07-11T06:45:39.863257Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}