{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":1144103,"sourceType":"datasetVersion","datasetId":645458},{"sourceId":4053053,"sourceType":"datasetVersion","datasetId":2398033},{"sourceId":5220538,"sourceType":"datasetVersion","datasetId":3037137},{"sourceId":7284438,"sourceType":"datasetVersion","datasetId":4223999},{"sourceId":7413748,"sourceType":"datasetVersion","datasetId":4079395},{"sourceId":7513583,"sourceType":"datasetVersion","datasetId":4348159},{"sourceId":126502507,"sourceType":"kernelVersion"},{"sourceId":150248402,"sourceType":"kernelVersion"}],"dockerImageVersionId":30580,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Prepare","metadata":{}},{"cell_type":"code","source":"!pip install monai lovely-numpy -q --no-index --find-links=../input/vesuvis-downloads\n!python -m pip install -q --no-index --find-links=/kaggle/input/pip-download-for-segmentation-models-pytorch segmentation-models-pytorch","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-24T10:17:06.276782Z","iopub.execute_input":"2024-01-24T10:17:06.277056Z","iopub.status.idle":"2024-01-24T10:17:37.719310Z","shell.execute_reply.started":"2024-01-24T10:17:06.277021Z","shell.execute_reply":"2024-01-24T10:17:37.718085Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!python -m pip install -q /kaggle/input/omegaconf222py3/omegaconf-2.2.2-py3-none-any.whl --no-index --find-links=/kaggle/input/omegaconf222py3/","metadata":{"execution":{"iopub.status.busy":"2024-01-24T10:17:37.721530Z","iopub.execute_input":"2024-01-24T10:17:37.722330Z","iopub.status.idle":"2024-01-24T10:17:49.724653Z","shell.execute_reply.started":"2024-01-24T10:17:37.722291Z","shell.execute_reply":"2024-01-24T10:17:49.723591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip -q install '/kaggle/input/connected-components-3d/pbr-5.11.1-py2.py3-none-any.whl'\n!pip install /kaggle/input/installation-connected-components-3d/connected_components_3d-3.12.4-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n\nimport cc3d","metadata":{"execution":{"iopub.status.busy":"2024-01-24T10:17:49.725924Z","iopub.execute_input":"2024-01-24T10:17:49.726201Z","iopub.status.idle":"2024-01-24T10:18:53.221795Z","shell.execute_reply.started":"2024-01-24T10:17:49.726176Z","shell.execute_reply":"2024-01-24T10:18:53.220877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append(\"../input/ttach-kaggle/\")\n\nimport ttach as tta","metadata":{"execution":{"iopub.status.busy":"2024-01-24T10:18:53.224301Z","iopub.execute_input":"2024-01-24T10:18:53.225001Z","iopub.status.idle":"2024-01-24T10:18:56.384267Z","shell.execute_reply.started":"2024-01-24T10:18:53.224963Z","shell.execute_reply":"2024-01-24T10:18:56.383497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport cv2\nfrom glob import glob\nfrom tqdm import tqdm\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nimport gc\nimport monai\n\nimport re","metadata":{"execution":{"iopub.status.busy":"2024-01-24T10:18:56.385497Z","iopub.execute_input":"2024-01-24T10:18:56.385956Z","iopub.status.idle":"2024-01-24T10:19:36.805357Z","shell.execute_reply.started":"2024-01-24T10:18:56.385923Z","shell.execute_reply":"2024-01-24T10:19:36.804615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\n\n\nclass ImagePadder:\n    def __init__(self, tensor):\n        self.tensor = tensor\n        self.original_size = tensor.shape[-2:]\n        self.pad_size = self.calculate_padding(self.original_size)\n\n    def calculate_padding(self, size):\n        \"\"\"\n        Calculate how much padding is needed to make the width and height divisible by 32.\n        \"\"\"\n        height, width = size\n        pad_height = (32 - height % 32) % 32\n        pad_width = (32 - width % 32) % 32\n        return (pad_height, pad_width)\n\n    def pad(self):\n        \"\"\"\n        Pad the image so that its height and width are divisible by 32.\n        \"\"\"\n        pad_height, pad_width = self.pad_size\n        # Apply padding equally on both sides of the height and width\n        padding = (pad_width // 2, pad_width - pad_width // 2, pad_height // 2, pad_height - pad_height // 2)\n        self.tensor = F.pad(self.tensor, padding)\n        return self.tensor\n\n    def unpad(self, tensor=None):\n        \"\"\"\n        Remove the padding from the image.\n        \"\"\"\n        pad_height, pad_width = self.pad_size\n        # Crop the image to remove the padding\n        return tensor[:, :, pad_height // 2: -pad_height // 2, pad_width // 2: -pad_width // 2]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predict_on = [\n# #     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 3, \"UnetPlusPlus_tu-tf_efficientnet_b5_src.models.components.losses.BoundaryDoULoss_size_3_512_bs32_hard_minmax_organ_novoi_fx\", 1., None, \"senet-hoa\"], #085 !!!!!#minmax organ norm\n# # #     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 3, \"UnetPlusPlus_tu-tf_efficientnet_b5_src.models.components.losses.BoundaryDoULoss_size_3_512_bs32_hard_minmax_organ_pseudo\", 1., None, \"senet-hoa\"], #0845 !!!!!#minmax organ norm\n#     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 3, \"UnetPlusPlus_tu-tf_efficientnet_b5_src.models.components.losses.BoundaryDoULoss_size_3_512_bs32_hard_minmax_organ_pseudo_th_v0\", 0.75, None, \"senet-hoa\"], #0871 !!!!!#minmax organ norm\n# #     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b6\", 3, \"UnetPlusPlus_tu-tf_efficientnet_b6_src.models.components.losses.BoundaryDoULoss_size_3_512_bs32_hard_minmax_organ_novoi_fx\", 1.5, None, \"senet-hoa\"], #0857 !!!!!#minmax organ norm\n# #    [\"UnetPlusPlus\", \"tu-tf_efficientnet_b6\", 3, \"UnetPlusPlus_tu-tf_efficientnet_b6_src.models.components.losses.BoundaryDoULoss_size_3_512_bs32_hard_minmax_organ_pseudo_th\", 1., None, \"senet-hoa\"], #0869\n# #    [\"UnetPlusPlus\", \"tu-tf_efficientnet_b6\", 3, \"UnetPlusPlus_tu-tf_efficientnet_b6_src.models.components.losses.BoundaryDoULoss_size_3_512_bs32_hard_minmax_organ_sd\", 1., None, \"senet-hoa\"],\n#      [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 3, \"UnetPlusPlus_tu-tf_efficientnet_b5_src.models.components.losses.BoundaryDoULoss_size_3_512_bs32_hard_minmax_organ_pseudo_th\", 1.25, \"scse\", \"senet-hoa\"],  #0873\n# #     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 3, \"UnetPlusPlus_tu-tf_efficientnet_b5_src.models.components.losses.BoundaryDoULoss_size_3_512_bs32_severe_minmax_organ_pseudo_th\", 1., \"scse\", \"senet-hoa\"]  #0855\n# #     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 3, \"UnetPlusPlus_tu-tf_efficientnet_b5_src.models.components.losses.BoundaryDoULoss_size_3_512_bs32_severe_v2_minmax_organ_pseudo_th\", 1., \"scse\", \"senet-hoa\"]  #0855\n# #     [\"UnetPlusPlus\", \"tu-tf_efficientnetv2_m\", 3, \"UnetPlusPlus_tu-tf_efficientnetv2_m_src.models.components.losses.BoundaryDoULoss_size_3_512_bs32_hard_minmax_organ_pseudo_th\", 1., \"scse\", \"senet-hoa\"] #0864\n# #     [\"UnetPlusPlus\", \"tu-maxvit_tiny_tf_512\", 3, \"UnetPlusPlus_tu-maxvit_tiny_tf_512.in1k_size_3_512_bs8_hard_pseudo\", 1., \"scse\", \"senet-hoa\"], #0859\n#      [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 1, \"UnetPlusPlus_tu-tf_efficientnet_b5_size_1_512_bs32_hard_pseudo\", 1.25, \"scse\", \"senet-hoa\"],  #0873\n# #     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 1, \"UnetPlusPlus_tu-tf_efficientnet_b5_size_1_512_bs32_hard_pseudo_a\", 0.75, \"scse\", \"senet-hoa\"],  #0862 @0.3, 0827 @0.05\n# #     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b7\", 1, \"UnetPlusPlus_tu-tf_efficientnet_b7_size_1_512_bs16_hard_pseudo_a\", 1., \"scse\", \"senet-hoa\"],  #0\n#     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 1, \"UnetPlusPlus_tu-tf_efficientnet_b5_size_1_512_bs16_hard_pseudo_b\", 0.75, \"scse\", \"senet-hoa\"],  #087   \n# #     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 1, \"UnetPlusPlus_tu-tf_efficientnet_b5_size_1_512_bs32_hard_pseudo_noval\", 1., \"scse\", \"senet-hoa\"],  #0864\n#     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 3, \"UnetPlusPlus_tu-tf_efficientnet_b5_size_3_512_bs32_hard_pseudo_noval\", 1.5, \"scse\", \"senet-hoa\"],  #0872\n# #     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 1, \"UnetPlusPlus_tu-tf_efficientnet_b5_size_1_512_bs32_hard_pseudo_rot\", 1., \"scse\", \"senet-hoa\"],  #0\n #   [\"Unet\", \"mit_b5\", 3, \"Unet_mit_b5_size_3_512_bs32_hard_pseudo\", 1., \"scse\", \"senet-models\"],  #0842\n#      [\"UnetPlusPlus\", \"tu-tf_efficientnet_b7\", 3, \"UnetPlusPlus_tu-tf_efficientnet_b7_size_3_512_bs32_hard_pseudo\", 1., \"scse\", \"senet-models\"],  #0847\n#      [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 3, \"DUAT_size_3_512_bs32_hard_pseudo_no_bce\", 1., \"scse\", \"senet-models\"],  #0855\n#      [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 1, \"UnetPlusPlus_tu-tf_efficientnet_b5_size_1_512_bs32_hard_pseudo_v2\", 1., \"scse\", \"senet-models\"],  #0855\n     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 1, \"UnetPlusPlus_tu-tf_efficientnet_b5_BoundaryDoULoss_size_1_512_bs32_hard_pseudo_v2\", 1., \"scse\", \"senet-models\"],  #0 new sampling + cutmix\n    \n    \n]\n\n\ntta_models = []\nweights = []\n\nuse_top_only = False #True\nuse_best = False # False\nfolds2predict = [0, 1, -1]\n# folds2predict = [1]\n\nuse_tta = True\n\nTH = 0.05\n\nDATASET_FOLDER = \"/kaggle/input/blood-vessel-segmentation\"\n\nis_test = not len(glob(os.path.join(DATASET_FOLDER, \"test\", \"*\", \"*\", \"*.tif\"))) == 6\n# is_test = not len(glob(os.path.join(DATASET_FOLDER, \"train\", \"*\", \"*\", \"*.tif\"))) == 6","metadata":{"execution":{"iopub.status.busy":"2024-01-24T10:27:03.158768Z","iopub.execute_input":"2024-01-24T10:27:03.159884Z","iopub.status.idle":"2024-01-24T10:27:03.182830Z","shell.execute_reply.started":"2024-01-24T10:27:03.159843Z","shell.execute_reply":"2024-01-24T10:27:03.181915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rename_keys(original_dict, pattern):\n    new_dict = {}\n    \n    for old_key, value in original_dict.items():\n        new_key = re.sub(pattern, '', old_key)\n        \n        new_dict[new_key] = value\n    \n    return new_dict\n\n\ndef rle_encode(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    rle = ' '.join(str(x) for x in runs)\n    if rle=='':\n        rle = '1 0'\n    return rle\n\n\ndef to_device(x: torch.Tensor, cuda_id: int = 0) -> torch.Tensor:\n    return x.cuda(cuda_id) if torch.cuda.is_available() else x\n\n\ndef load_jit_model(model_path: str, cuda_id: int = 0) -> torch.nn.Module:\n    model = torch.jit.load(\n        model_path,\n        map_location=f\"cuda:{cuda_id}\" if torch.cuda.is_available() else \"cpu\",\n    )\n    return model","metadata":{"execution":{"iopub.status.busy":"2024-01-24T10:27:03.373158Z","iopub.execute_input":"2024-01-24T10:27:03.373519Z","iopub.status.idle":"2024-01-24T10:27:03.382890Z","shell.execute_reply.started":"2024-01-24T10:27:03.373490Z","shell.execute_reply":"2024-01-24T10:27:03.381965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class BuildDataset(torch.utils.data.Dataset):\n    def __init__(self, dataset, in_channels=3, is_test=False):\n        self.window = in_channels // 2\n        self.is_test = is_test\n        self.ids = []\n\n        self.data_tensor = self.load_volume(dataset)\n        self.shape_orig = self.data_tensor.shape\n\n        padding = (\n            (self.window, self.window),\n        ) * self.data_tensor.ndim\n\n        self.padding = tuple(\n            (max(0, before), max(0, after)) for (before, after) in padding\n        )\n        self.data_tensor = np.pad(\n            self.data_tensor, padding, mode=\"constant\", constant_values=0\n        )\n\n    def __len__(self):\n        return sum(self.shape_orig) if self.is_test else self.shape_orig[0]\n\n    def normilize(self, image):\n        image = (image - self.xmin) / (\n                self.xmax - self.xmin)\n        image = np.clip(image, 0, 1)\n        return image.astype(np.float32)\n    \n    @staticmethod\n    def norm_by_percentile(volume, low=10, high=99.8):\n        xmin = np.percentile(volume, low)\n        print(xmin)\n        xmax = np.max([np.percentile(volume, high), 1])\n        print(xmax)\n        return xmin, xmax\n\n    def load_volume(self, dataset):\n        path = os.path.join(dataset, \"images\", \"*.tif\")\n        dataset = sorted(glob(path))\n        for p_img in tqdm(dataset):\n            path_ = p_img.split(os.path.sep)\n            slice_id, _ = os.path.splitext(path_[-1])\n            self.ids.append(f\"{path_[-3]}_{slice_id}\")\n            \n        volume = None\n\n        for z, path in enumerate(tqdm(dataset)):\n            image = cv2.imread(path, cv2.IMREAD_ANYDEPTH)\n            image = np.array(image, dtype=np.uint16)\n            if volume is None:\n                volume = np.zeros((len(dataset), *image.shape[-2:]), dtype=np.uint16)\n            volume[z] = image\n        self.xmin, self.xmax = self.norm_by_percentile(volume)\n        return volume\n\n    def __getitem__(self, idx):\n        # Determine which axis to sample from based on the index\n        if idx < self.shape_orig[0]:\n            idx = idx + self.window\n            slice_data = self.normilize(\n                self.data_tensor[\n                    idx - self.window : 1 + idx + self.window, :, :\n                ].transpose(1, 2, 0)[self.window:-self.window, self.window:-self.window, :]\n            )\n            axis = \"X\"\n            idx -= self.window\n        elif idx < self.shape_orig[0] + self.shape_orig[1]:\n            idx -= (self.shape_orig[0] - self.window)\n            slice_data = self.normilize(\n                self.data_tensor[\n                    :, idx - self.window : 1 + idx + self.window, :\n                ].transpose(0, 2, 1)[self.window:-self.window, self.window:-self.window, :]\n            )\n            axis = \"Y\"\n            idx -= self.window\n\n            \n        else:\n            idx -= (\n                self.shape_orig[0]\n                + self.shape_orig[1]\n                - self.window\n            ) \n            \n            slice_data = self.normilize(\n                self.data_tensor[\n                    :, :, idx - self.window : 1 + idx + self.window\n                ][self.window:-self.window, self.window:-self.window, :]\n            )\n            axis = \"Z\"\n            idx -= self.window\n\n        slice_data = torch.tensor(slice_data.transpose(2, 0, 1))\n\n        return {\n            \"slice\": slice_data,\n            \"slice_index\": idx,\n            \"axis\": axis\n        }","metadata":{"execution":{"iopub.status.busy":"2024-01-24T10:27:03.859655Z","iopub.execute_input":"2024-01-24T10:27:03.860017Z","iopub.status.idle":"2024-01-24T10:27:03.879553Z","shell.execute_reply.started":"2024-01-24T10:27:03.859989Z","shell.execute_reply":"2024-01-24T10:27:03.878628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_dataset = BuildDataset(sorted(glob(f\"{DATASET_FOLDER}/test/*\"))[-1], is_test=is_test, in_channels=5) # TODO: refactor this\n","metadata":{"execution":{"iopub.status.busy":"2024-01-24T10:27:04.205393Z","iopub.execute_input":"2024-01-24T10:27:04.206144Z","iopub.status.idle":"2024-01-24T10:27:04.209962Z","shell.execute_reply.started":"2024-01-24T10:27:04.206108Z","shell.execute_reply":"2024-01-24T10:27:04.209046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# len(test_dataset)","metadata":{"execution":{"iopub.status.busy":"2024-01-24T10:27:04.458413Z","iopub.execute_input":"2024-01-24T10:27:04.458992Z","iopub.status.idle":"2024-01-24T10:27:04.462758Z","shell.execute_reply.started":"2024-01-24T10:27:04.458960Z","shell.execute_reply":"2024-01-24T10:27:04.461895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_dataset[3]","metadata":{"execution":{"iopub.status.busy":"2024-01-24T10:27:04.730986Z","iopub.execute_input":"2024-01-24T10:27:04.731700Z","iopub.status.idle":"2024-01-24T10:27:04.735388Z","shell.execute_reply.started":"2024-01-24T10:27:04.731670Z","shell.execute_reply":"2024-01-24T10:27:04.734386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"def find_highest_score_filename(file_list):\n    highest_score = float('-inf')\n    highest_score_filename = None\n\n    for filename in file_list:\n        # Extract the score from the filename using regular expression\n        match = re.search(r'dice_(\\d+\\.\\d+)', filename)\n        if match:\n            current_score = float(match.group(1))\n            if current_score > highest_score:\n                highest_score = current_score\n                highest_score_filename = filename\n\n    return highest_score_filename","metadata":{"execution":{"iopub.status.busy":"2024-01-24T10:27:05.365995Z","iopub.execute_input":"2024-01-24T10:27:05.366348Z","iopub.status.idle":"2024-01-24T10:27:05.372066Z","shell.execute_reply.started":"2024-01-24T10:27:05.366320Z","shell.execute_reply":"2024-01-24T10:27:05.371182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"in_chans = []\n\nfor model_config in tqdm(predict_on):\n#     for fold in range(3):\n    for fold in folds2predict:\n        try:\n            if use_top_only:\n                model_path = sorted(glob(f\"/kaggle/input/{model_config[6]}/{model_config[3]}/{fold}/checkpoints/epoch*.ckpt\"))[-1]\n                print(f\"use_top_only, loading: {model_path}\")\n                state_dict = rename_keys(torch.load(model_path, map_location=\"cpu\")[\"state_dict\"], \"net.\")\n                model = to_device(smp.create_model(arch=model_config[0], encoder_name=model_config[1], in_channels=model_config[2], encoder_weights=None, decoder_attention_type=model_config[5]))\n                model.load_state_dict(state_dict)\n                model.eval()\n\n                if use_tta:\n                    tta_models.append(tta.SegmentationTTAWrapper(model, tta.aliases.hflip_transform(), merge_mode='mean')) #flip_transform d4_transform\n                else:\n                    tta_models.append(model)\n\n                weights.append(model_config[4])\n                in_chans.append(model_config[2])\n                print(weights)\n\n            elif use_best:\n                model_path = find_highest_score_filename(sorted(glob(f\"/kaggle/input/{model_config[6]}/{model_config[3]}/{fold}/checkpoints/epoch*.ckpt\")))\n                print(f\"use_best, loading: {model_path}\")\n                state_dict = rename_keys(torch.load(model_path, map_location=\"cpu\")[\"state_dict\"], \"net.\")\n                model = to_device(smp.create_model(arch=model_config[0], encoder_name=model_config[1], in_channels=model_config[2], encoder_weights=None, decoder_attention_type=model_config[5]))\n                model.load_state_dict(state_dict)\n                model.eval()\n\n                if use_tta:\n                    tta_models.append(tta.SegmentationTTAWrapper(model, tta.aliases.hflip_transform(), merge_mode='mean')) #flip_transform d4_transform\n                else:\n                    tta_models.append(model)\n\n                weights.append(model_config[4])\n                in_chans.append(model_config[2])\n\n            else:\n                for model_path in sorted(glob(f\"/kaggle/input/{model_config[6]}/{model_config[3]}/{fold}/checkpoints/*.ckpt\")):\n                    print(f\"use all, loading: {model_path}\")\n                    state_dict = rename_keys(torch.load(model_path, map_location=\"cpu\")[\"state_dict\"], \"net.\")\n                    model = to_device(smp.create_model(arch=model_config[0], encoder_name=model_config[1], in_channels=model_config[2], encoder_weights=None, decoder_attention_type=model_config[5]))\n                    model.load_state_dict(state_dict)\n                    model.eval()\n    #                 tta_models.append(tta.SegmentationTTAWrapper(model, tta.aliases.d4_transform(), merge_mode='mean'))\n                    if use_tta:\n                        tta_models.append(tta.SegmentationTTAWrapper(model, tta.aliases.hflip_transform(), merge_mode='mean')) #flip_transform d4_transform\n                    else:\n                        tta_models.append(model)\n                    weights.append(model_config[4])\n                    in_chans.append(model_config[2])\n        except:\n            pass","metadata":{"execution":{"iopub.status.busy":"2024-01-24T10:27:05.782874Z","iopub.execute_input":"2024-01-24T10:27:05.783285Z","iopub.status.idle":"2024-01-24T10:27:12.774548Z","shell.execute_reply.started":"2024-01-24T10:27:05.783258Z","shell.execute_reply":"2024-01-24T10:27:12.773686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"datasets = sorted(glob(f\"{DATASET_FOLDER}/test/*\"))[::-1]","metadata":{"execution":{"iopub.status.busy":"2024-01-09T09:20:20.721235Z","iopub.execute_input":"2024-01-09T09:20:20.721607Z","iopub.status.idle":"2024-01-09T09:20:20.726755Z","shell.execute_reply.started":"2024-01-09T09:20:20.721559Z","shell.execute_reply":"2024-01-09T09:20:20.725773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rles, ids = [], []\nwith torch.no_grad():\n    for dataset in datasets:\n#         test_dataset[2][\"slice\"][1,...].unsqueeze(0).shape\n        test_dataset = BuildDataset(dataset, is_test=is_test, in_channels=3) # TODO: refactor this\n        test_loader = DataLoader(test_dataset, batch_size=1, num_workers=4, shuffle=False, pin_memory=False)\n\n        y_preds = np.zeros(test_dataset.shape_orig, dtype=np.half)\n        ids += test_dataset.ids\n\n        pbar = tqdm(enumerate(test_loader), total=len(test_loader), desc=f'Inference {dataset}')\n        for step, batch in pbar:\n            images = to_device(batch[\"slice\"])\n        \n#             padder = ImagePadder(images)\n            \n#             padded_images = padder.pad()\n            \n            axis = batch[\"axis\"][0]\n            idx = batch[\"slice_index\"].numpy()[0]\n\n            preds = 0\n            for tta_model, weight, in_chan in zip(tta_models, weights, in_chans):\n#                 preds += padder.unpad(tta_model(padded_images))\n                preds += weight * monai.inferers.sliding_window_inference(\n                    inputs=images if in_chan != 1 else images[:, 1,...].unsqueeze(0), # TODO: Refactor this\n                    predictor=tta_model,\n                    sw_batch_size=4,\n                    roi_size=(512, 512),\n                    overlap=0.25,\n                    padding_mode=\"reflect\",\n                    mode=\"gaussian\",\n                    sw_device=\"cuda\",\n                    device=\"cuda\",\n                    progress=False,\n                )\n            if axis == \"X\":\n#                 print(preds.shape)\n                y_preds[idx, :, :] += ((preds / sum(weights)).squeeze().sigmoid().cpu().numpy() / 3.).astype(np.half)\n            elif axis == \"Y\":\n                y_preds[:, idx, :] += ((preds / sum(weights)).squeeze().sigmoid().cpu().numpy() / 3.).astype(np.half)\n            elif axis == \"Z\":\n                y_preds[:, :, idx] += ((preds / sum(weights)).squeeze().sigmoid().cpu().numpy() / 3.).astype(np.half)\n        \n       # y_preds = cc3d.dust(\n         #       (y_preds > TH).astype(np.uint8),\n             #   connectivity=18,\n            #    threshold=-1,\n           #     in_place=False\n         #   )\n        \n        for pred in y_preds:\n            rles.append(rle_encode((pred > TH).astype(np.uint8)))\n            # rles.append(rle_encode((pred)))\n\n        del test_dataset, test_loader, y_preds\n        gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-01-09T09:20:21.111127Z","iopub.execute_input":"2024-01-09T09:20:21.111537Z","iopub.status.idle":"2024-01-09T09:20:45.200090Z","shell.execute_reply.started":"2024-01-09T09:20:21.111503Z","shell.execute_reply":"2024-01-09T09:20:45.199025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.DataFrame.from_dict({\n    \"id\": ids,\n    \"rle\": rles\n})\nsubmission.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2024-01-09T09:20:45.202768Z","iopub.execute_input":"2024-01-09T09:20:45.203082Z","iopub.status.idle":"2024-01-09T09:20:45.215402Z","shell.execute_reply.started":"2024-01-09T09:20:45.203052Z","shell.execute_reply":"2024-01-09T09:20:45.214490Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2024-01-09T09:20:45.217277Z","iopub.execute_input":"2024-01-09T09:20:45.217559Z","iopub.status.idle":"2024-01-09T09:20:45.235865Z","shell.execute_reply.started":"2024-01-09T09:20:45.217535Z","shell.execute_reply":"2024-01-09T09:20:45.234958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}