{"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":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":7328643,"sourceType":"datasetVersion","datasetId":3896252},{"sourceId":146934283,"sourceType":"kernelVersion"}],"dockerImageVersionId":30580,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## UBC Ovarian Cancer Subtype Classification and Outlier Detection (UBC-OCEAN)","metadata":{}},{"cell_type":"markdown","source":"## 1. Setup","metadata":{}},{"cell_type":"code","source":"!yes | sudo dpkg -i /kaggle/input/libvips-pyvips-installation-and-getting-started/libvips/*.deb\n!pip install /kaggle/input/libvips-pyvips-installation-and-getting-started/pyvips/pyvips-2.2.1-py2.py3-none-any.whl --no-index --find-links /kaggle/input/libvips-pyvips-installation-and-getting-started/pyvips\n!pip install /kaggle/input/ubc-ocean-dataset/packages/imagesize-1.4.1-py2.py3-none-any.whl","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-09-27T17:34:05.377288Z","iopub.execute_input":"2024-09-27T17:34:05.377656Z","iopub.status.idle":"2024-09-27T17:35:21.144232Z","shell.execute_reply.started":"2024-09-27T17:34:05.377628Z","shell.execute_reply":"2024-09-27T17:35:21.143164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport yaml\nfrom pathlib import Path\nfrom tqdm import tqdm\n\nimport numpy as np\nimport pandas as pd\n\nos.environ['OPENCV_IO_MAX_IMAGE_PIXELS'] = str(pow(2, 40))\nimport cv2\nimport pyvips\nimport imagesize\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\n\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\n\nimport matplotlib.pyplot as plt","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-09-27T17:35:21.146252Z","iopub.execute_input":"2024-09-27T17:35:21.146552Z","iopub.status.idle":"2024-09-27T17:35:31.469867Z","shell.execute_reply.started":"2024-09-27T17:35:21.146523Z","shell.execute_reply":"2024-09-27T17:35:31.468791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"competition_dataset = Path('/kaggle/input/UBC-OCEAN')\nexternal_dataset = Path('/kaggle/input/ubc-ocean-dataset')","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:35:31.471353Z","iopub.execute_input":"2024-09-27T17:35:31.47251Z","iopub.status.idle":"2024-09-27T17:35:31.477267Z","shell.execute_reply.started":"2024-09-27T17:35:31.472457Z","shell.execute_reply":"2024-09-27T17:35:31.476292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(competition_dataset / 'test.csv')\nis_submission = df.shape[0] != 1\n\ndf['image_type'] = 'wsi'\ndf.loc[(df['image_width'] < 4000) & (df['image_height'] < 4000), 'image_type'] = 'tma'\ndf['image_path'] = df['image_id'].apply(lambda x: str(competition_dataset / 'test_images' / f'{str(x)}.png'))\n\nprint(f'Dataset Shape: {df.shape}')","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:35:31.479525Z","iopub.execute_input":"2024-09-27T17:35:31.479835Z","iopub.status.idle":"2024-09-27T17:35:31.519241Z","shell.execute_reply.started":"2024-09-27T17:35:31.479808Z","shell.execute_reply":"2024-09-27T17:35:31.51826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Image Utilities","metadata":{}},{"cell_type":"code","source":"def read_image(image_path):\n    \n    \"\"\"\n    Read image using libvips\n\n    Parameters\n    ----------\n    image_path: str\n        Path of the image\n\n    Returns\n    -------\n    image: numpy.ndarray of shape (height, width, 3)\n        Image array\n    \"\"\"\n    \n    image = pyvips.Image.new_from_file(image_path, access='sequential')\n\n    return np.ndarray(\n        buffer=image.write_to_memory(),\n        dtype=np.uint8,\n        shape=[image.height, image.width, image.bands]\n    )\n\n\ndef resize_with_aspect_ratio(image, longest_edge, interpolation=cv2.INTER_LINEAR):\n\n    \"\"\"\n    Resize image while preserving its aspect ratio\n\n    Parameters\n    ----------\n    image: numpy.ndarray of shape (height, width, 3)\n        Image array\n\n    longest_edge: int\n        Desired number of pixels on the longest edge\n\n    interpolation: int\n        OpenCV interpolation enum\n\n    Returns\n    -------\n    image: numpy.ndarray of shape (resized_height, resized_width, 3)\n        Resized image array\n    \"\"\"\n\n    height, width = image.shape[:2]\n    scale = longest_edge / max(height, width)\n    image = cv2.resize(image, dsize=(int(np.ceil(width * scale)), int(np.ceil(height * scale))), interpolation=interpolation)\n\n    return image\n\n\ndef drop_low_std(image, threshold):\n\n    \"\"\"\n    Drop rows and columns that are below the given standard deviation threshold\n\n    Parameters\n    ----------\n    image: numpy.ndarray of shape (height, width, 3)\n        Image array\n\n    threshold: int\n        Standard deviation threshold\n\n    Returns\n    -------\n    image: numpy.ndarray of shape (cropped_height, cropped_width, 3)\n        Cropped image array\n    \"\"\"\n\n    vertical_stds = image.std(axis=(1, 2))\n    horizontal_stds = image.std(axis=(0, 2))\n    cropped_image = image[vertical_stds > threshold, :, :]\n    cropped_image = cropped_image[:, horizontal_stds > threshold, :]\n\n    return cropped_image\n\n\ndef get_largest_contour(image, threshold):\n\n    \"\"\"\n    Get the largest contour from the image\n\n    Parameters\n    ----------\n    image: numpy.ndarray of shape (height, width)\n        Image array\n\n    threshold: int\n        Binarization threshold\n\n    Returns\n    -------\n    bounding_box: list of shape (4)\n        Bounding box with x1, y1, x2, y2 values\n    \"\"\"\n    \n    image_shape = image.shape[:2]\n    image = cv2.threshold(image, threshold, 255, cv2.THRESH_BINARY)[1]\n    contours, _ = cv2.findContours(image, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)\n    del image\n    \n    if len(contours) == 0:\n        x1 = 0\n        x2 = image_shape[1] + 1\n        y1 = 0\n        y2 = image_shape[0] + 1\n    else:\n        contour = max(contours, key=cv2.contourArea)\n        mask = np.zeros(image_shape, np.uint8)\n        cv2.drawContours(mask, [contour], -1, 255, cv2.FILLED)\n\n        y1, y2 = np.min(contour[:, :, 1]), np.max(contour[:, :, 1])\n        x1, x2 = np.min(contour[:, :, 0]), np.max(contour[:, :, 0])\n\n        x1 = int(0.999 * x1)\n        x2 = int(1.001 * x2)\n        y1 = int(0.999 * y1)\n        y2 = int(1.001 * y2)\n\n    bounding_box = [x1, y1, min(x2, image_shape[1]), min(y2, image_shape[0])]\n\n    return bounding_box\n","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:35:31.520797Z","iopub.execute_input":"2024-09-27T17:35:31.521129Z","iopub.status.idle":"2024-09-27T17:35:31.537539Z","shell.execute_reply.started":"2024-09-27T17:35:31.521073Z","shell.execute_reply":"2024-09-27T17:35:31.536497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_image(image, title, mask=None, path=None):\n\n    \"\"\"\n    Visualize the given image\n\n    Parameters\n    ----------\n    image: numpy.ndarray of shape (height, width, channel)\n        Image array\n\n    title: str\n        Title of the plot\n\n    mask: numpy.ndarray of shape (height, width)\n        Mask array\n\n    path: str or None\n        Path of the output file or None (if path is None, plot is displayed with selected backend)\n    \"\"\"\n\n    fig, ax = plt.subplots(figsize=(8, 8))\n    ax.imshow(image)\n    if mask is not None:\n        ax.imshow(mask, alpha=0.5)\n    ax.set_xlabel('')\n    ax.set_ylabel('')\n    ax.tick_params(axis='x', labelsize=15, pad=10)\n    ax.tick_params(axis='y', labelsize=15, pad=10)\n    ax.set_title(title, size=15, pad=12.5, loc='center', wrap=True)\n\n    if path is None:\n        plt.show()\n    else:\n        plt.savefig(path)\n        plt.close(fig)\n","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:35:31.539021Z","iopub.execute_input":"2024-09-27T17:35:31.539425Z","iopub.status.idle":"2024-09-27T17:35:31.555605Z","shell.execute_reply.started":"2024-09-27T17:35:31.539389Z","shell.execute_reply":"2024-09-27T17:35:31.554506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. Models","metadata":{}},{"cell_type":"code","source":"class Decoder(nn.Module):\n\n    def __init__(self, encoder_dim=(32, 64, 128, 256), upscale=4, num_classes=1):\n\n        super(Decoder, self).__init__()\n\n        self.conv = nn.ModuleList([\n            nn.Sequential(\n                nn.Conv2d(encoder_dim[i] + encoder_dim[i - 1], encoder_dim[i - 1], 3, 1, 1, bias=False),\n                nn.BatchNorm2d(encoder_dim[i - 1]),\n                nn.ReLU(inplace=True)\n            ) for i in range(1, len(encoder_dim))\n        ])\n\n        self.logit = nn.Conv2d(encoder_dim[0], num_classes, 1, 1, 0)\n        self.up = nn.Upsample(scale_factor=upscale, mode='bilinear')\n\n    def forward(self, feature):\n\n        for i in range(len(feature) - 1, 0, -1):\n            f_up = F.interpolate(feature[i], scale_factor=2, mode='bilinear')\n            f = torch.cat([feature[i - 1], f_up], dim=1)\n            f_down = self.conv[i - 1](f)\n            feature[i - 1] = f_down\n\n        x = self.logit(feature[0])\n        out = self.up(x)\n\n        return out\n\n\nclass SegModel(nn.Module):\n\n    def __init__(self, num_classes):\n\n        super(SegModel, self).__init__()\n\n        self.encoder = timm.create_model(\n            'maxvit_tiny_tf_512.in1k',\n            features_only=True,\n            pretrained=False,\n            out_indices=[1, 2, 3, 4]\n        )\n\n        self.decoder = Decoder(\n            self.encoder.feature_info.channels(),\n            upscale=self.encoder.feature_info.reduction()[0],\n            num_classes=num_classes\n        )\n\n    def forward(self, x):\n\n        feat_maps = self.encoder(x)\n        logits = self.decoder(feat_maps)\n        masks = torch.sigmoid(logits)\n\n        return masks","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:35:31.556856Z","iopub.execute_input":"2024-09-27T17:35:31.557239Z","iopub.status.idle":"2024-09-27T17:35:31.570651Z","shell.execute_reply.started":"2024-09-27T17:35:31.557205Z","shell.execute_reply":"2024-09-27T17:35:31.569701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ClassificationHead(nn.Module):\n\n    def __init__(self, input_dimensions, cancer_output_dimensions):\n\n        super(ClassificationHead, self).__init__()\n\n        self.cancer_head = nn.Linear(input_dimensions, cancer_output_dimensions, bias=True)\n\n    def forward(self, x):\n\n        cancer_output = self.cancer_head(x)\n\n        return cancer_output\n\n\nclass TimmConvImageClassificationModel(nn.Module):\n\n    def __init__(self, model_name, pretrained, backbone_args, pooling_type, dropout_rate, freeze_parameters, head_args):\n\n        super(TimmConvImageClassificationModel, self).__init__()\n\n        self.backbone = timm.create_model(\n            model_name=model_name,\n            pretrained=pretrained,\n            **backbone_args\n        )\n\n        if freeze_parameters:\n            for parameter in self.backbone.parameters():\n                parameter.requires_grad = False\n\n        input_features = self.backbone.get_classifier().in_features\n\n        self.pooling_type = pooling_type\n        self.dropout = nn.Dropout(dropout_rate) if dropout_rate > 0 else nn.Identity()\n        self.head = ClassificationHead(input_dimensions=input_features, **head_args)\n\n    def forward(self, x):\n\n        x = self.backbone.forward_features(x)\n\n        if self.pooling_type == 'avg':\n            x = F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.pooling_type == 'max':\n            x = F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.pooling_type == 'concat':\n            x = torch.cat([\n                F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1),\n                F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n            ], dim=-1)\n\n        x = self.dropout(x)\n        cancer_output = self.head(x)\n\n        return cancer_output\n","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:35:31.572248Z","iopub.execute_input":"2024-09-27T17:35:31.572704Z","iopub.status.idle":"2024-09-27T17:35:31.587846Z","shell.execute_reply.started":"2024-09-27T17:35:31.572669Z","shell.execute_reply":"2024-09-27T17:35:31.586988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_segmentation_model(model_directory, model_file_name, device):\n    \n    \"\"\"\n    Load model and pretrained weights from the given model directory\n\n    Parameters\n    ----------\n    model_directory: pathlib.Path\n        Path of the model directory\n\n    model_file_name: str\n        Name of the model weights file\n\n    device: torch.device\n        Location of the model\n\n    Returns\n    -------\n    model: torch.nn.Module\n        Model with weights loaded\n    \"\"\"\n\n    model = SegModel(num_classes=1)\n    model.load_state_dict(torch.load(model_directory / model_file_name), strict=True)\n    model.to(device)\n    model.eval()\n    print(f'{model.__class__.__name__} model\\'s weights are loaded from {model_directory / model_file_name}')\n\n    return model\n\n\ndef load_classification_model(model_directory, model_file_names, device):\n    \n    \"\"\"\n    Load model and pretrained weights from the given model directory\n\n    Parameters\n    ----------\n    model_directory: pathlib.Path\n        Path of the model directory\n\n    model_file_names: list\n        List of names of the model weights files\n\n    device: torch.device\n        Location of the model\n\n    Returns\n    -------\n    model: dict\n        Dictionary of models with weights loaded\n    \"\"\"\n\n    config = yaml.load(open(model_directory / 'config.yaml', 'r'), Loader=yaml.FullLoader)\n    config['model']['model_args']['pretrained'] = False\n        \n    models = {}\n\n    for model_file_name in model_file_names:\n        model = eval(config['model']['model_class'])(**config['model']['model_args'])\n        model.load_state_dict(torch.load(model_directory / model_file_name))\n        model.to(device)\n        model.eval()\n        models[model_file_name] = model\n        print(f'{model.__class__.__name__} model\\'s weights are loaded from {model_directory / model_file_name}')\n\n    return models, config\n","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:35:31.589175Z","iopub.execute_input":"2024-09-27T17:35:31.589518Z","iopub.status.idle":"2024-09-27T17:35:31.602645Z","shell.execute_reply.started":"2024-09-27T17:35:31.589457Z","shell.execute_reply":"2024-09-27T17:35:31.601612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"segmentation_model = load_segmentation_model(\n    model_directory=external_dataset / 'segmentation',\n    model_file_name='maxvit_tiny_512_v1_final_epoch_13.pt',\n    device=torch.device('cuda')\n)\n\nmodels, config = load_classification_model(\n    model_directory=external_dataset / 'efficientnetv2s_1024_crop_16_v2.1',\n    model_file_names=[\n        'model_fold1_epoch_8.pt',\n        'model_fold2_epoch_15.pt',\n        'model_fold3_epoch_15.pt',\n        'model_fold4_epoch_15.pt',\n        'model_fold5_epoch_14.pt',\n    ],\n    device=torch.device('cuda')\n)\n","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:35:31.60585Z","iopub.execute_input":"2024-09-27T17:35:31.606221Z","iopub.status.idle":"2024-09-27T17:35:41.892507Z","shell.execute_reply.started":"2024-09-27T17:35:31.606191Z","shell.execute_reply":"2024-09-27T17:35:41.891516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4. Inference","metadata":{}},{"cell_type":"code","source":"segmentation_size = 384\nsegmentation_device = torch.device('cuda')\nsegmentation_amp = True\n\nsegmentation_transforms = A.Compose([\n    A.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),\n    ToTensorV2()\n])","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:35:41.89382Z","iopub.execute_input":"2024-09-27T17:35:41.894157Z","iopub.status.idle":"2024-09-27T17:35:41.901156Z","shell.execute_reply.started":"2024-09-27T17:35:41.894128Z","shell.execute_reply":"2024-09-27T17:35:41.90012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classification_crop_size = 1024\nclassification_top_n_crops = 16\n\nclassification_device = torch.device('cuda')\nclassification_amp = True\n\ntta = False\ntta_flip_dimensions = [(2,), (3,), (2, 3)]\n\nclassification_transforms = A.Compose([\n    A.Resize(\n        height=1024,\n        width=1024,\n        interpolation=cv2.INTER_NEAREST,\n        always_apply=True\n    ),\n    A.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225],\n        max_pixel_value=255,\n        always_apply=True\n    ),\n    ToTensorV2(always_apply=True)\n])","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:35:41.90242Z","iopub.execute_input":"2024-09-27T17:35:41.903047Z","iopub.status.idle":"2024-09-27T17:35:41.915726Z","shell.execute_reply.started":"2024-09-27T17:35:41.90301Z","shell.execute_reply":"2024-09-27T17:35:41.914797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_max_quo(num, div):\n    quo, rem = divmod(num, div)\n    max_quo = quo + 1 if rem > 0 else quo\n    return max_quo\n","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:35:41.916901Z","iopub.execute_input":"2024-09-27T17:35:41.917613Z","iopub.status.idle":"2024-09-27T17:35:41.926595Z","shell.execute_reply.started":"2024-09-27T17:35:41.917584Z","shell.execute_reply":"2024-09-27T17:35:41.925806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = []\nimage_ids = []\n\nfor idx, row in df.iterrows():\n    \n    image_id = row['image_id']\n    image_type = row['image_type']\n    image_path = row['image_path']\n    \n    if image_type == 'tma':\n\n        image = cv2.imread(image_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        # Drop low standard deviation rows and columns (white areas with less tissue)\n        image = drop_low_std(image=image, threshold=10)\n        \n        tma_predictions = torch.zeros(1, 6)\n        \n        inputs = classification_transforms(image=image)['image']\n        inputs = inputs.to(classification_device)\n\n        if tta:\n            inputs = torch.stack((\n                inputs,\n                torch.flip(inputs, dims=(1,)),\n                torch.flip(inputs, dims=(2,)),\n                torch.flip(inputs, dims=(1, 2))\n            ), dim=0)\n        else:\n            inputs = torch.unsqueeze(inputs, dim=0)\n\n        for model_idx, models in enumerate([models]):\n            for model in models.values():\n                with torch.no_grad():\n                    if classification_amp:\n                        with torch.autocast(device_type='cuda', dtype=torch.float16):\n                            outputs = model(inputs)\n                    else:\n                        outputs = model(inputs)\n\n                outputs = outputs.cpu()\n                if tta:\n                    outputs = torch.mean(outputs, dim=0)\n                else:\n                    outputs = torch.squeeze(outputs, dim=0)\n\n                tma_predictions += outputs / len(models)\n                \n        tma_predictions = torch.softmax(tma_predictions, dim=-1).numpy()\n        predictions.append(tma_predictions)\n        image_ids.append(np.array([image_id]))\n        \n    else:\n        \n        image_thumbnail_path = str(competition_dataset / 'test_thumbnails' / f'{image_id}_thumbnail.png')\n        image_thumbnail = cv2.imread(image_thumbnail_path)\n        image_thumbnail = cv2.cvtColor(image_thumbnail, cv2.COLOR_BGR2RGB)\n        \n        if is_submission is False:\n            print(f'Image {image_id} {image_thumbnail.shape} thumbnail is loaded from {image_thumbnail_path}')\n            visualize_image(image=image_thumbnail, title=f'Image {image_id} Thumbnail {image_thumbnail.shape}', path='./thumbnail_raw.png')\n            \n        image_thumbnail_height, image_thumbnail_width = image_thumbnail.shape[:2]\n        thumbnail_height_padding = segmentation_size - image_thumbnail_height % segmentation_size\n        thumbnail_width_padding = segmentation_size - image_thumbnail_width % segmentation_size\n        image_thumbnail_padded = np.pad(image_thumbnail, ((0 + 64, thumbnail_height_padding + 64), (0 + 64, thumbnail_width_padding + 64), (0, 0)))\n        mask_padded = np.zeros((image_thumbnail_height + thumbnail_height_padding, image_thumbnail_width + thumbnail_width_padding), dtype=np.float32)\n        \n        if is_submission is False:\n            print(f'Image {image_id} {image_thumbnail_padded.shape} thumbnail and mask {mask_padded.shape} are padded')\n            visualize_image(image=image_thumbnail_padded, title=f'Image {image_id} Thumbnail Padded {image_thumbnail_padded.shape}', path='./thumbnail_padded.png')\n        \n        for height in range((image_thumbnail_height + thumbnail_height_padding) // segmentation_size):\n            height_1, height_2 = height * segmentation_size, (height + 1) * segmentation_size\n            for width in range((image_thumbnail_width + thumbnail_width_padding) // segmentation_size):\n                width_1, width_2 = width * segmentation_size, (width + 1) * segmentation_size\n                \n                image_thumbnail_tile = image_thumbnail_padded[height_1:height_2 + 128, width_1:width_2 + 128]\n                inputs = segmentation_transforms(image=image_thumbnail_tile)['image'].to(segmentation_device)\n                inputs = torch.stack((\n                    inputs,\n                    torch.flip(inputs, dims=(1,)),\n                    torch.flip(inputs, dims=(2,)),\n                    torch.flip(inputs, dims=(1, 2))\n                ), dim=0)\n                \n                with torch.no_grad():\n                    if segmentation_amp:\n                        with torch.autocast(device_type='cuda', dtype=torch.float16):\n                            outputs = segmentation_model(inputs)\n                    else:\n                        outputs = segmentation_model(inputs)\n                    \n                outputs = outputs.cpu()\n                outputs = torch.stack((\n                    outputs[0],\n                    torch.flip(outputs[1], dims=(1,)),\n                    torch.flip(outputs[2], dims=(2,)),\n                    torch.flip(outputs[3], dims=(1, 2)),\n                ), dim=0)\n                outputs = torch.mean(outputs, dim=0).squeeze().cpu().numpy()\n                mask_padded[height_1:height_2, width_1:width_2] = outputs[64:448, 64:448]\n                    \n                if is_submission is False:\n                    print(f'Predicted x ({width_1}-{width_2}) y ({height_1}-{height_2})')\n                        \n        mask = mask_padded[:image_thumbnail_height, :image_thumbnail_width]\n        del mask_padded\n        \n        if is_submission is False:\n            print(f'Image {image_id} {mask.shape} mask is predicted')\n            visualize_image(image=image_thumbnail, title=f'Image {image_id} Thumbnail and Mask {mask.shape}', mask=mask, path='./thumbnail_mask.png')\n            \n        # Extract image size from the image headers\n        image_width, image_height = imagesize.get(image_path)\n        # Cast mask soft predictions to 8-bit integer and upsample\n        mask = np.uint8(mask * 255)\n        mask = cv2.resize(mask, (image_width, image_height), cv2.INTER_NEAREST)\n        \n        if is_submission is False:\n            print(f'Image {image_id} mask {mask.shape} is upsampled')\n            visualize_image(image=mask, title=f'Image {image_id} Mask {mask.shape}', path='./resized_mask.png')\n\n        image_height_max_quotient = get_max_quo(image_height, classification_crop_size)\n        image_width_max_quotient = get_max_quo(image_width, classification_crop_size)\n        crops = []\n\n        for h in range(image_height_max_quotient):\n            for w in range(image_width_max_quotient):\n                width_2, height_2 = min(image_width, (w + 1) * classification_crop_size), min(image_height, (h + 1) * classification_crop_size)\n                width_1, height_1 = width_2 - classification_crop_size, height_2 - classification_crop_size\n\n                mask_crop = mask[height_1:height_2, width_1:width_2]\n                area = np.sum(mask_crop) / 255.0\n                crops.append([area, width_1, height_1, width_2, height_2])\n                \n        del mask\n        image = read_image(image_path=image_path)\n        image_crops = []\n        \n        # Sort crops by mask area in descending order\n        crops = np.array(crops)\n        crops = crops[crops[:, 0].argsort()[::-1]].astype(np.int32)\n        \n        if is_submission is False:\n            print(f'{crops.shape[0]} crops are extracted from the mask')\n        \n        # Crop the raw image and delete it\n        for crop_idx, crop in enumerate(crops[:classification_top_n_crops]):\n            area, width_1, height_1, width_2, height_2 = crop\n            image_crops.append(image[height_1:height_2, width_1:width_2, :].copy())    \n        del image\n                \n        crop_predictions = torch.zeros(classification_top_n_crops, 6)\n        for crop_idx, image_crop in enumerate(image_crops):\n            \n            if is_submission is False:\n                print(f'Image {image_id} crop {crop_idx}')\n                visualize_image(image=image_crop, title=f'Image {image_id} Crop {crop_idx}', path=f'./cropped_{crop_idx}.png')\n\n            inputs = classification_transforms(image=image_crop)['image']\n            inputs = inputs.to(classification_device)\n\n            if tta:\n                inputs = torch.stack((\n                    inputs,\n                    torch.flip(inputs, dims=(1,)),\n                    torch.flip(inputs, dims=(2,)),\n                    torch.flip(inputs, dims=(1, 2))\n                ), dim=0)\n            else:\n                inputs = torch.unsqueeze(inputs, dim=0)\n\n            for model_idx, models in enumerate([models]):\n                for model in models.values():\n                    with torch.no_grad():\n                        if classification_amp:\n                            with torch.autocast(device_type='cuda', dtype=torch.float16):\n                                outputs = model(inputs)\n                        else:\n                            outputs = model(inputs)\n\n                    outputs = outputs.cpu()\n                    if tta:\n                        outputs = torch.mean(outputs, dim=0)\n                    else:\n                        outputs = torch.squeeze(outputs, dim=0)\n\n                    crop_predictions[crop_idx] += outputs / len(models)\n\n        crop_predictions = torch.softmax(crop_predictions, dim=-1).numpy()\n        predictions.append(crop_predictions)\n        image_ids.append(np.array([image_id] * classification_top_n_crops))\n        \npredictions = np.concatenate(predictions)\nimage_ids = np.concatenate(image_ids)","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:44:04.528246Z","iopub.execute_input":"2024-09-27T17:44:04.528634Z","iopub.status.idle":"2024-09-27T17:44:47.479817Z","shell.execute_reply.started":"2024-09-27T17:44:04.528601Z","shell.execute_reply":"2024-09-27T17:44:47.478959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5. Post processing","metadata":{}},{"cell_type":"code","source":"df_predictions = pd.DataFrame(data=image_ids, columns=['image_id'])\n\nlabel_mapping = {\n    0: 'HGSC',\n    1: 'EC',\n    2: 'CC',\n    3: 'LGSC',\n    4: 'MC',\n    5: 'Other',\n}\n\nfor label_idx, label in label_mapping.items():\n    df_predictions[f'{label}_prediction'] = predictions[:, label_idx]\n    \ndf_predictions","metadata":{"execution":{"iopub.status.busy":"2023-12-28T05:25:29.880156Z","iopub.execute_input":"2023-12-28T05:25:29.881281Z","iopub.status.idle":"2023-12-28T05:25:29.914085Z","shell.execute_reply.started":"2023-12-28T05:25:29.881217Z","shell.execute_reply":"2023-12-28T05:25:29.91282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"prediction_columns = [f'{label}_prediction' for label in list(label_mapping.values())]\n\ndf_predictions = df_predictions.groupby('image_id')[prediction_columns].mean().reset_index()\ndf_predictions","metadata":{"execution":{"iopub.status.busy":"2023-12-28T05:25:38.264429Z","iopub.execute_input":"2023-12-28T05:25:38.265236Z","iopub.status.idle":"2023-12-28T05:25:38.294498Z","shell.execute_reply.started":"2023-12-28T05:25:38.265199Z","shell.execute_reply":"2023-12-28T05:25:38.293395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_predictions['label'] = np.argmax(df_predictions[prediction_columns], axis=-1)\ndf_predictions['label'] = df_predictions['label'].map(label_mapping)\n\ndf_predictions","metadata":{"execution":{"iopub.status.busy":"2023-12-28T05:25:47.459292Z","iopub.execute_input":"2023-12-28T05:25:47.460162Z","iopub.status.idle":"2023-12-28T05:25:47.474897Z","shell.execute_reply.started":"2023-12-28T05:25:47.460127Z","shell.execute_reply":"2023-12-28T05:25:47.474039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 6. Submission","metadata":{}},{"cell_type":"code","source":"df_submission = pd.read_csv(competition_dataset / 'sample_submission.csv').drop(columns=['label'])\ndf_submission","metadata":{"execution":{"iopub.status.busy":"2023-12-28T05:25:55.312101Z","iopub.execute_input":"2023-12-28T05:25:55.312491Z","iopub.status.idle":"2023-12-28T05:25:55.325564Z","shell.execute_reply.started":"2023-12-28T05:25:55.312458Z","shell.execute_reply":"2023-12-28T05:25:55.32461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission = df_submission.merge(df_predictions.loc[:, ['image_id', 'label']], on='image_id', how='left')\ndf_submission","metadata":{"execution":{"iopub.status.busy":"2023-12-28T05:25:56.950593Z","iopub.execute_input":"2023-12-28T05:25:56.951322Z","iopub.status.idle":"2023-12-28T05:25:56.968065Z","shell.execute_reply.started":"2023-12-28T05:25:56.951289Z","shell.execute_reply":"2023-12-28T05:25:56.967121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-12-28T05:25:58.660082Z","iopub.execute_input":"2023-12-28T05:25:58.660474Z","iopub.status.idle":"2023-12-28T05:25:58.668039Z","shell.execute_reply.started":"2023-12-28T05:25:58.660444Z","shell.execute_reply":"2023-12-28T05:25:58.667079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}