{"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":"nvidiaTeslaT4","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":7445815,"sourceType":"datasetVersion","datasetId":4333929},{"sourceId":7526908,"sourceType":"datasetVersion","datasetId":4079395},{"sourceId":7529557,"sourceType":"datasetVersion","datasetId":4327897},{"sourceId":7530527,"sourceType":"datasetVersion","datasetId":4308668},{"sourceId":7544279,"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-02-03T08:22:47.471645Z","iopub.execute_input":"2024-02-03T08:22:47.471881Z","iopub.status.idle":"2024-02-03T08:23:20.033791Z","shell.execute_reply.started":"2024-02-03T08:22:47.471858Z","shell.execute_reply":"2024-02-03T08:23:20.032578Z"},"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-02-03T08:23:20.035544Z","iopub.execute_input":"2024-02-03T08:23:20.035862Z","iopub.status.idle":"2024-02-03T08:23:32.379789Z","shell.execute_reply.started":"2024-02-03T08:23:20.035835Z","shell.execute_reply":"2024-02-03T08:23:32.378514Z"},"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-02-03T08:23:32.381385Z","iopub.execute_input":"2024-02-03T08:23:32.381705Z","iopub.status.idle":"2024-02-03T08:24:36.095028Z","shell.execute_reply.started":"2024-02-03T08:23:32.381677Z","shell.execute_reply":"2024-02-03T08:24:36.093838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q /kaggle/input/tensordict/tensordict-0.2.1-cp310-cp310-manylinux1_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:24:36.098113Z","iopub.execute_input":"2024-02-03T08:24:36.099124Z","iopub.status.idle":"2024-02-03T08:25:08.228508Z","shell.execute_reply.started":"2024-02-03T08:24:36.099082Z","shell.execute_reply":"2024-02-03T08:25:08.227046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append(\"../input/ttach-kaggle/\")","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:25:08.230269Z","iopub.execute_input":"2024-02-03T08:25:08.230654Z","iopub.status.idle":"2024-02-03T08:25:08.236618Z","shell.execute_reply.started":"2024-02-03T08:25:08.230624Z","shell.execute_reply":"2024-02-03T08:25:08.235588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\n\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 tensordict\n\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nimport gc\nimport monai\nimport ttach as tta\nfrom typing import Union, Dict, Tuple\nfrom scipy.ndimage import binary_dilation, binary_erosion, binary_fill_holes, binary_opening, binary_closing\n\n\nimport re","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:25:08.237756Z","iopub.execute_input":"2024-02-03T08:25:08.238010Z","iopub.status.idle":"2024-02-03T08:25:52.422047Z","shell.execute_reply.started":"2024-02-03T08:25:08.237988Z","shell.execute_reply":"2024-02-03T08:25:52.421284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATASET_FOLDER = \"/kaggle/input/blood-vessel-segmentation\"\n\nis_test = not len(glob(os.path.join(DATASET_FOLDER, \"test\", \"*\", \"*\", \"*.tif\"))) == 6\nif is_test:\n    datasets = sorted(glob(f\"{DATASET_FOLDER}/test/*\"))[::-1]\nelse:\n    datasets = sorted(glob(f\"{DATASET_FOLDER}/train/kidney_2\"))\n\nprint(len(datasets))\n\nos.makedirs(\"preds_3d\")\nos.makedirs(\"kidney_masks\")","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:25:52.423249Z","iopub.execute_input":"2024-02-03T08:25:52.423914Z","iopub.status.idle":"2024-02-03T08:25:52.444833Z","shell.execute_reply.started":"2024-02-03T08:25:52.423887Z","shell.execute_reply":"2024-02-03T08:25:52.443806Z"},"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":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-02-03T08:25:52.446064Z","iopub.execute_input":"2024-02-03T08:25:52.446964Z","iopub.status.idle":"2024-02-03T08:25:52.465814Z","shell.execute_reply.started":"2024-02-03T08:25:52.446924Z","shell.execute_reply":"2024-02-03T08:25:52.465071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predict_on = [\n    [\n        \"Unet3d\",\n        192,\n        \"baseline_3d_unet_192_bs4_d4_scaled_pseudo_0.1_random\",\n        1.0,\n    ],\n]\n\n\nclass BuildDataset:\n    def __init__(self, dataset: str, is_test: bool = True):\n        self.ids = []\n        self.is_test = is_test\n\n        self.xmin, self.xmax = 0, 0\n\n        self.data_tensor = self.load_volume(dataset)\n        self.shape_orig = self.data_tensor.shape\n        \n    def normilize(self, image: np.ndarray) -> np.ndarray:\n        if image.dtype != np.half:\n            image = image.astype(np.half, copy=False)\n            \n        image -= self.xmin\n        image /= (self.xmax - self.xmin)\n        \n        np.clip(image, 0, 1, out=image)\n        return image\n    \n    @staticmethod\n    def norm_by_percentile(\n        volume: np.ndarray, low: float = 10, high: float = 99.8\n    ) -> Tuple:\n        xmin = np.percentile(volume, low)\n        print(xmin)\n        xmax = np.max([np.percentile(volume, high), 1])\n        print(xmax)\n        return int(xmin), int(xmax)\n\n    def load_volume(self, dataset: str) -> np.ndarray:\n        path = os.path.join(dataset, \"images\", \"*.tif\")\n        \n        dataset = sorted(glob(path)) if self.is_test else sorted(glob(path))[:192]\n\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).astype(np.half, copy=False)\n            \n            if volume is None:\n                volume = np.zeros((len(dataset), *image.shape[-2:]), dtype=np.float16)\n            volume[z, :, :] = image\n            \n        self.xmin, self.xmax = self.norm_by_percentile(volume)\n        return volume\n    \n    \nclass ModelWrapper(torch.nn.Module):\n    def __init__(self, base_model):\n        super(ModelWrapper, self).__init__()\n        self.base_model = base_model\n\n    def forward(self, x):\n        return torch.sigmoid(self.base_model(x)).half()","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:25:52.467819Z","iopub.execute_input":"2024-02-03T08:25:52.468120Z","iopub.status.idle":"2024-02-03T08:25:52.483424Z","shell.execute_reply.started":"2024-02-03T08:25:52.468094Z","shell.execute_reply":"2024-02-03T08:25:52.482581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tta_models = []\nweights = []\n\nfolds2predict = [0, 1]\n\nfor model_config in tqdm(predict_on):\n    for fold in folds2predict:\n        model_path = sorted(\n            glob(\n                f\"/kaggle/input/senet-3d-models/{model_config[2]}/{fold}/checkpoints/epoch*.ckpt\"\n            )\n        )[-1]\n        print(model_path)\n        state_dict = rename_keys(\n            torch.load(model_path, map_location=\"cpu\")[\"state_dict\"], \"net.\"\n        )\n        model_base = to_device(\n            monai.networks.nets.DynUNet(spatial_dims=3, in_channels=1, out_channels=1, kernel_size=[ [ 3, 3, 3 ], [ 3, 3, 3 ], [ 3, 3, 3 ], [ 3, 3, 3 ], [ 3, 3, 3 ], [ 3, 3, 3 ] ], strides=[ [ 1, 1, 1 ], [ 2, 2, 2 ], [ 2, 2, 2 ], [ 2, 2, 2 ], [ 2, 2, 2 ], [ 2, 2, 2 ] ], upsample_kernel_size=[[ 2, 2, 2 ], [ 2, 2, 2 ], [ 2, 2, 2 ], [ 2, 2, 2 ], [ 2, 2, 2 ]], dropout=0.2)\n        )\n        model_base.load_state_dict(state_dict)\n        model = ModelWrapper(model_base)\n\n        model.eval()\n        model = torch.nn.DataParallel(model)\n        \n        if is_test:\n            tta_models.append(\n                tta.SegmentationTTAWrapper(\n                    model.half(), tta.aliases.d4_transform(), merge_mode=\"mean\"\n                )\n            )\n        else:\n            tta_models.append(model.half())\n\n        weights.append(model_config[-1])","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:25:52.488387Z","iopub.execute_input":"2024-02-03T08:25:52.488697Z","iopub.status.idle":"2024-02-03T08:25:58.434825Z","shell.execute_reply.started":"2024-02-03T08:25:52.488675Z","shell.execute_reply":"2024-02-03T08:25:58.433895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rles, ids = [], []\nwith torch.no_grad():\n    for dataset in datasets:\n        folder = dataset.split(\"/\")[-1]\n\n        test_dataset = BuildDataset(dataset, is_test=is_test)\n\n        ids += test_dataset.ids\n        \n        preds = 0\n        input_tensor = tensordict.MemmapTensor.from_tensor(torch.from_numpy(test_dataset.normilize(test_dataset.data_tensor.astype(np.half))).unsqueeze(0).unsqueeze(0))\n        for tta_model, weight in zip(tta_models, weights):\n            preds += weight * monai.inferers.sliding_window_inference(\n#                 inputs=torch.from_numpy(test_dataset.normilize(test_dataset.data_tensor.astype(np.half))).unsqueeze(0).unsqueeze(0), # if is_test else torch.rand(1, 1, 512, 512, 512),\n                inputs=input_tensor, # if is_test else torch.rand(1, 1, 512, 512, 512),\n                predictor=tta_model,\n                sw_batch_size=2,\n                roi_size=(256, 256, 256),\n                overlap=0.25,\n                padding_mode=\"reflect\",\n                mode=\"gaussian\",\n                sw_device=\"cuda\",\n                device=\"cpu\",\n                progress=True,\n            ).squeeze().cpu().numpy().astype(np.half) / sum(weights)\n\n        for idx, pred in enumerate(preds):\n            cv2.imwrite(f\"preds_3d/{test_dataset.ids[idx]}.png\", (255*pred).astype(np.uint8))\n\n            \n        if is_test:\n            del input_tensor, test_dataset, preds\n            gc.collect()\n            torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:25:58.435905Z","iopub.execute_input":"2024-02-03T08:25:58.436163Z","iopub.status.idle":"2024-02-03T08:28:11.326675Z","shell.execute_reply.started":"2024-02-03T08:25:58.436139Z","shell.execute_reply":"2024-02-03T08:28:11.325847Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# if is_test:\ndel tta_models\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:28:11.327838Z","iopub.execute_input":"2024-02-03T08:28:11.328120Z","iopub.status.idle":"2024-02-03T08:28:11.737005Z","shell.execute_reply.started":"2024-02-03T08:28:11.328095Z","shell.execute_reply":"2024-02-03T08:28:11.736264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predict_on = [\n     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 1, \"UnetPlusPlus_tu-tf_efficientnet_b5_BoundaryDoULoss_size_1_512_bs32_hard_pseudo_v2\", 3., \"scse\", \"senet-models\", 800],  #0878 new sampling + cutmix\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., \"scse\", \"senet-hoa\", 800],  #0873\n#      [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 1, \"UnetPlusPlus_tu-tf_efficientnet_b5_size_1_512_bs32_hard_pseudo\", 1., \"scse\", \"senet-hoa\", 800],  #0873\n#     [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 1, \"UnetPlusPlus_tu-tf_efficientnet_b5_size_1_512_bs32_hard_pseudo_v2_cutmix\", 1., \"scse\", \"senet-models\", 800], #0.873 #0 old sampling + only d4, severe zoom and gamma + brightness \n #    [\"UnetPlusPlus\", \"tu-tf_efficientnet_b5\", 1, \"UnetPlusPlus_tu-tf_efficientnet_b5_size_1_1024_bs32_hard_pseudo_v2_cutmix\", 1., \"scse\", \"senet-models\", 1280], #0.875 #0 old sampling + only d4, severe zoom and gamma + brightness \n]\n\n\ntta_models = []\nweights = []\n\nfolds2predict = [0, 1]\n\nuse_tta = True\n\nTH3d = 0.5 #()\nTH2d = 0.05\n\n","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:28:11.738232Z","iopub.execute_input":"2024-02-03T08:28:11.738590Z","iopub.status.idle":"2024-02-03T08:28:11.744968Z","shell.execute_reply.started":"2024-02-03T08:28:11.738558Z","shell.execute_reply":"2024-02-03T08:28:11.744018Z"},"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 generate_kidney_mask(\n#         normalized_image: np.array, \n#         kidney_tresh: float = 0.6, \n#         opening_factor: int = 20,\n#         dialation_factor: int = 30\n#     ) -> np.ndarray:\n#         kidney_mask = normalized_image >= kidney_tresh\n#         kidney_mask = binary_fill_holes(kidney_mask)\n#         kidney_mask = binary_opening(kidney_mask, iterations=opening_factor)\n#         kidney_mask = binary_dilation(kidney_mask, iterations=dialation_factor)\n#         kidney_mask = binary_fill_holes(kidney_mask)\n#         return (255. * kidney_mask).astype(np.uint8)\n    \n#     @staticmethod\n#     def generate_kidney_mask(\n#         normalized_image: np.array, \n#         kidney_tresh: float = 0.6, \n#         opening_factor: int = 30,\n#     ) -> np.ndarray:\n#         mask = (255. * (normalized_image > kidney_tresh)).astype(np.uint8)\n\n#         edged = cv2.Canny(mask, 30, 200)\n#         contours, hierarchy = cv2.findContours(edged,  \n#             cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE) \n\n#         contour_img = np.zeros_like(mask)\n\n#         cv2.drawContours(contour_img, contours, -1, (1), thickness=10)\n\n#         kidney_mask = binary_fill_holes(contour_img > 0.5)\n#         kidney_mask = binary_opening(kidney_mask, iterations=opening_factor)\n#         return (255. * kidney_mask).astype(np.uint8)\n    \n    @staticmethod\n    def generate_kidney_mask(\n        normalized_image: np.array, \n        kidney_tresh: float = 0.5, \n        opening_factor: int = 45,\n        dialation_factor: int = 30\n    ) -> np.ndarray:\n        mask = (255*normalized_image).astype(np.uint8)\n\n        edged = cv2.Canny(mask, 150, 200, L2gradient=True) \n\n        contours, hierarchy = cv2.findContours(edged,  \n            cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE) \n\n\n        contour_img = np.zeros_like(mask)\n        cv2.drawContours(contour_img, contours, -1, (1), thickness=2)\n        kidney_mask = contour_img > 0.5\n\n        kidney_mask = binary_dilation(kidney_mask, iterations=dialation_factor)\n\n        kidney_mask = binary_fill_holes(kidney_mask)\n        kidney_mask = binary_opening(kidney_mask, iterations=opening_factor)\n        kidney_mask = binary_erosion(kidney_mask, iterations=dialation_factor)\n        return (255. * kidney_mask).astype(np.uint8)\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)) if self.is_test else sorted(glob(path))[:192]\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            \n            kidney_mask = self.generate_kidney_mask(slice_data[..., 1])\n            \n            axis = \"X\"\n            idx -= self.window \n            cv2.imwrite(f\"kidney_masks/{self.ids[idx]}.png\", kidney_mask)\n\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.half(),\n            \"slice_index\": idx,\n            \"axis\": axis\n        }","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:46:14.386258Z","iopub.execute_input":"2024-02-03T08:46:14.386597Z","iopub.status.idle":"2024-02-03T08:46:14.411089Z","shell.execute_reply.started":"2024-02-03T08:46:14.386571Z","shell.execute_reply":"2024-02-03T08:46:14.410170Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"in_chans = []\nresolutions = []\n\nfor model_config in tqdm(predict_on):\n    for fold in folds2predict:\n        try:\n            files = glob(f\"/kaggle/input/{model_config[6]}/{model_config[3]}/{fold}/checkpoints/last*.ckpt\")\n            if len(files) == 0:\n                files = glob(f\"/kaggle/input/{model_config[6]}/{model_config[3]}/{fold}/checkpoints/epoch*.ckpt\")\n            model_path = sorted(files)[-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            model = torch.nn.DataParallel(model)\n\n            if is_test and use_tta:\n                tta_models.append(tta.SegmentationTTAWrapper(model.half(), tta.aliases.d4_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            resolutions.append(model_config[7])\n            \n        except:\n            pass","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:46:15.057814Z","iopub.execute_input":"2024-02-03T08:46:15.058140Z","iopub.status.idle":"2024-02-03T08:46:17.028069Z","shell.execute_reply.started":"2024-02-03T08:46:15.058114Z","shell.execute_reply":"2024-02-03T08:46:17.027149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del state_dict","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:46:17.093895Z","iopub.execute_input":"2024-02-03T08:46:17.094173Z","iopub.status.idle":"2024-02-03T08:46:17.100301Z","shell.execute_reply.started":"2024-02-03T08:46:17.094149Z","shell.execute_reply":"2024-02-03T08:46:17.099579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"def merge_preds(mask1, mask2, kidney_mask):\n    binary_mask = (255 * (mask1 > TH2d)).astype(np.uint8)\n    edged = cv2.Canny(binary_mask, 30, 200)\n    contours, hierarchy = cv2.findContours(edged,  \n        cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE) \n    interest_mask = np.zeros_like(binary_mask)\n\n    if len(contours) > 0:\n        all_contours = np.vstack(contours[i] for i in range(len(contours)))\n        hull = cv2.convexHull(all_contours)\n        cv2.drawContours(interest_mask, [hull], -1, (1), thickness=cv2.FILLED)\n\n        interest_mask = cv2.dilate(interest_mask, np.ones((5, 5), np.uint8), iterations=5) \n#         return ((interest_mask * mask2) > 0.5).astype(np.uint8)  \n#         return ((kidney_mask * (interest_mask * mask2 + mask1)) > TH2d + TH3d).astype(np.uint8)   \n        return ((interest_mask * (mask2 + mask1)) > TH2d + TH3d).astype(np.uint8)   \n    else:\n        return (interest_mask * mask1 > TH2d).astype(np.uint8)        \n#         return (mask1 > TH2d).astype(np.uint8)","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:46:18.471076Z","iopub.execute_input":"2024-02-03T08:46:18.471436Z","iopub.status.idle":"2024-02-03T08:46:18.479921Z","shell.execute_reply.started":"2024-02-03T08:46:18.471410Z","shell.execute_reply":"2024-02-03T08:46:18.478961Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rles, ids = [], []\n\n\nwith torch.no_grad():\n    for dataset in datasets:\n        test_dataset = BuildDataset(dataset, is_test=is_test, in_channels=3)\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            axis = batch[\"axis\"][0]\n            idx = batch[\"slice_index\"].numpy()[0]\n\n            preds = 0\n            for tta_model, weight, in_chan, resolution in zip(tta_models, weights, in_chans, resolutions):\n                preds += weight * monai.inferers.sliding_window_inference(\n                    inputs=images.half() if in_chan != 1 else images[:, 1,...].unsqueeze(0).half(), # TODO: Refactor this\n                    predictor=tta_model.half(),\n                    sw_batch_size=8,\n                    roi_size=(resolution, resolution),\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                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        for idx, pred_2d in enumerate(y_preds):\n            pred_3d = cv2.imread(f\"preds_3d/{test_dataset.ids[idx]}.png\", 0) / 255.\n            kidney_mask = cv2.imread(f\"kidney_masks/{test_dataset.ids[idx]}.png\", 0) / 255.\n            \n            rles.append(rle_encode(merge_preds(pred_2d, pred_3d, kidney_mask)))\n            \n        if is_test:\n            del test_dataset, test_loader, y_preds\n            gc.collect()\n            torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:47:05.689111Z","iopub.execute_input":"2024-02-03T08:47:05.689858Z","iopub.status.idle":"2024-02-03T08:50:00.427345Z","shell.execute_reply.started":"2024-02-03T08:47:05.689826Z","shell.execute_reply":"2024-02-03T08:50:00.426186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# idx = 150\n\n# mask1 = y_preds[idx,:,:]\n# mask2 = cv2.imread(f\"preds_3d/{test_dataset.ids[idx]}.png\", 0) / 255.\n# kidney_mask = cv2.imread(f\"kidney_masks/{test_dataset.ids[idx]}.png\", 0) / 255.\n\n# binary_mask = (255 * (mask1 > TH2d)).astype(np.uint8)\n# edged = cv2.Canny(binary_mask, 30, 200)\n# contours, hierarchy = cv2.findContours(edged,  \n#     cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE) \n# interest_mask = np.zeros_like(binary_mask)\n\n# if len(contours) > 0:\n#     all_contours = np.vstack(contours[i] for i in range(len(contours)))\n#     hull = cv2.convexHull(all_contours)\n#     cv2.drawContours(interest_mask, [hull], -1, (1), thickness=cv2.FILLED)\n\n#     interest_mask = cv2.dilate(interest_mask, np.ones((5, 5), np.uint8), iterations=5) \n# #         return ((interest_mask * mask2) > 0.5).astype(np.uint8)  \n#     x = ((kidney_mask * (interest_mask * mask2 + mask1)) > TH2d + TH3d).astype(np.uint8)   \n# else:\n#     x = (kidney_mask * mask1 > TH2d).astype(np.uint8)  \n","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:50:00.429450Z","iopub.execute_input":"2024-02-03T08:50:00.429830Z","iopub.status.idle":"2024-02-03T08:50:00.556568Z","shell.execute_reply.started":"2024-02-03T08:50:00.429796Z","shell.execute_reply":"2024-02-03T08:50:00.555501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# kidney_mask = cv2.imread(f\"kidney_masks/{test_dataset.ids[idx]}.png\", 0)","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:50:00.558130Z","iopub.execute_input":"2024-02-03T08:50:00.558431Z","iopub.status.idle":"2024-02-03T08:50:00.569678Z","shell.execute_reply.started":"2024-02-03T08:50:00.558407Z","shell.execute_reply":"2024-02-03T08:50:00.568786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# data = test_dataset[150][\"slice\"].numpy()","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:50:00.572138Z","iopub.execute_input":"2024-02-03T08:50:00.572741Z","iopub.status.idle":"2024-02-03T08:50:00.843826Z","shell.execute_reply.started":"2024-02-03T08:50:00.572710Z","shell.execute_reply":"2024-02-03T08:50:00.842876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import matplotlib.pyplot as plt\n# plt.figure(figsize=(20, 10))\n# plt.subplot(2,2,1)\n# plt.imshow(x)\n# plt.subplot(2,2,2)\n# plt.imshow(mask1 > 0.05)\n# plt.subplot(2,2,3)\n# plt.imshow(mask2 > 0.5)\n# plt.subplot(2,2,4)\n# plt.imshow(kidney_mask)","metadata":{"execution":{"iopub.status.busy":"2024-02-03T08:54:30.704128Z","iopub.execute_input":"2024-02-03T08:54:30.704506Z","iopub.status.idle":"2024-02-03T08:54:32.040227Z","shell.execute_reply.started":"2024-02-03T08:54:30.704474Z","shell.execute_reply":"2024-02-03T08:54:32.039113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if is_test:\n    del tta_models, tta_model, batch, preds, images, model\n    gc.collect()\n    torch.cuda.empty_cache()\n","metadata":{"execution":{"iopub.status.busy":"2024-01-22T10:33:12.202646Z","iopub.execute_input":"2024-01-22T10:33:12.203002Z","iopub.status.idle":"2024-01-22T10:33:12.604875Z","shell.execute_reply.started":"2024-01-22T10:33:12.202968Z","shell.execute_reply":"2024-01-22T10:33:12.604117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -rf preds_3d kidney_masks","metadata":{},"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}