{"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":160190612,"sourceType":"kernelVersion"}],"dockerImageVersionId":30626,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\nimport csv\nimport sys\nimport albumentations as A\nimport cv2\nimport json\nimport random\nimport re\nimport matplotlib.pyplot as plt\nimport shutil\nimport tifffile\n# Install smp with resnet checkpoints\n!mkdir -p /root/.cache/torch/hub/checkpoints\n!cp /kaggle/input/install-segmentation-models-pytorch-0-3-3-ckpt/smp_ckpt/* /root/.cache/torch/hub/checkpoints/\n!python -m pip install --no-index --find-links=/kaggle/input/install-segmentation-models-pytorch-0-3-3-ckpt/smp_whl segmentation-models-pytorch\nimport segmentation_models_pytorch as smp\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nimport torchvision.transforms as transforms\n\nfrom abc import ABC, abstractmethod\nfrom typing import Any, List, Optional, Tuple\nfrom albumentations.pytorch import ToTensorV2\nfrom collections import OrderedDict, defaultdict\nfrom dataclasses import dataclass\nfrom PIL import Image\nfrom segmentation_models_pytorch.losses.dice import DiceLoss\nfrom segmentation_models_pytorch.losses.focal import FocalLoss\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# ref.: https://www.kaggle.com/stainsby/fast-tested-rle\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    if np.max(pixels) == np.min(pixels) == 0:\n        return \"1 0\"\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n \ndef rle_decode(mask_rle, shape):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (height,width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n    '''\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape)\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        #print(os.path.join(dirname, filename))\n        pass\n        \n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8cc998af-97ac-4148-8b66-4c396f19fe52","_cell_guid":"71499b90-892c-4290-8d9c-1c1dd45c8f11","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-05T00:00:50.339513Z","iopub.execute_input":"2024-02-05T00:00:50.340479Z","iopub.status.idle":"2024-02-05T00:02:44.497474Z","shell.execute_reply.started":"2024-02-05T00:00:50.340439Z","shell.execute_reply":"2024-02-05T00:02:44.496083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = (\n    \"cuda\"\n    if torch.cuda.is_available()\n    else \"cpu\"\n)\nprint(f\"Using {device} device\")","metadata":{"_uuid":"6a8cefc2-5147-4811-baa0-1b8ee9c389d5","_cell_guid":"5a7cad3f-de9f-4fa1-a685-d7c02d31cf9a","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-04T23:18:33.307123Z","iopub.execute_input":"2024-02-04T23:18:33.307626Z","iopub.status.idle":"2024-02-04T23:18:33.368648Z","shell.execute_reply.started":"2024-02-04T23:18:33.307599Z","shell.execute_reply":"2024-02-04T23:18:33.367752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!nvidia-smi","metadata":{"_uuid":"27ca6721-8699-4236-9b60-01d357b00be6","_cell_guid":"6449a15e-f211-4455-aef7-5337e1ceb148","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-04T23:18:33.370008Z","iopub.execute_input":"2024-02-04T23:18:33.370751Z","iopub.status.idle":"2024-02-04T23:18:34.405261Z","shell.execute_reply.started":"2024-02-04T23:18:33.370706Z","shell.execute_reply":"2024-02-04T23:18:34.403929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> ## Configuration","metadata":{"_uuid":"55d4b759-8b17-43aa-8124-730fc2263a26","_cell_guid":"f6bffb33-57df-436d-9ed5-4c6dd4999630","trusted":true}},{"cell_type":"code","source":"@dataclass\nclass ConfigParams:\n    # TRAIN AND AUGMENTATIONS\n    batch_size = 4\n    epochs = 50\n    patience = 5\n    train_augmentation = \"my_aug_v2b\"\n    \n    # OPTIMIZER AND LOSS\n    optimizer = \"ADAM\"\n    learning_rate = 0.001\n    loss_function = \"dice_loss\"\n    # loss_function = \"smp_focal_loss\"\n    val_metrics_to_log = [\"dice_score\", \"smp_dice_score\"]\n    val_metric_to_monitor = \"dice_score\"\n    \n    # FOCAL_LOSS\n    focal_loss_alpha = None\n    focal_loss_gamma = 2.0\n    \n    # TRAIN DEBUG\n    num_batches_train_loss_aggregation = 10\n    # num_batches_preds_train_visualization_period = 50\n    # num_batches_preds_val_visualization_period = 25\n\n    # DATASET\n    input_root_dir = \"/kaggle/input/blood-vessel-segmentation\"\n    output_dir =\"/kaggle/working\"\n    # train_dirs = [\"train/kidney_1_dense\"]  # Used for the challenge\n    train_dirs = [\"train/kidney_1_dense\", \"train/kidney_1_voi\", \"train/kidney_2\"]  # Further experiments\n    # train_dirs = [\"train/kidney_1_dense\", \"train/kidney_3_sparse\"]  # As final train already tuned\n    val_dirs = [\"train/kidney_3_sparse\"]\n    test_dirs = [\"test/kidney_5\", \"test/kidney_6\"]\n    # test_dirs = [\"train/kidney_3_sparse\"]  # To test TTA\n\n    # MODEL\n    model_name = \"unet_smp_se-resnext50\"\n    model_train_input_size = 512\n    model_input_channels = 1\n    model_smp_model = \"unet\"\n    model_smp_encoder = \"se_resnext50_32x4d\"\n    model_smp_encoder_weights = \"imagenet\"\n    # Not used for smp\n    model_batch_norm = True\n    model_dropout = True\n                          \n    # INFERENCE\n    threshold = 0.1  # 0.1 for dice loss, 0.4 for focal loss\n    tta_mode = \"5max\"\n    model_inference_input_size_width = 768  # 768 used for the challenge, 1024 is better on private dataset +0.03 with dice loss\n    model_inference_input_size_height = 864  # 864 used for the challenge, 1024 is better on private dataset +0.03 with dice loss\n    \nconfig = ConfigParams()","metadata":{"_uuid":"32fda894-3846-4c24-a3ee-3574c0d782db","_cell_guid":"73405eb3-dc41-46eb-bf42-fb5afc1a2b11","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-04T23:56:42.896360Z","iopub.execute_input":"2024-02-04T23:56:42.896749Z","iopub.status.idle":"2024-02-04T23:56:42.905581Z","shell.execute_reply.started":"2024-02-04T23:56:42.896719Z","shell.execute_reply":"2024-02-04T23:56:42.904572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Set seed for reproducibility\nFrom: https://www.kaggle.com/code/vinayaktiwari28/easy-to-understand-clean-baseline-code-train","metadata":{}},{"cell_type":"code","source":"seed = 23\nnp.random.seed(seed)\nrandom.seed(seed)\ntorch.manual_seed(seed)\ntorch.cuda.manual_seed(seed)\n# When running on the CuDNN backend, two further options must be set\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False\n# Set a fixed value for the hash seed\nos.environ[\"PYTHONHASHSEED\"] = str(seed)","metadata":{"execution":{"iopub.status.busy":"2024-02-04T23:55:41.194830Z","iopub.execute_input":"2024-02-04T23:55:41.195209Z","iopub.status.idle":"2024-02-04T23:55:41.201299Z","shell.execute_reply.started":"2024-02-04T23:55:41.195171Z","shell.execute_reply":"2024-02-04T23:55:41.200308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model definition with preprocessing functions","metadata":{"_uuid":"de588d34-3707-45ba-931c-5a60baa50055","_cell_guid":"3561bac1-8003-4983-ac28-3758a8aeec0d","trusted":true}},{"cell_type":"code","source":"def preprocess_min_max(x: torch.Tensor):\n    x_max = x.max()\n    x_min = x.min()\n\n    assert x_max <= 1.0 and x_min >= 0.0\n\n    return (x - x_min) / (x_max - x_min)\n\n\ndef inverse_preprocess_min_max(\n    x_norm: torch.Tensor, new_min: float = 0.0, new_max: float = 1.0\n):\n    \"\"\"\n    Operation to restore an normalized input to the original range in [0.0, 1.0]\n    corresponding to [0, 255] as uint8. The operation with default values has no actions\n    and you need original max and min to restore the original image\n    \"\"\"\n    x = x_norm * (new_max - new_min) + new_min\n    x_max = x.max()\n    x_min = x.min()\n    # TODO: Check closeness to new_max and new_min\n    assert x_max <= new_max and x_min >= new_min\n    return x\n\n\ndef preprocess_mean_std_grayscale(\n    x: torch.Tensor, mean: float = 0.449, std: float = 0.226\n):\n    \"\"\"\n    Same preprocess_input of smp for resnext50 with imagenet weights, but adapted to grayscale\n    \"\"\"\n    x_max = x.max()\n    x_min = x.min()\n\n    assert x_max <= 1.0 and x_min >= 0.0\n\n    x = x - mean\n    x = x / std\n\n    return x\n\n\ndef inverse_preprocess_mean_str_grayscale(\n    x_norm: torch.Tensor, mean: float = 0.449, std: float = 0.226\n):\n    \"\"\"\n    Operation to restore a normalized input to the original range in [0.0, 1.0]\n    corresponding to [0, 255] as uint8\n    \"\"\"\n    x = x_norm * std\n    x = x + mean\n    x_max = x.max()\n    x_min = x.min()\n\n    assert x_max <= 1.0 and x_min >= 0.0\n\n    return x\n\nclass UnetAfolabi(nn.Module):\n    \"\"\"\n    Unet model from \"Fundus Images using Modified U-net Convolutional Neural Network\" (Afolabi, 2020)\n    \"\"\"\n    def __init__(self, batch_norm: bool = True, dropout: bool = True):\n        super(UnetAfolabi, self).__init__()\n        self.batch_norm = batch_norm\n        self.dropout = dropout\n        self.ds_block_1 = ConvBlock(\n            in_channels=1, out_channels=64, num_blocks=1, batch_norm=batch_norm\n        )\n        self.ds_block_2 = ConvBlock(\n            in_channels=64, out_channels=64, num_blocks=3, batch_norm=batch_norm\n        )\n        self.ds_block_3 = ConvBlock(\n            in_channels=64, out_channels=64, num_blocks=3, batch_norm=batch_norm\n        )\n        self.ds_block_4 = ConvBlock(\n            in_channels=64, out_channels=64, num_blocks=3, batch_norm=batch_norm\n        )\n        self.bottom = ConvBlock(\n            in_channels=64, out_channels=64, num_blocks=3, batch_norm=batch_norm\n        )\n        self.max_pooling = nn.MaxPool2d(kernel_size=2, stride=2)\n        self.upsampling_2x = nn.UpsamplingNearest2d(scale_factor=2.0)\n        self.us_block_4 = ConvBlock(\n            in_channels=128, out_channels=32, num_blocks=3, batch_norm=batch_norm\n        )\n        self.us_block_3 = ConvBlock(\n            in_channels=96, out_channels=32, num_blocks=3, batch_norm=batch_norm\n        )\n        self.us_block_2 = ConvBlock(\n            in_channels=96, out_channels=32, num_blocks=3, batch_norm=batch_norm\n        )\n        self.us_block_1 = ConvBlock(\n            in_channels=96, out_channels=32, num_blocks=3, batch_norm=batch_norm\n        )\n        self.last_conv = nn.Conv2d(\n            in_channels=32, out_channels=1, kernel_size=(1, 1), padding=0\n        )\n\n    def forward(self, x):\n        # Input: 1x512x512\n        x_1 = self.ds_block_1(x)  # 64x512x512\n        x_1_down = self.max_pooling(x_1)  # 64x256x256\n        x_2 = self.ds_block_2(x_1_down)  # 64x256x256\n        x_2_down = self.max_pooling(x_2)  # 64x128x128\n        x_3 = self.ds_block_3(x_2_down)  # 64x128x128\n        x_3_down = self.max_pooling(x_3)  # 64x64x64\n        x_4 = self.ds_block_4(x_3_down)  # 64x64x64\n        x_4_down = self.max_pooling(x_4)  # 64x32x32\n        x_bottom = self.bottom(x_4_down)  # 64x32x32\n        x_4_up = self.upsampling_2x(x_bottom)  # 64x64x64\n        x_4_cat = torch.cat([x_4, x_4_up], dim=1)  # 128x64x64\n        x_4_up_conv = self.us_block_4(x_4_cat)  # 32x64x64\n        x_3_up = self.upsampling_2x(x_4_up_conv)  # 32x128x128\n        x_3_cat = torch.cat([x_3, x_3_up], dim=1)  # 96x128x128\n        x_3_up_conv = self.us_block_3(x_3_cat)  # 32x128x128\n        x_2_up = self.upsampling_2x(x_3_up_conv)  # 32x256x256\n        x_2_cat = torch.cat([x_2, x_2_up], dim=1)  # 96x256x256\n        x_2_up_conv = self.us_block_2(x_2_cat)  # 32x256x256\n        x_1_up = self.upsampling_2x(x_2_up_conv)  # 32x512x512\n        x_1_cat = torch.cat([x_1, x_1_up], dim=1)  # 96x512x512\n        x_1_up_conv = self.us_block_1(x_1_cat)  # 32x512x512\n        if self.dropout:\n            x_1_dropout = nn.Dropout()(\n                x_1_up_conv\n            )  # Not explicitly mentioned in the paper\n            return self.last_conv(x_1_dropout)  # 1x512x512\n        else:\n            return self.last_conv(x_1_up_conv)  # 1x512x512\n\n\nclass ConvBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, num_blocks, batch_norm=True):\n        \"\"\"\n        Args:\n            in_channels: number of input channels only for the first  convolution layer\n            out_channels: number of the output channels of all the convolution layers\n            which is also the in_channels for the next step\n            num_blocks: number of conv+leaky_relu+bn blocks\n            batch_norm: apply batch normalization or not\n        \"\"\"\n        super(ConvBlock, self).__init__()\n        self.blocks = nn.ModuleList()\n        for i in range(num_blocks):\n            input_channels = in_channels if i == 0 else out_channels\n            if batch_norm:\n                self.blocks.append(\n                    nn.Sequential(\n                        nn.Conv2d(\n                            in_channels=input_channels,\n                            out_channels=out_channels,\n                            kernel_size=(3, 3),\n                            padding=1,\n                        ),\n                        nn.LeakyReLU(negative_slope=0.018),\n                        nn.BatchNorm2d(out_channels),\n                    )\n                )\n            else:\n                self.blocks.append(\n                    nn.Sequential(\n                        nn.Conv2d(\n                            in_channels=input_channels,\n                            out_channels=out_channels,\n                            kernel_size=(3, 3),\n                            padding=1,\n                        ),\n                        nn.LeakyReLU(negative_slope=0.018),\n                    )\n                )\n\n    def forward(self, x):\n        for block in self.blocks:\n            x = block(x)\n        return x\n\ndef init_smp_model(config: ConfigParams) -> nn.Module:\n    if config.model_smp_model != \"unet\":\n        raise Exception(\"Only unet model is supported\")\n    model = smp.Unet(\n        encoder_name=config.model_smp_encoder,  # choose encoder, e.g. mobilenet_v2 or efficientnet-b7\n        encoder_weights=config.model_smp_encoder_weights,  # use `imagenet` pre-trained weights for encoder initialization\n        in_channels=config.model_input_channels,  # model input channels (1 for gray-scale images, 3 for RGB, etc.)\n        classes=1,  # model output channels (number of classes in your dataset)\n    )\n    return model\n\ndef init_model(config: ConfigParams) -> Tuple[nn.Module, Any, Any]:\n    if config.model_name == \"unet_afolabi\":\n        model = UnetAfolabi(\n            batch_norm=config.model_batch_norm, dropout=config.model_dropout\n        )\n        preprocessing_fn = preprocess_min_max\n        inverse_preprocessing_fn = inverse_preprocess_min_max\n    elif config.model_smp_model is not None:\n        model = init_smp_model(config)\n        preprocessing_fn = preprocess_mean_std_grayscale\n        inverse_preprocessing_fn = inverse_preprocess_mean_str_grayscale\n    else:\n        raise Exception(\"Unable to initialize the model, please check the config\")\n    return model, preprocessing_fn, inverse_preprocessing_fn","metadata":{"_uuid":"f369116f-6cce-4bff-ada4-f4083e14ecd8","_cell_guid":"f1b125c6-5e09-4e7e-9db2-ed588a6f6860","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-04T23:55:41.963471Z","iopub.execute_input":"2024-02-04T23:55:41.963924Z","iopub.status.idle":"2024-02-04T23:55:41.999517Z","shell.execute_reply.started":"2024-02-04T23:55:41.963887Z","shell.execute_reply":"2024-02-04T23:55:41.998551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training code","metadata":{"_uuid":"0a762cb2-cfb6-4952-a7b0-8fbb47dd4e85","_cell_guid":"cd79f219-08a2-4c1f-9efa-2317db51f5c7","trusted":true}},{"cell_type":"code","source":"model, preprocess_function, _ = init_model(config)\nmodel.to(device)\ntotal_parameters = 0\nfor parameter in model.parameters():\n    if parameter.requires_grad:\n        total_parameters += parameter.numel()\nprint(f\"Trainable parameters: {total_parameters}\")\nprint(f\"Estimated size: {(total_parameters * 4 / 1024 / 1024):.2f} MB\")\nprint(model)","metadata":{"_uuid":"dc343dfb-5704-419b-9d89-b0cd0307615e","_cell_guid":"43c1e8e2-96e0-4341-ae1d-f3fc043492d2","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-04T23:18:34.479049Z","iopub.execute_input":"2024-02-04T23:18:34.479999Z","iopub.status.idle":"2024-02-04T23:18:35.260355Z","shell.execute_reply.started":"2024-02-04T23:18:34.479973Z","shell.execute_reply":"2024-02-04T23:18:35.259290Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Decode rle test, the decoded rle range is [0, 255]","metadata":{"_uuid":"0865bf8b-84a7-46e9-a576-9ec6ce7ccb3a","_cell_guid":"908658c9-53f3-454d-ba24-398821dced67","trusted":true}},{"cell_type":"code","source":"train_rles_filepath = os.path.join(config.input_root_dir, \"train_rles.csv\")\nrows_to_test = 5\nnum_rows = 0\nwith open(train_rles_filepath, \"r\") as in_fp:\n    reader = csv.reader(in_fp, delimiter=\",\")\n    next(reader)\n    for row in reader:\n        full_data_id, rle = row[0], row[1]\n        full_data_id_parts = full_data_id.split(\"_\")\n        subset_name = \"_\".join(full_data_id_parts[:-1])\n        image_name = full_data_id_parts[-1]\n        label_filepath = os.path.join(config.input_root_dir, \"train\", subset_name, \n                                      \"labels\", f\"{image_name}.tif\")\n        print(f\"Label filepath: {label_filepath}, rle {rle}\")\n        label = cv2.imread(label_filepath, cv2.IMREAD_GRAYSCALE)\n        decoded_rle = rle_decode(rle, (label.shape[0], label.shape[1]))\n        print(f\"Decoded rle shape: {decoded_rle.shape}, dtype {decoded_rle.dtype} \" + \n              f\"min {np.min(label)}, max {np.max(label)}\")\n        label_norm = (label / 255).astype(np.uint8)\n        re_encoded_rle = rle_encode(label_norm)\n        #print(f\"Label shape {label_norm.shape}, dtype {label_norm.dtype}\")\n        assert re_encoded_rle == rle, f\"re-encoded {re_encoded_rle} != original {rle}\"\n        num_rows += 1 \n        if num_rows == rows_to_test:\n            break","metadata":{"_uuid":"194c1bee-e488-45b5-8bfe-205fc6b93340","_cell_guid":"24b6af5a-1d4a-4966-b389-4f8778137c80","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-04T23:18:35.261409Z","iopub.execute_input":"2024-02-04T23:18:35.261674Z","iopub.status.idle":"2024-02-04T23:18:35.355265Z","shell.execute_reply.started":"2024-02-04T23:18:35.261651Z","shell.execute_reply":"2024-02-04T23:18:35.354318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Datasets and dataloaders definition","metadata":{"_uuid":"48a0575c-3721-4fcd-9e98-7e9177189650","_cell_guid":"f3078973-3aea-4dde-af45-16741b4428a0","trusted":true}},{"cell_type":"code","source":"def get_train_transform(config: ConfigParams):\n    if config.train_augmentation == \"2.5d_aug\":\n        # Augmentations from https://www.kaggle.com/code/yoyobar/2-5d-cutting-model-baseline-training/notebook\n        # but this requires a specific TTA or resolution because the crop is always zoomed in.\n        return A.Compose(\n            [\n                # 2.5d augmentation\n                A.Rotate(limit=45, p=0.5),\n                # Always zoomed in by upscaling ~2x + random crop\n                A.RandomScale(scale_limit=(0.8, 1.25), p=0.5),\n                A.RandomCrop(\n                    config.model_train_input_size, config.model_train_input_size, p=1\n                ),\n                A.RandomGamma(p=0.75),\n                A.RandomBrightnessContrast(p=0.5),\n                A.GaussianBlur(p=0.5),\n                A.MotionBlur(p=0.5),\n                A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.5),\n                A.ToFloat(max_value=255),\n                ToTensorV2(transpose_mask=True),\n            ]\n        )\n    elif config.train_augmentation == \"my_2.5d_aug\":\n        # Augmentations from https://www.kaggle.com/code/yoyobar/2-5d-cutting-model-baseline-training/notebook\n        # but this requires a specific TTA or resolution because the crop is always zoomed in.\n        # Added Invert augmentation.\n        return A.Compose(\n            [\n                # 2.5d augmentation\n                A.Rotate(limit=45, p=0.5),\n                # Always zoomed in with random scale ~2x (50%) or crop (100%)\n                A.RandomScale(scale_limit=(0.8, 1.25), p=0.5),\n                A.RandomCrop(\n                    config.model_train_input_size, config.model_train_input_size, p=1\n                ),\n                A.RandomGamma(p=0.75),\n                A.RandomBrightnessContrast(p=0.5),\n                A.GaussianBlur(p=0.5),\n                A.MotionBlur(p=0.5),\n                A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.5),\n                A.InvertImg(p=0.5),\n                A.ToFloat(max_value=255),\n                ToTensorV2(transpose_mask=True),\n            ]\n        )\n    elif config.train_augmentation == \"overfit\":\n        # Overfit test\n        return A.Compose(\n            [\n                A.Resize(\n                    config.model_train_input_size,\n                    config.model_train_input_size,\n                    interpolation=cv2.INTER_NEAREST,\n                ),\n                A.ToFloat(max_value=255),\n                ToTensorV2(transpose_mask=True),\n            ]\n        )\n    elif config.train_augmentation == \"my_aug_v2a\":\n        return A.Compose(\n            [\n                A.Rotate(limit=180, p=1.0),\n                # Zoom level similar to validation, no TTA strictly necessary, but it helps\n                A.Resize(\n                    config.model_train_input_size,\n                    config.model_train_input_size,\n                    interpolation=cv2.INTER_NEAREST,\n                ),\n                # This is applied only to the image\n                A.RandomBrightnessContrast(\n                    brightness_limit=0.33,\n                    contrast_limit=0.33,\n                    brightness_by_max=True,\n                    p=1.0,\n                ),\n                # This is applied only to the image\n                A.InvertImg(p=0.5),\n                A.HorizontalFlip(p=0.5),\n                A.VerticalFlip(p=0.5),\n                A.GridDistortion(p=0.5),\n                A.ToFloat(max_value=255),\n                ToTensorV2(),\n            ]\n        )\n    elif config.train_augmentation == \"my_aug_v2b\":\n        # Same of my_aug_v2a without Invert\n        return A.Compose(\n            [\n                A.Rotate(limit=180, p=1.0),\n                # Zoom level similar to validation, no TTA strictly necessary, but it helps\n                A.Resize(\n                    config.model_train_input_size,\n                    config.model_train_input_size,\n                    interpolation=cv2.INTER_NEAREST,\n                ),\n                # This is applied only to the image\n                A.RandomBrightnessContrast(\n                    brightness_limit=0.33,\n                    contrast_limit=0.33,\n                    brightness_by_max=True,\n                    p=1.0,\n                ),\n                A.HorizontalFlip(p=0.5),\n                A.VerticalFlip(p=0.5),\n                A.GridDistortion(p=0.5),\n                A.ToFloat(max_value=255),\n                ToTensorV2(),\n            ]\n        )\n    elif config.train_augmentation == \"my_aug_v3\":\n        return A.Compose(\n            [\n                A.Rotate(\n                    limit=180,\n                    interpolation=cv2.INTER_NEAREST,\n                    rotate_method=\"largest_box\",\n                    border_mode=cv2.BORDER_REFLECT_101,\n                    p=1.0,\n                ),\n                # Zoom by resizing the image to\n                # [0.75 * image_size, 1.25 * image_size] + crop or full resize to model input size\n                # Reference default value: v1d image size 910x1303 and input size 512x512\n                # 0.75 image size = 683x977 ~ 1.33x,1.9x crop size\n                # 1.25 image size = 1138x1628 ~ 2.2x,3.2x crop size\n                # so the crop of 512x512 could be\n                A.RandomScale(\n                    scale_limit=(-0.25, 0.25),\n                    interpolation=cv2.INTER_NEAREST,\n                    p=1.0,\n                ),\n                # Crop (3x probability of being applied) or full image to match tta at eval time\n                A.OneOf(\n                    [\n                        A.RandomCrop(\n                            config.model_train_input_size,\n                            config.model_train_input_size,\n                            p=3.0,\n                        ),\n                        A.Resize(\n                            config.model_train_input_size,\n                            config.model_train_input_size,\n                            interpolation=cv2.INTER_NEAREST,\n                            p=1.0,\n                        ),\n                    ],\n                    p=1.0,\n                ),\n                # This is applied only to the image\n                A.RandomBrightnessContrast(\n                    brightness_limit=0.2,\n                    contrast_limit=0.2,\n                    brightness_by_max=True,\n                    p=1.0,\n                ),\n                A.HorizontalFlip(p=0.5),\n                A.VerticalFlip(p=0.5),\n                A.GridDistortion(interpolation=cv2.INTER_NEAREST, normalized=True),\n                A.ToFloat(max_value=255),\n                ToTensorV2(),\n            ]\n        )\n    else:\n        raise Exception(f\"Augmentation {config.train_augmentation} not suppported\")\n\n\ndef get_val_transform(config: ConfigParams):\n    return A.Compose(\n        [\n            A.Resize(\n                config.model_train_input_size,\n                config.model_train_input_size,\n                interpolation=cv2.INTER_NEAREST,\n            ),\n            A.ToFloat(max_value=255),\n            ToTensorV2(),\n        ],\n    )\n\n\ndef get_test_transform(input_size_height: int, input_size_width: int):\n    # Assuming all test images are of the same size, if we want different resize for different images\n    # we need to change the logic and resize it in the dataset itself\n    return A.Compose(\n        [\n            A.Resize(\n                input_size_height, input_size_width, interpolation=cv2.INTER_NEAREST\n            ),\n            A.ToFloat(max_value=255),\n            ToTensorV2(),\n        ],\n    )\n\n\nclass BloodVesselDataset(Dataset):\n    def __init__(self, selected_dirs, transform, preprocess_function, dataset_with_gt):\n        self.selected_dirs = selected_dirs\n        self.transform = transform\n        self.preprocess_function = preprocess_function\n        self.dataset_with_gt = dataset_with_gt\n        self.samples = []\n        for selected_dir in selected_dirs:\n            images_dir = os.path.join(selected_dir, \"images\")\n            assert os.path.exists(images_dir), f\"{images_dir} does not exist\"\n            images_filepaths = sorted(os.listdir(images_dir))\n            assert len(images_filepaths) > 0\n            labels_filepaths = []\n            if dataset_with_gt:\n                labels_dir = os.path.join(selected_dir, \"labels\")\n                assert os.path.exists(labels_dir), f\"{labels_dir} does not exist\"\n                labels_filepaths = os.listdir(os.path.join(selected_dir, \"labels\"))\n                assert len(images_filepaths) == len(labels_filepaths), (\n                    f\"Number of images {len(images_filepaths)} != number of labels {len(labels_filepaths)} \"\n                    f\"for dir {selected_dir}\"\n                )\n            print(\n                f\"{selected_dir}: images {len(images_filepaths)}, labels {len(labels_filepaths)}\"\n            )\n            for image_filepath in tqdm(images_filepaths, desc=\"image\"):\n                full_image_filepath = os.path.join(\n                    selected_dir, \"images\", image_filepath\n                )\n                if self.dataset_with_gt:\n                    full_label_filepath = os.path.join(\n                        selected_dir, \"labels\", image_filepath\n                    )\n                else:\n                    full_label_filepath = None\n                self.samples.append([full_image_filepath, full_label_filepath])\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n\n        image_path = sample[0]\n        image = cv2.imread(image_path, cv2.IMREAD_GRAYSCALE)\n        image = np.expand_dims(image, axis=-1)\n\n        if sample[1]:\n            label_path = sample[1]\n            label_full_size = cv2.imread(label_path, cv2.IMREAD_GRAYSCALE)\n            # HW, [0.0, 1.0] range, no channel dimension\n            label_full_size = label_full_size.astype(np.float32) / 255.0\n            # Transform image and label\n            transformed = self.transform(image=image, mask=label_full_size)\n            # Add the channel dimension to the label to CHW\n            if len(transformed[\"mask\"].shape) == 2:\n                transformed[\"mask\"] = torch.unsqueeze(transformed[\"mask\"], dim=0)\n            image_preprocessed = self.preprocess_function(transformed[\"image\"])\n            return {\n                \"image\": image_preprocessed,\n                \"label_model_size\": transformed[\"mask\"],\n                # \"label_full_size\": label_full_size,  # Converted to Torch\n                \"file\": image_path,\n                \"shape\": list(image.shape),\n            }\n\n        else:\n            # Transform image\n            transformed = self.transform(image=image)\n            image_preprocessed = self.preprocess_function(transformed[\"image\"])\n            return {\n                \"image\": image_preprocessed,\n                \"file\": image_path,\n                \"shape\": list(image.shape),\n            }\n\n\nclass BloodVesselDatasetTest(BloodVesselDataset):\n    \"\"\"\n    Dataset with specific features for test like time augmentation support (not implemented with albumentation)\n    and returns label at full size\n    \"\"\"\n\n    def __init__(\n        self,\n        selected_dirs,\n        transform,\n        preprocess_function,\n        dataset_with_gt,\n        input_size_width,\n        input_size_height,\n        tta_mode=None,\n    ):\n        super().__init__(selected_dirs, transform, preprocess_function, dataset_with_gt)\n        self.input_size_width = input_size_width\n        self.input_size_height = input_size_height\n        self.tta_mode = tta_mode\n        assert (\n            self.tta_mode in [None, \"\"]\n            or re.search(\"^4.*max\", self.tta_mode)\n            or re.search(\"^5.*max\", self.tta_mode)\n        ), f\"tta_mode {self.tta_mode} not supported\"\n        if self.tta_mode in [None, \"\"]:\n            self.tta_mode = None\n        else:\n            assert (\n                input_size_width % 2 == 0\n            ), f\"input_size_width {input_size_width} is not divisible by 2\"\n            assert (\n                input_size_height % 2 == 0\n            ), f\"input_size_height {input_size_height} is not divisible by 2\"\n\n    def __getitem__(self, idx):\n        if self.tta_mode is None or self.tta_mode == \"\":\n            # Not TTA\n            return super().__getitem__(idx)\n\n        elif re.search(\"^4.*max\", self.tta_mode) or re.search(\"^5.*max\", self.tta_mode):\n            # One crop: full image resized at input_size x input_size\n            # 4 crops option: 4 crops of input_size x input_size at\n            #   top_left, top_right, bottom_left, bottom_right\n            #   No overlapping between the 4 crops\n            # 5 crops option: 5 crops of input_size x input_size at\n            #   top_left, top_right, bottom_left, bottom_right and additional center\n\n            # To create the crops the full image is first resized to (2 * input_size) x (2 * input_size)\n            sample = self.samples[idx]\n\n            image_path = sample[0]\n            full_image_2dims = cv2.imread(image_path, cv2.IMREAD_GRAYSCALE)\n            full_image_3dims = np.expand_dims(full_image_2dims, axis=-1)\n\n            # Get crops\n            image_double_size_2dims = cv2.resize(\n                full_image_2dims,\n                (self.input_size_width * 2, self.input_size_height * 2),\n                cv2.INTER_NEAREST,\n            )\n            image_double_size_3dims = np.expand_dims(image_double_size_2dims, axis=-1)\n            top_left_3dims = image_double_size_3dims[\n                0 : self.input_size_height, 0 : self.input_size_width, :\n            ]\n            top_right_3dims = image_double_size_3dims[\n                0 : self.input_size_height, self.input_size_width :, :\n            ]\n            bottom_left_3dims = image_double_size_3dims[\n                self.input_size_height :, 0 : self.input_size_width, :\n            ]\n            bottom_right_3dims = image_double_size_3dims[\n                self.input_size_height :, self.input_size_width :, :\n            ]\n            center_3dims = image_double_size_3dims[\n                int(self.input_size_height / 2) : -int(self.input_size_height / 2),\n                int(self.input_size_width / 2) : -int(self.input_size_width / 2),\n                :,\n            ]\n            top_left_preprocessed = self.preprocess_function(\n                self.transform(image=top_left_3dims)[\"image\"]\n            )\n            top_right_preprocessed = self.preprocess_function(\n                self.transform(image=top_right_3dims)[\"image\"]\n            )\n            bottom_left_preprocessed = self.preprocess_function(\n                self.transform(image=bottom_left_3dims)[\"image\"]\n            )\n            bottom_right_preprocessed = self.preprocess_function(\n                self.transform(image=bottom_right_3dims)[\"image\"]\n            )\n            # Return center preprocessed even if not used\n            center_preprocessed = self.preprocess_function(\n                self.transform(image=center_3dims)[\"image\"]\n            )\n\n            # Get label and transform it if available\n            if sample[1]:\n                label_path = sample[1]\n                label_full_size = cv2.imread(label_path, cv2.IMREAD_GRAYSCALE)\n                # HW, [0.0, 1.0] range, no channel dimension\n                label_full_size = label_full_size.astype(np.float32) / 255.0\n                # Transform full image and label\n                transformed_full_image_and_label_dict = self.transform(\n                    image=full_image_3dims, mask=label_full_size\n                )\n                # Add the channel dimension to the label to CHW\n                if len(transformed_full_image_and_label_dict[\"mask\"].shape) == 2:\n                    transformed_full_image_and_label_dict[\"mask\"] = torch.unsqueeze(\n                        transformed_full_image_and_label_dict[\"mask\"], dim=0\n                    )\n                full_image_preprocessed = self.preprocess_function(\n                    transformed_full_image_and_label_dict[\"image\"]\n                )\n                return {\n                    \"image\": full_image_preprocessed,\n                    \"label_model_size\": transformed_full_image_and_label_dict[\"mask\"],\n                    # \"label_full_size\": label_full_size,  # Converted to Torch\n                    \"file\": image_path,\n                    \"shape\": list(full_image_2dims.shape),\n                    \"top_left\": top_left_preprocessed,\n                    \"top_right\": top_right_preprocessed,\n                    \"bottom_left\": bottom_left_preprocessed,\n                    \"bottom_right\": bottom_right_preprocessed,\n                    \"center\": center_preprocessed,\n                }\n            else:\n                # Transform image\n                transformed_full_image_dict = self.transform(image=full_image_3dims)\n                full_image_preprocessed = self.preprocess_function(\n                    transformed_full_image_dict[\"image\"]\n                )\n\n                return {\n                    \"image\": full_image_preprocessed,\n                    \"file\": image_path,\n                    \"shape\": list(full_image_2dims.shape),\n                    \"top_left\": top_left_preprocessed,\n                    \"top_right\": top_right_preprocessed,\n                    \"bottom_left\": bottom_left_preprocessed,\n                    \"bottom_right\": bottom_right_preprocessed,\n                    \"center\": center_preprocessed,\n                }\n        else:\n            raise Exception(f\"tta_mode {self.tta_mode} not supported\")\n","metadata":{"_uuid":"9d2ae9d4-110f-449d-95f7-97762f882178","_cell_guid":"81269741-2edc-4a47-9f79-d8ff62fb13df","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-04T23:18:35.356932Z","iopub.execute_input":"2024-02-04T23:18:35.357610Z","iopub.status.idle":"2024-02-04T23:18:35.408372Z","shell.execute_reply.started":"2024-02-04T23:18:35.357571Z","shell.execute_reply":"2024-02-04T23:18:35.407529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Datasets and dataloaders initialization","metadata":{"_uuid":"d5c88588-6925-46d9-b228-c24858a59ff0","_cell_guid":"0e1d5757-4e46-4e12-ba52-316dc7e8821e","trusted":true}},{"cell_type":"code","source":"data_transform_train = get_train_transform(config)\ndata_transform_val = get_val_transform(config)\ntrain_dataset = BloodVesselDataset(\n    [os.path.join(config.input_root_dir, train_dir) for train_dir in config.train_dirs],\n    data_transform_train,\n    preprocess_function=preprocess_function,\n    dataset_with_gt=True,\n)\nval_dataset = BloodVesselDataset(\n    [os.path.join(config.input_root_dir, val_dir) for val_dir in config.val_dirs],\n    data_transform_val,\n    preprocess_function=preprocess_function,\n    dataset_with_gt=True,\n)\ntrain_dataloader = DataLoader(\n    train_dataset, batch_size=config.batch_size, num_workers=4, shuffle=True\n)\ntrain_batches = len(train_dataset) // config.batch_size + int(\n    len(train_dataset) % config.batch_size > 0\n)\nprint(\n    f\"Train dataset samples: {len(train_dataset)}, \"\n    f\"num_batches with batch size {config.batch_size}: {train_batches}\"\n)\nval_dataloader = DataLoader(\n    val_dataset, batch_size=config.batch_size, num_workers=4, shuffle=False\n)\nval_batches = len(val_dataset) // config.batch_size + int(\n    len(val_dataset) % config.batch_size > 0\n)\nprint(\n    f\"Validation dataset samples: {len(val_dataset)}, \"\n    f\"num_batches with batch size {config.batch_size}: {val_batches}\"\n)","metadata":{"_uuid":"d4771bb7-55ad-43cf-b32a-1456951beca9","_cell_guid":"97aff02d-74a0-4bd6-813c-9a4e1ce16333","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-04T23:18:35.411298Z","iopub.execute_input":"2024-02-04T23:18:35.411613Z","iopub.status.idle":"2024-02-04T23:18:35.459803Z","shell.execute_reply.started":"2024-02-04T23:18:35.411587Z","shell.execute_reply":"2024-02-04T23:18:35.459036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Loss","metadata":{"_uuid":"59f81dff-166f-4352-9d93-3b2806102d06","_cell_guid":"1b08d26a-638a-4095-9b40-533b15972994","trusted":true}},{"cell_type":"code","source":"def dice_loss(output, target, eps=1e-7, from_logits=True) -> torch.Tensor:\n    \"\"\"\n    Dice Loss without log\n\n    Args:\n        output: model prediction, shape NCHW\n        target: ground truth, shape NCHW, range [0, 1]\n        eps: small value to avoid dividing by zero\n        from_logits: True if input is logit, False if already in range [0, 1]\n\n    Returns:\n        Dice loss\n    \"\"\"\n    # Comment from smp DiceLoss:\n    # Using Log-Exp as this gives more numerically stable result and does not cause vanishing gradient on\n    # extreme values 0 and 1\n    # The output values are not exactly the same\n    if from_logits:\n        output_stable = nn.LogSigmoid()(output).exp()\n    else:\n        output_stable = output\n    intersection_mul = torch.sum(output_stable * target)\n\n    # The difference w.r.t. smp DiceLoss is the additional square as mentioned in\n    # \"Fundus Images using Modified U-net Convolutional Neural Network\" and the fact\n    # that we don't consider if the ground truth is zero (the loss is zero in that case in smp dice loss)\n    union = torch.sum(output_stable + target)\n    union_clamped = torch.clamp_min(union, eps)\n    # In case of GT and prediction empty, dice score is 1.0\n    if union == 0.0:\n        dice_score = torch.Tensor([1.0]).to(output.device)\n    else:\n        dice_score = 2 * intersection_mul / union_clamped\n    return 1 - dice_score\n\n\ndef dice_log_loss(output, target, eps=1e-7, from_logits=True) -> torch.Tensor:\n    \"\"\"\n    Dice Loss with log\n\n    Args:\n        output: model prediction, shape NCHW\n        target: ground truth, shape NCHW, range [0, 1]\n        eps: small value to avoid dividing by zero\n        from_logits: True if input is logit, False if already in range [0, 1]\n\n    Returns:\n        Dice loss\n    \"\"\"\n    # Comment from smp DiceLoss:\n    # Using Log-Exp as this gives more numerically stable result and does not cause vanishing gradient on\n    # extreme values 0 and 1\n    # The output values are not exactly the same\n    if from_logits:\n        output_stable = nn.LogSigmoid()(output).exp()\n    else:\n        output_stable = output\n    intersection_mul = torch.sum(output_stable * target)\n\n    # The difference w.r.t. smp DiceLoss is the additional square as mentioned in\n    # \"Fundus Images using Modified U-net Convolutional Neural Network\" and the fact\n    # that we don't consider if the ground truth is zero (the loss is zero in that case in smp dice loss)\n    union = torch.sum(output_stable + target)\n    union_clamped = torch.clamp_min(union, eps)\n    # In case of GT and prediction empty, dice score is 1.0\n    if union == 0.0:\n        dice_score = 1.0\n    else:\n        dice_score = 2 * intersection_mul / union_clamped\n    loss_batch = -torch.log(dice_score)\n    return loss_batch\n\n\ndef dice_log_loss_with_square(\n    output, target, eps=1e-7, from_logits=True\n) -> torch.Tensor:\n    \"\"\"\n    Dice Loss with log and squares at denominator\n\n    Args:\n        output: model prediction, shape NCHW\n        target: ground truth, shape NCHW, range [0, 1]\n        eps: small value to avoid dividing by zero\n        from_logits: True if input is logit, False if already in range [0, 1]\n\n    Returns:\n        Dice loss with squares in the formula\n    \"\"\"\n    # Comment from smp DiceLoss:\n    # Using Log-Exp as this gives more numerically stable result and does not cause vanishing gradient on\n    # extreme values 0 and 1\n    # The output values are not exactly the same\n    if from_logits:\n        output_stable = nn.LogSigmoid()(output).exp()\n    else:\n        output_stable = output\n    intersection_mul = torch.sum(output_stable * target)\n\n    # The difference w.r.t. smp DiceLoss is the additional square as mentioned in\n    # \"Fundus Images using Modified U-net Convolutional Neural Network\" and the fact\n    # that we don't consider if the ground truth is zero (the loss is zero in that case in smp dice loss)\n    union_squared = torch.sum(torch.square(output_stable) + torch.square(target))\n    union_clamped = torch.clamp_min(\n        union_squared,\n        eps,\n    )\n    # In case of GT and prediction empty, dice score is 1.0\n    if union_squared == 0.0:\n        dice_score = 1.0\n    else:\n        dice_score = 2 * intersection_mul / union_clamped\n    loss_batch = -torch.log(dice_score)\n    return loss_batch\n\n\ndef init_loss(config):\n    if config.loss_function == \"BCE\":\n        criterion = nn.BCEWithLogitsLoss()\n    elif config.loss_function == \"dice_loss\":\n        criterion = dice_loss\n    elif config.loss_function == \"log_dice_loss\":\n        criterion = dice_log_loss\n    elif config.loss_function == \"log_dice_loss_squared\":\n        criterion = dice_log_loss_with_square\n    elif config.loss_function == \"smp_dice_loss\":\n        criterion = DiceLoss(mode=\"binary\", log_loss=False, from_logits=True)\n    elif config.loss_function == \"smp_log_dice_loss\":\n        criterion = DiceLoss(mode=\"binary\", log_loss=True, from_logits=True)\n    elif config.loss_function == \"smp_focal_loss\":\n        criterion = FocalLoss(mode=\"binary\", alpha=config.focal_loss_alpha, gamma=config.focal_loss_gamma)\n    else:\n        raise Exception(\"Loss function not set, please check the config\")\n    return criterion","metadata":{"_uuid":"271bbd12-e46b-4a94-829d-a22b5afc4662","_cell_guid":"607a76a4-081e-4a84-b0e5-2912bb9ade22","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-04T23:18:35.461260Z","iopub.execute_input":"2024-02-04T23:18:35.461500Z","iopub.status.idle":"2024-02-04T23:18:35.477951Z","shell.execute_reply.started":"2024-02-04T23:18:35.461478Z","shell.execute_reply":"2024-02-04T23:18:35.477012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Optimizer","metadata":{"_uuid":"e6417503-9ec5-484e-982b-d34bcbfb50a6","_cell_guid":"3d0db6f7-a268-4ef4-a85b-707139a89c7f","trusted":true}},{"cell_type":"code","source":"def init_optimizer(config: ConfigParams, model: nn.Module):\n    if config.optimizer == \"SGD\":\n        optimizer = optim.SGD(\n            model.parameters(), lr=config.learning_rate, momentum=config.momentum\n        )\n    elif config.optimizer == \"ADAM\":\n        optimizer = optim.Adam(model.parameters(), lr=config.learning_rate)\n    else:\n        raise Exception(\"Missing optimizer\")\n    return optimizer","metadata":{"_uuid":"cb52438c-5a10-40f3-8da4-8d56178696f0","_cell_guid":"4553b803-929b-4047-8638-0ea453c67857","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-04T23:18:35.479074Z","iopub.execute_input":"2024-02-04T23:18:35.479353Z","iopub.status.idle":"2024-02-04T23:18:35.493804Z","shell.execute_reply.started":"2024-02-04T23:18:35.479329Z","shell.execute_reply":"2024-02-04T23:18:35.493081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Metrics","metadata":{"_uuid":"55537a15-0320-4682-9fa2-70840dea3886","_cell_guid":"e6156789-f8be-49b3-81b0-1d36a3ab156b","trusted":true}},{"cell_type":"code","source":"class Metric(ABC):\n    def __init__(self, to_monitor: bool):\n        self._name = None\n        self._to_monitor = to_monitor\n\n    @abstractmethod\n    def evaluate(self, output, target) -> float:\n        \"\"\"\n        Evaluate the metric on the model prediction and the target\n\n        Args:\n            output: output in range [0, 1] (sigmoid must be already applied)\n            target: ground truth in the range [0, 1]\n\n        Returns:\n            Metric score\n        \"\"\"\n        raise NotImplementedError\n\n    @abstractmethod\n    def is_improved(self, new_value, old_value: Optional) -> bool:\n        raise NotImplementedError\n\n    @property\n    def name(self):\n        return self._name\n\n    @property\n    def to_monitor(self):\n        return self._to_monitor\n\n\nclass DiceScore(Metric):\n    \"\"\"\n    DiceScore based on my custom dice loss class\n\n    1 - dice_loss (without log)\n\n    The diffence with the smp score is that the score is 0\n    when the GT is 0 and the prediction no (instead of 1)\n    \"\"\"\n\n    def __init__(self, to_monitor: bool):\n        super().__init__(to_monitor)\n        self._name = \"dice_score\"\n        self.dice_loss_function = dice_loss\n\n    def evaluate(self, output, target) -> float:\n        return 1 - self.dice_loss_function(output, target, from_logits=False).item()\n\n    def is_improved(self, new_value, old_value: Optional) -> bool:\n        return new_value > old_value if old_value is not None else True\n\n\nclass SMPDiceScore(Metric):\n    \"\"\"\n    Segmentation Models PyTorch dice score based on DiceLoss class\n\n    1 - dice_loss (without log)\n\n    WARNING: dice_loss is 0 when the GT is 0 even with FP predictions\n    \"\"\"\n\n    def __init__(self, to_monitor: bool):\n        super().__init__(to_monitor)\n        self._name = \"smp_dice_score\"\n        self.dice_loss_function = DiceLoss(\n            mode=\"binary\",\n            log_loss=False,\n            from_logits=False,\n        )\n\n    def evaluate(self, output, target) -> float:\n        return 1 - self.dice_loss_function(output, target).item()\n\n    def is_improved(self, new_value, old_value: Optional) -> bool:\n        return new_value > old_value if old_value is not None else True\n\n\nclass SMPDiceLossMetric(Metric):\n    \"\"\"\n    Segmentation Models PyTorch DiceLoss (without log)\n\n    WARNING: dice_loss is 0 when the GT is 0 even with FP predictions\n    \"\"\"\n\n    def __init__(self, to_monitor: bool):\n        super().__init__(to_monitor)\n        self._name = \"smp_dice_loss\"\n        self.dice_loss_function = DiceLoss(\n            mode=\"binary\",\n            log_loss=False,\n            from_logits=False,\n        )\n\n    def evaluate(self, output, target) -> float:\n        return self.dice_loss_function(output, target).item()\n\n    def is_improved(self, new_value, old_value: Optional) -> bool:\n        return new_value < old_value if old_value is not None else True\n\n\ndef init_metrics(config: ConfigParams) -> List[Metric]:\n    \"\"\"\n    Initialize metrics classe\n\n    Args:\n        config: configuration parameters\n\n    Returns:\n        List of metrics to calculate\n    \"\"\"\n    assert (\n        config.val_metric_to_monitor in config.val_metrics_to_log\n    ), f\"config.val_metrics_monitored {config.val_metric_to_monitor} not present in config.val_metrics_logged\"\n    metrics_list = []\n    for val_metric_name in config.val_metrics_to_log:\n        if val_metric_name == \"dice_score\":\n            if val_metric_name == config.val_metric_to_monitor:\n                metrics_list.append(DiceScore(to_monitor=True))\n            else:\n                metrics_list.append(SMPDiceScore(to_monitor=False))\n        elif val_metric_name == \"smp_dice_score\":\n            if val_metric_name == config.val_metric_to_monitor:\n                metrics_list.append(SMPDiceScore(to_monitor=True))\n            else:\n                metrics_list.append(SMPDiceScore(to_monitor=False))\n        elif val_metric_name == \"smp_dice_loss\":\n            if val_metric_name == config.val_metric_to_monitor:\n                metrics_list.append(SMPDiceLossMetric(to_monitor=True))\n            else:\n                metrics_list.append(SMPDiceLossMetric(to_monitor=False))\n        else:\n            raise Exception(f\"Metric {val_metric_name} unknown\")\n    return metrics_list","metadata":{"_uuid":"f4d5048c-6786-4bab-ad02-5a8324c1459a","_cell_guid":"7385a4e3-2b0d-4811-b211-932e55d2d507","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-04T23:18:35.495340Z","iopub.execute_input":"2024-02-04T23:18:35.495707Z","iopub.status.idle":"2024-02-04T23:18:35.514660Z","shell.execute_reply.started":"2024-02-04T23:18:35.495676Z","shell.execute_reply":"2024-02-04T23:18:35.513802Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training loop","metadata":{"_uuid":"1b6565f9-fbfb-460c-8df8-79589b2e56dc","_cell_guid":"f1282eda-7825-4082-abad-4dd1f9238585","trusted":true}},{"cell_type":"code","source":"model.to(device)\n\nloss_criterion = init_loss(config)\noptimizer = init_optimizer(config, model)\nval_metrics = init_metrics(config)\n\nval_metric_to_monitor = None\nfor val_metric in val_metrics:\n    if val_metric.to_monitor:\n        val_metric_to_monitor = val_metric\nassert val_metric_to_monitor is not None, \"No val_metric_to_monitor\"\n\ntraining_metrics = OrderedDict()\n\nbest_monitored_metric_value = None\nbest_epoch_1_index = 0\nconsecutive_no_improvements = 0\n\nfor epoch_id in tqdm(range(config.epochs), desc=\"epoch\"):\n    model.train(True)\n\n    train_total_loss = 0.0\n    running_loss = 0.0\n\n    # Iterate on train batches and update weights using loss\n    for batch_id, data in tqdm(\n        enumerate(train_dataloader), desc=\"batch\", total=train_batches\n    ):\n        global_step = epoch_id * train_batches + batch_id\n        # get the input images and labels\n        images = data[\"image\"]\n        labels = data[\"label_model_size\"]\n        \n        # Move to GPU\n        labels_device = labels.to(device)\n        images_device = images.to(device)\n\n        # forward pass to get outputs\n        preds_logits = model(images_device)\n        preds_sigmoid = nn.Sigmoid()(preds_logits)\n\n        # zero the parameter (weight) gradients\n        optimizer.zero_grad()\n        \n        # calculate the loss between predicted and output image\n        loss = loss_criterion(preds_logits, labels_device)\n\n        # backward pass to calculate the weight gradients\n        loss.backward()\n\n        # update the weights\n        optimizer.step()\n\n        # print loss statistics\n        running_loss += loss.item()\n        if (\n            batch_id % config.num_batches_train_loss_aggregation\n            == config.num_batches_train_loss_aggregation - 1\n        ):\n            avg_sample_loss = running_loss / config.num_batches_train_loss_aggregation\n            print(\n                \"Epoch: {}, Batch: {}, train last {} batches avg. loss: {}\".format(\n                    epoch_id + 1,\n                    batch_id + 1,\n                    config.num_batches_train_loss_aggregation,\n                    avg_sample_loss,\n                )\n            )\n            running_loss = 0.0\n\n        train_total_loss += loss.item()\n\n    train_loss = train_total_loss / train_batches\n    print(\"Epoch: {}, Train avg. loss: {}\".format(epoch_id + 1, train_loss))\n    training_metrics[epoch_id + 1] = {\"train_loss\": train_loss}\n    \n    # Iterate on validation batches\n    model.eval()\n    print(f\"Epoch: {epoch_id + 1}, calculating validation metrics...\")\n    with torch.no_grad():\n        val_total_metrics = defaultdict(float)\n        for batch_id, data in tqdm(enumerate(val_dataloader), total=val_batches):\n            global_step = epoch_id * train_batches + batch_id\n            # get the input images and labels\n            images = data[\"image\"]\n            labels = data[\"label_model_size\"]\n\n            # Move to GPU\n            labels_device = labels.to(device)\n            images_device = images.to(device)\n\n            # forward pass to get outputs\n            preds_logits = model(images_device)\n            preds_sigmoid = nn.Sigmoid()(preds_logits)\n\n            # calculate the metrics on validation batch\n            for single_metric in val_metrics:\n                metric_value = single_metric.evaluate(\n                    preds_sigmoid, labels_device\n                )\n                val_total_metrics[single_metric.name] += metric_value\n\n            # TODO: Calculate surface dice metric\n\n    early_stop = False\n    for single_metric in val_metrics:\n        single_metric_name = single_metric.name\n        single_metric_avg = val_total_metrics[single_metric_name] / val_batches\n        print(\n            f\"Epoch: {epoch_id + 1}, Validation avg. {single_metric_name}: {single_metric_avg}\"\n        )\n        training_metrics[epoch_id + 1][f\"val_{single_metric_name}\"] = single_metric_avg\n        if single_metric.to_monitor:\n            monitored_metric_value = single_metric_avg\n            if single_metric.is_improved(\n                new_value=monitored_metric_value,\n                old_value=best_monitored_metric_value,\n            ):\n                print(\n                    f\"Epoch: {epoch_id + 1}, \"\n                    f\"validation avg. {single_metric_name} improvement from {best_monitored_metric_value} \"\n                    f\"to {monitored_metric_value}\"\n                )\n                best_monitored_metric_value = monitored_metric_value\n                best_epoch_1_index = epoch_id + 1\n                output_model_filename = (\n                    f\"{config.output_dir}/{config.model_name}_{epoch_id + 1}.pt\"\n                )\n                torch.save(model.state_dict(), output_model_filename)\n                print(f\"Model saved to {output_model_filename}\")\n                consecutive_no_improvements = 0\n            else:\n                print(\n                    f\"Epoch: {epoch_id + 1}, \"\n                    f\"NO validation avg. {single_metric_name} improvement from {best_monitored_metric_value} \"\n                    f\"to {monitored_metric_value}\"\n                )\n                consecutive_no_improvements += 1\n                if consecutive_no_improvements > config.patience:\n                    print(\n                        f\"Early stop, patience: {config.patience}, \"\n                        f\"consecutive no improvements: {consecutive_no_improvements}\"\n                    )\n                    early_stop = True\n            \n    if early_stop:\n        break\n\nprint(\n    f\"Train completed, best epoch {best_epoch_1_index} \"\n    f\"with val avg. {val_metric_to_monitor.name} {best_monitored_metric_value}\"\n)\nprint(f\"List of losses for each epoch: {training_metrics}\")\nwith open(os.path.join(config.output_dir, \"training_metrics.json\"), \"w\") as out_fp:\n    json.dump(training_metrics, out_fp, indent=4)","metadata":{"_uuid":"6d456715-fb11-4ab6-886b-1a0a55cbb025","_cell_guid":"d09cc8dc-04ad-41d8-a82f-65fd56fcf363","collapsed":false,"scrolled":true,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-04T23:18:35.516171Z","iopub.execute_input":"2024-02-04T23:18:35.516534Z","iopub.status.idle":"2024-02-04T23:41:41.556323Z","shell.execute_reply.started":"2024-02-04T23:18:35.516502Z","shell.execute_reply":"2024-02-04T23:41:41.555193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Evaluation","metadata":{"_uuid":"56dfe65c-415e-44cb-9ddb-e7d066b0d48d","_cell_guid":"6a4f00fe-97e9-419a-bb71-129f75a022b2","trusted":true}},{"cell_type":"markdown","source":"### Visualize predictions in validation\n\nTest prediction on validation using the train settings (e.g. train input size)","metadata":{"_uuid":"515eb552-37d2-4b54-b688-d70b5e57dae9","_cell_guid":"66b58f5d-cdc0-4384-a9c7-62e7f7514148","trusted":true}},{"cell_type":"code","source":"val_inference_threshold = config.threshold\n\nvalidation_sample_id = \"0178\"\n# Actual validation\nvalidation_image_path = f\"/kaggle/input/blood-vessel-segmentation/train/kidney_3_sparse/images/{validation_sample_id}.tif\"\nvalidation_label_path = f\"/kaggle/input/blood-vessel-segmentation/train/kidney_3_sparse/labels/{validation_sample_id}.tif\"\n\nval_model_filename = f\"{config.model_name}_{best_epoch_1_index}\"\nval_model, val_preprocess_function, _ = init_model(config)\nval_model.load_state_dict(torch.load(os.path.join(config.output_dir, f\"{val_model_filename}.pt\")))\nval_model.eval()\n# Use CPU device for this single image test\nval_model.to(\"cpu\")\n\n# Manual load with dataloader\ninput_image = cv2.imread(validation_image_path, cv2.IMREAD_GRAYSCALE)\noriginal_image_shape_hw = input_image.shape\ninput_image = cv2.resize(input_image, (config.model_train_input_size, config.model_train_input_size), cv2.INTER_NEAREST)\nprint(f\"Original shape: {original_image_shape_hw}, model input shape: {input_image.shape}\")\ninput_image = input_image.astype(np.float32)\ninput_image = input_image / 255.0\ninput_image_preprocessed = val_preprocess_function(input_image)\ninput_image_tensor = torch.from_numpy(input_image_preprocessed)\ninput_image_tensor = torch.reshape(input_image_tensor, (1, 1, config.model_train_input_size, config.model_train_input_size))\n#input_image_tensor.to(device)\n                                   \nprediction = nn.Sigmoid()(val_model(input_image_tensor))\nprint(f\"Prediction max: {torch.max(prediction)}, min: {torch.min(prediction)}\")\nprediction = torch.squeeze(prediction)     \nprediction_thresholded = torch.as_tensor(prediction > val_inference_threshold, dtype=prediction.dtype)\nprint(f\"Prediction thresholded: max {torch.max(prediction_thresholded)}, min {torch.min(prediction_thresholded)}\")\n\n# Prepare prediction raw and prediction thresholded\nprediction = prediction * 255.0\nprediction_npy = prediction.detach().numpy(force=True)\nprediction_npy = prediction_npy.astype(np.uint8)\n# Upsample to the original size\nprediction_npy = cv2.resize(prediction_npy, (original_image_shape_hw[1], original_image_shape_hw[0]), cv2.INTER_NEAREST)\n\nprediction_thresholded = prediction_thresholded * 255.0\nprediction_thresholded_npy = prediction_thresholded.detach().numpy(force=True)\nprediction_thresholded_npy = prediction_thresholded_npy.astype(np.uint8)\n# Upsample to the original size\nprediction_thresholded_npy = cv2.resize(prediction_thresholded_npy, (original_image_shape_hw[1], original_image_shape_hw[0]), cv2.INTER_NEAREST)\n\nplt.figure(figsize=(12, 12))\n\nplt.subplot(2, 2, 1)\ninput_img = tifffile.imread(validation_image_path)\nplt.imshow(input_img, cmap='gray')\nplt.title(f'Input')\nplt.axis('off')\n\nplt.subplot(2, 2, 2)\nlabel_img = tifffile.imread(validation_label_path)\nplt.imshow(label_img, cmap='gray')\nplt.title(f'Ground truth')\nplt.axis('off')\n                                   \nplt.subplot(2, 2, 3)\nplt.imshow(prediction_npy, cmap='gray')\nplt.title(f'Prediction')\nplt.axis('off')\n\nplt.subplot(2, 2, 4)\nplt.imshow(prediction_thresholded_npy, cmap='gray')\nplt.title(f'Prediction_thresholded_{val_inference_threshold}')\nplt.axis('off')\n                          \nplt.show()","metadata":{"_uuid":"e972e5b1-16eb-497b-8382-fd2315fa098d","_cell_guid":"6b65ce91-3c59-457d-8ed5-d5085c16cdb6","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-04T23:41:41.558204Z","iopub.execute_input":"2024-02-04T23:41:41.558902Z","iopub.status.idle":"2024-02-04T23:41:44.547726Z","shell.execute_reply.started":"2024-02-04T23:41:41.558862Z","shell.execute_reply":"2024-02-04T23:41:44.546820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Evaluate the model on the test dataset and prepare the csv","metadata":{"_uuid":"3386b391-e38f-4d3f-8c47-7a8b237bbc62","_cell_guid":"c074ec2c-b587-4afb-8f35-8ae81f060e96","trusted":true}},{"cell_type":"markdown","source":"#### Test time augmentation functions\n\nPredictions are thresholded because rle_encoding used for the challenge requires only 0 or 1 values,\nwhile the validation metrics during the training are calculated on the sigmoid raw output.","metadata":{}},{"cell_type":"code","source":"def predict_crops_tta_max(\n    full_image_prediction_raw,\n    tl_prediction_raw,\n    tr_prediction_raw,\n    bl_prediction_raw,\n    br_prediction_raw,\n    center_prediction_raw,\n    threshold,\n    tta_mode: str,\n) -> Tuple[torch.Tensor, torch.Tensor]:\n    # One crop: full image\n    # Other crops are zoomed in parts: top_left, top_right, bottom_left, bottom_right and center\n    assert tta_mode in [\n        \"4+fullmax\",\n        \"5+fullmax\",\n        \"4max\",\n        \"5max\",\n    ], f\"tta_mode {tta_mode} not supported\"\n\n    # Reference image to put all the predictions resolution: (model_input_size * 2) x (model_input_size * 2)\n\n    # 4 crops to a single output map\n    not_overlapping_crops_pred_ref_top = torch.concat(\n        [tl_prediction_raw, tr_prediction_raw], dim=2\n    )\n    not_overlapping_crops_pred_ref_bottom = torch.concat(\n        [bl_prediction_raw, br_prediction_raw], dim=2\n    )\n    not_overlapping_crops_pred_ref = torch.concat(\n        [not_overlapping_crops_pred_ref_top, not_overlapping_crops_pred_ref_bottom],\n        dim=1,\n    )\n\n    # Pad the center crop to match the desired size\n    pad_size_height, pad_size_width = int(full_image_prediction_raw.shape[1] / 2), int(\n        full_image_prediction_raw.shape[2] / 2\n    )\n    center_crop_pred_ref = torch.nn.functional.pad(\n        center_prediction_raw,\n        (pad_size_width, pad_size_width, pad_size_height, pad_size_height),\n        mode=\"constant\",\n        value=0.0,\n    )\n\n    # Upsample the full image prediction: pass to 4D for the operation and then back to 3D\n    full_image_pred_ref = torch.squeeze(\n        torch.nn.UpsamplingNearest2d(\n            size=[\n                not_overlapping_crops_pred_ref.shape[1],\n                not_overlapping_crops_pred_ref.shape[2],\n            ]\n        )(torch.unsqueeze(full_image_prediction_raw, dim=0)),\n        dim=0,\n    )\n\n    # Only max aggregation between different crops and full image is supported: 4 sides crops or 5 including center\n    if tta_mode.startswith(\"5\"):\n        max_between_crops_pred_ref = torch.maximum(\n            not_overlapping_crops_pred_ref, center_crop_pred_ref\n        )\n    else:\n        # tta with 4 crops\n        max_between_crops_pred_ref = not_overlapping_crops_pred_ref\n\n    if \"full\" in tta_mode:\n        prediction_raw = torch.maximum(max_between_crops_pred_ref, full_image_pred_ref)\n    else:\n        prediction_raw = max_between_crops_pred_ref\n\n    prediction_thresholded = torch.as_tensor(\n        prediction_raw > threshold, dtype=prediction_raw.dtype\n    )\n\n    return prediction_raw, prediction_thresholded","metadata":{"execution":{"iopub.status.busy":"2024-02-04T23:57:03.465718Z","iopub.execute_input":"2024-02-04T23:57:03.466442Z","iopub.status.idle":"2024-02-04T23:57:03.477243Z","shell.execute_reply.started":"2024-02-04T23:57:03.466405Z","shell.execute_reply":"2024-02-04T23:57:03.476316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Run evaluation","metadata":{}},{"cell_type":"code","source":"test_inference_threshold = config.threshold\ntest_inference_input_size_width = config.model_inference_input_size_width\ntest_inference_input_size_height = config.model_inference_input_size_height\ntta_mode = config.tta_mode\ndata_transform_test = get_test_transform(\n    input_size_height=test_inference_input_size_height, input_size_width=test_inference_input_size_width\n)\ntest_dataset = BloodVesselDatasetTest(\n    input_size_width=test_inference_input_size_width,\n    input_size_height=test_inference_input_size_height,\n    selected_dirs=[os.path.join(config.input_root_dir, test_dir) for test_dir in config.test_dirs],\n    transform=data_transform_test,\n    preprocess_function=preprocess_function,\n    dataset_with_gt=False,\n    tta_mode=tta_mode,\n)\ntest_dataloader = DataLoader(\n    test_dataset, batch_size=config.batch_size, num_workers=4, shuffle=False\n)\ntest_batches = len(test_dataset) // config.batch_size + int(\n    len(test_dataset) % config.batch_size > 0\n)\nprint(\n    f\"Test dataset samples: {len(test_dataset)}, \"\n    f\"num_batches with batch size {config.batch_size}: {test_batches}\"\n)\n\n# Load the model from file\ntest_model_filename = f\"{config.model_name}_{best_epoch_1_index}\"\ntest_model, test_preprocess_function, _ = init_model(config)\ntest_model_path = os.path.join(config.output_dir, f\"{test_model_filename}.pt\")\nprint(f\"Loading model from {test_model_path}\")\nassert os.path.exists(test_model_path)\ntest_model.load_state_dict(torch.load(test_model_path))\ntest_model.eval()\ntest_model.to(device)\nsubmission_filepath = os.path.join(config.output_dir, f'{test_model_filename}_thresh_{test_inference_threshold}_submission.csv')\n\nwith torch.no_grad():\n                        \n    with open(submission_filepath, 'w') as out_fp:\n        field_names = ['id', 'rle']\n        writer = csv.DictWriter(out_fp, fieldnames=field_names)\n        writer.writeheader()\n\n        # Iterate on test batches\n        for batch_id, data in tqdm(enumerate(test_dataloader), desc=\"batch\", total=test_batches):\n\n            # get the input images and labels\n            images = data['image']\n            images_paths = data['file']\n            images_shapes = data['shape']\n\n            # print(images_paths)\n            # print(images_shapes)\n            # print(f\"Images shape: {images.shape}, labels shape: {labels.shape}\")\n\n            # Move to GPU\n            images = images.to(device)\n            # forward pass to get outputs\n            predictions_raw = nn.Sigmoid()(test_model(images))\n\n            # Prepare results for TTA\n            if tta_mode:\n                tl_predictions_raw = nn.Sigmoid()(test_model(data[\"top_left\"].to(device)))\n                tr_predictions_raw = nn.Sigmoid()(test_model(data[\"top_right\"].to(device)))\n                bl_predictions_raw = nn.Sigmoid()(test_model(data[\"bottom_left\"].to(device)))\n                br_predictions_raw = nn.Sigmoid()(test_model(data[\"bottom_right\"].to(device)))\n                center_predictions_raw = nn.Sigmoid()(test_model(data[\"center\"].to(device)))\n            \n            print(f\"\\nBatch predictions shape: {predictions_raw.shape}, type {predictions_raw.dtype}\")\n\n            for i in range(images.shape[0]):\n                print(images_paths[i])\n                # Shape format is a list of 3 elements (H-W-C). Each element is a list of batch_size length with the values\n                original_image_height, original_image_width = images_shapes[0][i].item(), images_shapes[1][i].item()\n                print(f\"Single image original shape h x w: {original_image_height} x {original_image_width}\")\n                \n                # Calculate prediction and labels upscaled for 3D metrics and 3D results export (only 0-1 values)\n                # the fast_surface_dice implementation requires two dataframes as input: prediction and label\n                if tta_mode:\n                    prediction_raw, prediction_thresholded = predict_crops_tta_max(\n                        full_image_prediction_raw=predictions_raw[i],\n                        tl_prediction_raw=tl_predictions_raw[i],\n                        tr_prediction_raw=tr_predictions_raw[i],\n                        bl_prediction_raw=bl_predictions_raw[i],\n                        br_prediction_raw=br_predictions_raw[i],\n                        center_prediction_raw=center_predictions_raw[i],\n                        threshold=test_inference_threshold,\n                        tta_mode=tta_mode,\n                    )\n                else:\n                    prediction_raw, prediction_thresholded = predict_no_tta(\n                        prediction_raw=predictions_raw[i], threshold=test_inference_threshold\n                    )\n                print(f\"Single prediction thresholded shape h x w: {prediction_thresholded.shape}\")\n                \n                # Manipulate prediction with torch, 4D input is necessary for upsampling,\n                # upsampling is made at the original image size\n                prediction_thresholded_4d = torch.unsqueeze(\n                    prediction_thresholded, dim=0\n                )  # 1 x 1 x height x width shape\n                prediction_upscaled_th = nn.UpsamplingNearest2d(\n                    size=[original_image_height, original_image_width]\n                )(prediction_thresholded_4d)\n                # Back to 2 dimensions HW\n                prediction_upscaled_th = torch.squeeze(prediction_upscaled_th)\n                # Numpy shape HW for rle encode\n                prediction_upscaled_npy = (\n                    prediction_upscaled_th.cpu().data.detach().numpy(force=True)\n                )\n                print(f\"Single prediction upscaled shape h x w: {prediction_upscaled_npy.shape}\")\n                print(f\"Single prediction upscaled: max {np.max(prediction_upscaled_npy)}, min {np.min(prediction_upscaled_npy)}\")\n                \n                filename_with_ext = os.path.basename(images_paths[i])\n                filename = filename_with_ext[:filename_with_ext.rfind(\".\")]\n                path_parts = images_paths[i].split(\"/\")\n                # Get the subset name (e.g. kidney_5) from the first dir after \"test\"\n                subset = path_parts[path_parts.index(\"test\") + 1]\n                writer.writerow({'id': f'{subset}_{filename}', 'rle': rle_encode(prediction_upscaled_npy)})\n\nprint(f\"\\nSubmission CSV {submission_filepath} file ready\")","metadata":{"_uuid":"b61e76dc-bda7-4659-9bb2-3e13d1b81834","_cell_guid":"5956c6fc-32ea-4801-9525-8e58c088a938","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-02-05T00:02:44.499833Z","iopub.execute_input":"2024-02-05T00:02:44.500201Z","iopub.status.idle":"2024-02-05T00:02:51.311558Z","shell.execute_reply.started":"2024-02-05T00:02:44.500165Z","shell.execute_reply":"2024-02-05T00:02:51.310429Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Copy the submission.csv file","metadata":{"_uuid":"f3699558-bdd3-4002-8e6c-55365b972363","_cell_guid":"a1e4893f-2ffb-4e67-902c-dad0fc6f5325","trusted":true}},{"cell_type":"code","source":"# !rm /kaggle/working/...\nshutil.copy(submission_filepath, \"submission.csv\")","metadata":{"_uuid":"88e1db2d-3e46-4743-9f10-037f2450c13c","_cell_guid":"ce889607-1fa2-4098-9e82-f6f3a74a2635","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-01-26T23:07:02.949945Z","iopub.status.idle":"2024-01-26T23:07:02.950272Z","shell.execute_reply.started":"2024-01-26T23:07:02.950111Z","shell.execute_reply":"2024-01-26T23:07:02.950127Z"},"trusted":true},"execution_count":null,"outputs":[]}]}