{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## RSNA 2023 Abdominal Trauma Detection","metadata":{}},{"cell_type":"markdown","source":"## 1. Setup","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/rsna-2023-abdominal-trauma-detection-dataset/packages/monai-1.2.0-202306081546-py3-none-any.whl --no-index --find-links /kaggle/input/rsna-2023-abdominal-trauma-detection-dataset/packages","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:51:33.483658Z","iopub.execute_input":"2023-10-13T09:51:33.484113Z","iopub.status.idle":"2023-10-13T09:51:44.483538Z","shell.execute_reply.started":"2023-10-13T09:51:33.484084Z","shell.execute_reply":"2023-10-13T09:51:44.482186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport yaml\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nimport monai.transforms as T","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-10-13T09:51:44.486103Z","iopub.execute_input":"2023-10-13T09:51:44.48725Z","iopub.status.idle":"2023-10-13T09:52:02.802426Z","shell.execute_reply.started":"2023-10-13T09:51:44.487206Z","shell.execute_reply":"2023-10-13T09:52:02.801233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"competition_dataset = Path('/kaggle/input/rsna-2023-abdominal-trauma-detection')\nexternal_dataset = Path('/kaggle/input/rsna-2023-abdominal-trauma-detection-dataset')","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:52:02.804079Z","iopub.execute_input":"2023-10-13T09:52:02.804745Z","iopub.status.idle":"2023-10-13T09:52:02.811521Z","shell.execute_reply.started":"2023-10-13T09:52:02.804704Z","shell.execute_reply":"2023-10-13T09:52:02.809816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(competition_dataset / 'sample_submission.csv')\nprint(f'Dataset Shape: {df.shape}')","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:52:02.815263Z","iopub.execute_input":"2023-10-13T09:52:02.81595Z","iopub.status.idle":"2023-10-13T09:52:02.845002Z","shell.execute_reply.started":"2023-10-13T09:52:02.815895Z","shell.execute_reply":"2023-10-13T09:52:02.843967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. DICOM Utilities","metadata":{}},{"cell_type":"code","source":"def shift_bits(image, dicom, bits_allocated=None, bits_stored=None):\n\n    \"\"\"\n    Shift bits using allocated and stored bits\n\n    Parameters\n    ----------\n    image: numpy.ndarray of shape (height, width)\n        Image array\n\n    dicom: pydicom.dataset.FileDataset\n        DICOM dataset\n\n    bits_allocated: int, str ('dataset') or None\n        Number of bits allocated\n\n    bits_stored: int, str ('dataset') or None\n        Number of bits stored\n\n    Returns\n    -------\n    image: numpy.ndarray of shape (height, width)\n        Image array with shifted bits\n    \"\"\"\n\n    if bits_allocated == 'dataset':\n        try:\n            bits_allocated = dicom.BitsAllocated\n        except AttributeError:\n            bits_allocated = None\n\n    if bits_stored == 'dataset':\n        try:\n            bits_stored = dicom.BitsStored\n        except AttributeError:\n            bits_stored = None\n\n    if bits_allocated is not None and bits_stored is not None:\n        bit_shift = bits_allocated - bits_stored\n    else:\n        bit_shift = None\n\n    if bit_shift is not None:\n        dtype = image.dtype\n        image = (image << bit_shift).astype(dtype) >> bit_shift\n\n    return image\n\n\ndef rescale_pixel_values(image, dicom, rescale_slope=None, rescale_intercept=None):\n\n    \"\"\"\n    Rescale pixel values using rescale slope and intercept as a linear function\n\n    Parameters\n    ----------\n    image: numpy.ndarray of shape (height, width)\n        Image array\n\n    dicom: pydicom.dataset.FileDataset\n        DICOM dataset\n\n    rescale_slope: int, str ('dataset') or None\n        Rescale slope for rescaling pixel values\n\n    rescale_intercept: int, str ('dataset') or None\n        Rescale intercept for rescaling pixel values\n\n    Returns\n    -------\n    image: numpy.ndarray of shape (height, width)\n        Image array with rescaled pixel values\n    \"\"\"\n\n    if rescale_slope == 'dataset':\n        try:\n            rescale_slope = dicom.RescaleSlope\n        except AttributeError:\n            rescale_slope = None\n\n    if rescale_intercept == 'dataset':\n        try:\n            rescale_intercept = dicom.RescaleIntercept\n        except AttributeError:\n            rescale_intercept = None\n\n    if rescale_slope is not None and rescale_intercept is not None:\n        image = image.astype(np.float32)\n        image = image * rescale_slope + rescale_intercept\n\n    return image\n\n\ndef window_pixel_values(image, dicom, window_center=None, window_width=None):\n\n    \"\"\"\n    Window pixel values using window center and width\n\n    Parameters\n    ----------\n    image: numpy.ndarray of shape (height, width)\n        Image array\n\n    dicom: pydicom.dataset.FileDataset\n        DICOM dataset\n\n    window_center: int, str ('dataset') or None\n        Window center for windowing pixel values\n\n    window_width: int, str ('dataset') or None\n        Window width for windowing pixel values\n\n    Returns\n    -------\n    image: numpy.ndarray of shape (height, width)\n        Image array with windowed pixel values\n    \"\"\"\n\n    if window_center == 'dataset':\n        try:\n            window_center = dicom.WindowCenter\n        except AttributeError:\n            window_center = None\n\n    if window_width == 'dataset':\n        try:\n            window_width = dicom.WindowWidth\n        except AttributeError:\n            window_width = None\n\n    if window_center is not None and window_width is not None:\n        image_min = window_center - window_width // 2\n        image_max = window_center + window_width // 2\n        image = np.clip(image.copy(), image_min, image_max)\n\n    return image\n\n\ndef invert_pixel_values(image, dicom, photometric_interpretation=None, max_pixel_value=255):\n\n    \"\"\"\n    Invert pixel values using given max pixel value\n\n    Parameters\n    ----------\n    image: numpy.ndarray of shape (height, width)\n        Image array\n\n    dicom: pydicom.dataset.FileDataset\n        DICOM dataset\n\n    photometric_interpretation: str or None\n        Interpretation of the pixel data\n\n    max_pixel_value: int or None\n        Max pixel value used for inverting pixel values\n\n    Returns\n    -------\n    image: numpy.ndarray of shape (height, width)\n        Image array with inverted pixel values\n    \"\"\"\n\n    if photometric_interpretation == 'dataset':\n        try:\n            photometric_interpretation = dicom.PhotometricInterpretation\n        except AttributeError:\n            photometric_interpretation = None\n\n    if photometric_interpretation == 'MONOCHROME1':\n        image = max_pixel_value - image\n\n    return image\n\n\ndef adjust_pixel_values(\n        image, dicom,\n        bits_allocated=None, bits_stored=None,\n        rescale_slope=None, rescale_intercept=None,\n        window_centers=None, window_widths=None,\n        photometric_interpretation=None, max_pixel_value=255\n):\n\n    \"\"\"\n    Adjust pixel values by shifting bits, windowing, rescaling and inverting\n\n    Parameters\n    ----------\n    image: numpy.ndarray of shape (height, width)\n        Image array\n\n    dicom: pydicom.dataset.FileDataset\n        DICOM dataset\n\n    bits_allocated: int, str ('dataset') or None\n        Number of bits allocated\n\n    bits_stored: int, str ('dataset') or None\n        Number of bits stored\n\n    rescale_slope: int, str ('dataset') or None\n        Rescale slope for rescaling pixel values\n\n    rescale_intercept: int, str ('dataset') or None\n        Rescale intercept for rescaling pixel values\n\n    window_centers: list of int, str ('dataset') or None\n        List of window center values for windowing pixel values\n\n    window_widths: list of int, str ('dataset') or None\n        List of window width values for windowing pixel values\n\n    photometric_interpretation: str or None\n        Interpretation of the pixel data\n\n    max_pixel_value: int or None\n        Max pixel value used for inverting pixel values\n\n    Returns\n    -------\n    image: numpy.ndarray of shape (height, width)\n        Image array with adjusted pixel values\n    \"\"\"\n\n    image = shift_bits(image=image, dicom=dicom, bits_allocated=bits_allocated, bits_stored=bits_stored)\n    image = rescale_pixel_values(image=image, dicom=dicom, rescale_slope=rescale_slope, rescale_intercept=rescale_intercept)\n\n    image = np.stack([\n        window_pixel_values(image=np.copy(image), dicom=dicom, window_center=window_center, window_width=window_width)\n        for window_center, window_width in zip(window_centers, window_widths)\n    ], axis=-1)\n\n    image_min = image.min(axis=(0, 1))\n    image_max = image.max(axis=(0, 1))\n    image = (image - image_min) / (image_max - image_min + 1e-6)\n    image = invert_pixel_values(image=image, dicom=dicom, photometric_interpretation=photometric_interpretation, max_pixel_value=max_pixel_value)\n    image = (image * 255.0).astype(np.uint8)\n\n    return image\n\n\ndef adjust_pixel_spacing(image, dicom, current_pixel_spacing=None, new_pixel_spacing=(1.0, 1.0)):\n\n    \"\"\"\n    Adjust pixel values by shifting bits, windowing, rescaling and inverting\n\n    Parameters\n    ----------\n    image: numpy.ndarray of shape (height, width)\n        Image array\n\n    dicom: pydicom.dataset.FileDataset\n        DICOM dataset\n\n    current_pixel_spacing: tuple, str ('dataset') or None\n        Physical distance in the patient between the center of each pixel\n\n    new_pixel_spacing: tuple\n        Desired pixel spacing after resize operation\n\n    Returns\n    -------\n    image: numpy.ndarray of shape (height, width)\n        Image array with adjusted pixel spacing\n    \"\"\"\n\n    if current_pixel_spacing == 'dataset':\n        try:\n            current_pixel_spacing = dicom.PixelSpacing\n        except AttributeError:\n            current_pixel_spacing = None\n\n    if current_pixel_spacing is not None:\n        resize_factor = np.array(current_pixel_spacing) / np.array(new_pixel_spacing)\n        rounded_shape = np.round(image.shape[:2] * resize_factor)\n        resize_factor = rounded_shape / image.shape[:2]\n        image = cv2.resize(image, dsize=None, fx=resize_factor[1], fy=resize_factor[0], interpolation=cv2.INTER_NEAREST)\n\n    return image\n\n\ndef read_image(dicom_file_path, output_directory, pixel_spacing=None):\n\n    \"\"\"\n    Read DICOM file and process the image\n\n    Parameters\n    ----------\n    dicom_file_path: str\n        Path of the DICOM file\n\n    output_directory: pathlib.Path\n        Path of the directory image will be written to\n\n    pixel_spacing: tuple\n        Image pixel spacing after normalization\n    \"\"\"\n\n    dicom = pydicom.dcmread(str(dicom_file_path))\n    image = dicom.pixel_array\n    image = adjust_pixel_values(\n        image=image, dicom=dicom,\n        bits_allocated='dataset', bits_stored='dataset',\n        rescale_slope='dataset', rescale_intercept='dataset',\n        window_centers=['dataset'], window_widths=['dataset'],\n        photometric_interpretation='dataset', max_pixel_value=1\n    )\n\n    if pixel_spacing is not None:\n        image = dicom_utilities.adjust_pixel_spacing(\n            image=image,\n            dicom=dicom,\n            current_pixel_spacing='dataset',\n            new_pixel_spacing=pixel_spacing\n        )\n\n    return image\n\n\ndef get_largest_contour(image):\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    Returns\n    -------\n    bounding_box: list of shape (4)\n        Bounding box with x1, y1, x2, y2 values\n    \"\"\"\n\n    thresholded_image = cv2.threshold(image, 20, 255, cv2.THRESH_BINARY)[1]\n    contours, _ = cv2.findContours(thresholded_image, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)\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.99 * x1)\n        x2 = int(1.01 * x2)\n        y1 = int(0.99 * y1)\n        y2 = int(1.01 * y2)\n\n    bounding_box = [x1, y1, x2, y2]\n\n    return bounding_box\n","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:52:04.951549Z","iopub.execute_input":"2023-10-13T09:52:04.952121Z","iopub.status.idle":"2023-10-13T09:52:05.00023Z","shell.execute_reply.started":"2023-10-13T09:52:04.952042Z","shell.execute_reply":"2023-10-13T09:52:04.997495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. Models","metadata":{}},{"cell_type":"code","source":"class ClassificationHead(nn.Module):\n\n    def __init__(self, input_dimensions):\n\n        super(ClassificationHead, self).__init__()\n\n        self.bowel_head = nn.Linear(input_dimensions, 1, bias=True)\n        self.extravasation_head = nn.Linear(input_dimensions, 1, bias=True)\n        self.kidney_head = nn.Linear(input_dimensions, 3, bias=True)\n        self.liver_head = nn.Linear(input_dimensions, 3, bias=True)\n        self.spleen_head = nn.Linear(input_dimensions, 3, bias=True)\n\n    def forward(self, x):\n\n        bowel_output = self.bowel_head(x)\n        extravasation_output = self.extravasation_head(x)\n        kidney_output = self.kidney_head(x)\n        liver_output = self.liver_head(x)\n        spleen_output = self.spleen_head(x)\n\n        return bowel_output, extravasation_output, kidney_output, liver_output, spleen_output\n","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:52:05.612821Z","iopub.execute_input":"2023-10-13T09:52:05.613568Z","iopub.status.idle":"2023-10-13T09:52:05.621033Z","shell.execute_reply.started":"2023-10-13T09:52:05.613526Z","shell.execute_reply":"2023-10-13T09:52:05.619478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GeM(nn.Module):\n\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM, self).__init__()\n        self.p = nn.Parameter(torch.ones(1) * p)\n        self.eps = eps\n\n    def forward(self, x):\n        return self.gem(x, p=self.p, eps=self.eps)\n\n    def gem(self, x, p=3, eps=1e-6):\n        return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1. / p)\n\n    def __repr__(self):\n        return self.__class__.__name__ + \\\n            '(' + 'p=' + '{:.4f}'.format(self.p.data.tolist()[0]) + \\\n            ', ' + 'eps=' + str(self.eps) + ')'\n\n\nclass Attention(nn.Module):\n\n    def __init__(self, sequence_length, dimensions, bias=True):\n\n        super(Attention, self).__init__()\n\n        weight = torch.zeros(dimensions, 1)\n        nn.init.xavier_uniform_(weight)\n        self.weight = nn.Parameter(weight)\n        self.bias = bias\n        if bias:\n            self.b = nn.Parameter(torch.zeros(sequence_length))\n\n    def forward(self, x):\n\n        input_batch_size, input_sequence_length, input_dimensions = x.shape\n\n        eij = torch.mm(\n            x.contiguous().view(-1, input_dimensions),\n            self.weight\n        ).view(-1, input_sequence_length)\n\n        if self.bias:\n            eij = eij + self.b\n\n        eij = torch.tanh(eij)\n        a = torch.exp(eij)\n        a = a / torch.sum(a, 1, keepdim=True) + 1e-10\n        weighted_input = x * torch.unsqueeze(a, -1)\n        output = torch.sum(weighted_input, 1)\n\n        return output\n","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:52:06.237215Z","iopub.execute_input":"2023-10-13T09:52:06.237561Z","iopub.status.idle":"2023-10-13T09:52:06.247844Z","shell.execute_reply.started":"2023-10-13T09:52:06.237519Z","shell.execute_reply":"2023-10-13T09:52:06.246521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MILClassificationModel(nn.Module):\n\n    def __init__(self, model_name, pretrained, backbone_args, mil_pooling_type, feature_pooling_type, dropout_rate, freeze_parameters):\n\n        super(MILClassificationModel, 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        self.mil_pooling_type = mil_pooling_type\n        self.feature_pooling_type = feature_pooling_type\n        input_features = self.backbone.get_classifier().in_features\n        self.backbone.classifier = nn.Identity()\n\n        if self.feature_pooling_type == 'gem':\n            self.pooling = GeM()\n        elif self.feature_pooling_type == 'attention':\n            self.pooling = nn.Sequential(\n                nn.LayerNorm(normalized_shape=input_features),\n                Attention(sequence_length=49, dimensions=input_features)\n            )\n        else:\n            self.pooling = nn.Identity()\n\n        self.dropout = nn.Dropout(dropout_rate) if dropout_rate > 0 else nn.Identity()\n        self.head = ClassificationHead(input_dimensions=input_features)\n\n    def forward(self, x):\n\n        input_batch_size, input_channel, input_depth, input_height, input_width = x.shape\n        x = x.view(input_batch_size * input_depth, input_channel, input_height, input_width)\n        x = self.backbone.forward_features(x)\n        feature_batch_size, feature_channel, feature_height, feature_width = x.shape\n\n        if self.mil_pooling_type == 'avg':\n            x = x.contiguous().view(input_batch_size, input_depth, feature_channel, feature_height, feature_width)\n            x = torch.mean(x, dim=1)\n        elif self.mil_pooling_type == 'max':\n            x = x.contiguous().view(input_batch_size, input_depth, feature_channel, feature_height, feature_width)\n            x = torch.max(x, dim=1)[0]\n        elif self.mil_pooling_type == 'concat':\n            x = x.contiguous().view(input_batch_size, input_depth * feature_channel, feature_height, feature_width)\n        else:\n            raise ValueError(f'Invalid MIL pooling type {self.mil_pooling_type}')\n\n        if self.feature_pooling_type == 'avg':\n            x = F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.feature_pooling_type == 'max':\n            x = F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.feature_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        elif self.feature_pooling_type == 'gem':\n            x = self.pooling(x).view(x.size(0), -1)\n        elif self.feature_pooling_type == 'attention':\n            input_batch_size, feature_channel = x.shape[:2]\n            x = x.contiguous().view(input_batch_size, feature_channel, -1).permute(0, 2, 1)\n            x = self.pooling(x)\n        else:\n            raise ValueError(f'Invalid feature pooling type {self.feature_pooling_type}')\n\n        x = self.dropout(x)\n        bowel_output, extravasation_output, kidney_output, liver_output, spleen_output = self.head(x)\n\n        return bowel_output, extravasation_output, kidney_output, liver_output, spleen_output\n","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:52:25.91821Z","iopub.execute_input":"2023-10-13T09:52:25.918606Z","iopub.status.idle":"2023-10-13T09:52:25.933662Z","shell.execute_reply.started":"2023-10-13T09:52:25.918578Z","shell.execute_reply":"2023-10-13T09:52:25.932657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RNNClassificationModel(nn.Module):\n\n    def __init__(self, model_name, pretrained, backbone_args, feature_pooling_type, rnn_class, rnn_args, dropout_rate, freeze_parameters):\n\n        super(RNNClassificationModel, 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        self.feature_pooling_type = feature_pooling_type\n        input_features = self.backbone.get_classifier().in_features\n        self.backbone.classifier = nn.Identity()\n\n        if self.feature_pooling_type == 'gem':\n            self.pooling = GeM()\n        else:\n            self.pooling = nn.Identity()\n\n        self.rnn = getattr(nn, rnn_class)(input_size=input_features, **rnn_args)\n\n        self.dropout = nn.Dropout(dropout_rate) if dropout_rate > 0 else nn.Identity()\n        input_dimensions = rnn_args['hidden_size'] * (int(rnn_args['bidirectional']) + 1)\n        self.head = ClassificationHead(input_dimensions=input_dimensions)\n\n    def forward(self, x):\n\n        input_batch_size, input_channel, input_depth, input_height, input_width = x.shape\n        x = x.view(input_batch_size * input_depth, input_channel, input_height, input_width)\n        x = self.backbone.forward_features(x)\n\n        feature_batch_size, feature_channel, feature_height, feature_width = x.shape\n\n        if self.feature_pooling_type == 'avg':\n            x = F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.feature_pooling_type == 'max':\n            x = F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.feature_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        elif self.feature_pooling_type == 'gem':\n            x = self.pooling(x).view(x.size(0), -1)\n        else:\n            raise ValueError(f'Invalid feature pooling type {self.feature_pooling_type}')\n\n        x = x.contiguous().view(input_batch_size, input_depth, feature_channel)\n        x, _ = self.rnn(x)\n        x = torch.max(x, dim=1)[0]\n        x = self.dropout(x)\n        bowel_output, extravasation_output, kidney_output, liver_output, spleen_output = self.head(x)\n\n        return bowel_output, extravasation_output, kidney_output, liver_output, spleen_output\n","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:52:27.270449Z","iopub.execute_input":"2023-10-13T09:52:27.271145Z","iopub.status.idle":"2023-10-13T09:52:27.282702Z","shell.execute_reply.started":"2023-10-13T09:52:27.271114Z","shell.execute_reply":"2023-10-13T09:52:27.281381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_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: str\n        Path of the model directory\n\n    model_file_names: str\n        Name of the model weights files\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    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 tqdm(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\n    return models, config\n","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:52:29.360181Z","iopub.execute_input":"2023-10-13T09:52:29.360511Z","iopub.status.idle":"2023-10-13T09:52:29.366819Z","shell.execute_reply.started":"2023-10-13T09:52:29.360485Z","shell.execute_reply":"2023-10-13T09:52:29.365775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mil_efficientnetb0_models, mil_efficientnetb0_config = load_model(\n    model_directory=external_dataset / 'mil_efficientnetb0_3d_1w_contour_cropped_96x256x256',\n    model_file_names=[\n        'model_fold1_best.pt',\n        'model_fold2_best.pt',\n        'model_fold3_best.pt',\n        'model_fold4_best.pt',\n        'model_fold5_best.pt',\n    ],\n    device=torch.device('cuda')\n)\n\nmil_densenet121_models, mil_densenet121_config = load_model(\n    model_directory=external_dataset / 'mil_densenet121_3d_1w_contour_cropped_96x256x256',\n    model_file_names=[\n        'model_fold1_best.pt',\n        'model_fold2_best.pt',\n        'model_fold3_best.pt',\n        'model_fold4_best.pt',\n        'model_fold5_best.pt',\n    ],\n    device=torch.device('cuda')\n)\n\nlstm_efficientnetb0_models, lstm_efficientnetb0_config = load_model(\n    model_directory=external_dataset / 'lstm_efficientnetb0_3d_1w_contour_cropped_96x256x256',\n    model_file_names=[\n        'model_fold1_best.pt',\n        'model_fold2_best.pt',\n        'model_fold3_best.pt',\n        'model_fold4_best.pt',\n        'model_fold5_best.pt',\n    ],\n    device=torch.device('cuda')\n)\n\nlstm_efficientnetv2t_models, lstm_efficientnetv2t_config = load_model(\n    model_directory=external_dataset / 'lstm_efficientnetv2t_3d_1w_contour_cropped_96x256x256',\n    model_file_names=[\n        'model_fold1_best.pt',\n        'model_fold2_best.pt',\n        'model_fold3_best.pt',\n        'model_fold4_best.pt',\n        'model_fold5_best.pt',\n    ],\n    device=torch.device('cuda')\n)","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:52:46.763238Z","iopub.execute_input":"2023-10-13T09:52:46.763588Z","iopub.status.idle":"2023-10-13T09:53:07.60073Z","shell.execute_reply.started":"2023-10-13T09:52:46.76356Z","shell.execute_reply":"2023-10-13T09:53:07.599682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4. Transforms","metadata":{}},{"cell_type":"code","source":"def get_3d_classification_transforms(**transform_parameters):\n\n    \"\"\"\n    Get transforms for classification dataset\n\n    Parameters\n    ----------\n    transform_parameters: dict\n        Dictionary of transform parameters\n\n    Returns\n    -------\n    transforms: dict\n        Transforms for training, validation and test sets\n    \"\"\"\n\n    training_transforms = T.Compose([\n        T.EnsureChannelFirst(channel_dim=0),\n        T.RandFlip(spatial_axis=0, prob=transform_parameters['random_z_flip_probability']),\n        T.RandFlip(spatial_axis=1, prob=transform_parameters['random_x_flip_probability']),\n        T.RandFlip(spatial_axis=2, prob=transform_parameters['random_y_flip_probability']),\n        T.RandRotate90(spatial_axes=(1, 2), max_k=3, prob=transform_parameters['random_axial_rotate_90_probability']),\n        T.RandRotate(\n            range_x=transform_parameters['random_rotate_range_x'],\n            range_y=transform_parameters['random_rotate_range_y'],\n            range_z=transform_parameters['random_rotate_range_z'],\n            prob=transform_parameters['random_rotate_probability']\n        ),\n        T.OneOf([\n            T.RandHistogramShift(num_control_points=transform_parameters['random_histogram_shift_num_control_points'], prob=transform_parameters['random_histogram_shift_probability']),\n            T.RandAdjustContrast(gamma=transform_parameters['random_contrast_gamma'], prob=transform_parameters['random_contrast_probability'])\n        ], weights=(0.5, 0.5)),\n        T.RandSpatialCrop(roi_size=transform_parameters['crop_roi_size'], max_roi_size=None, random_center=True, random_size=False),\n        T.RandCoarseDropout(\n            holes=transform_parameters['cutout_holes'],\n            spatial_size=transform_parameters['cutout_spatial_size'],\n            dropout_holes=True,\n            fill_value=0,\n            max_holes=transform_parameters['cutout_max_holes'],\n            max_spatial_size=transform_parameters['max_spatial_size'],\n            prob=transform_parameters['cutout_probability']\n        ),\n        T.ToTensor(dtype=torch.float32, track_meta=False)\n    ])\n\n    inference_transforms = T.Compose([\n        T.Resize(spatial_size=(96, 256, 256)),\n        T.CenterSpatialCrop(roi_size=(-1, 224, 224))\n    ])\n\n    classification_transforms = {'training': training_transforms, 'inference': inference_transforms}\n    return classification_transforms\n","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:53:27.709058Z","iopub.execute_input":"2023-10-13T09:53:27.70941Z","iopub.status.idle":"2023-10-13T09:53:27.719377Z","shell.execute_reply.started":"2023-10-13T09:53:27.709384Z","shell.execute_reply":"2023-10-13T09:53:27.718373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5. Inference","metadata":{}},{"cell_type":"code","source":"device = torch.device('cuda')\namp = True\ninference_transforms = T.Compose([\n    T.Resize(spatial_size=(96, 256, 256)),\n    T.CenterSpatialCrop(roi_size=(-1, 224, 224))\n])\ntta = True\ntta_flip_dimensions = [(2, 3, 4), (2, 3), (2, 4), (3, 4)]","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:54:12.336755Z","iopub.execute_input":"2023-10-13T09:54:12.337136Z","iopub.status.idle":"2023-10-13T09:54:12.343828Z","shell.execute_reply.started":"2023-10-13T09:54:12.337106Z","shell.execute_reply":"2023-10-13T09:54:12.342624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bowel_predictions = []\nextravasation_predictions = []\nkidney_predictions = []\nliver_predictions = []\nspleen_predictions = []\npatient_ids_predictions = []\nscan_ids_predictions = []\n\ndicom_dataset_directory = competition_dataset / 'test_images'\npatient_ids = sorted(os.listdir(dicom_dataset_directory), key=lambda filename: int(filename))\n\nfor patient_id in tqdm(patient_ids):\n\n    patient_directory = dicom_dataset_directory / str(patient_id)\n    patient_scans = sorted(os.listdir(patient_directory), key=lambda filename: int(filename))\n\n    for scan_id in patient_scans:\n\n        scan_directory = patient_directory / str(scan_id)\n        file_names = sorted(os.listdir(scan_directory), key=lambda x: int(str(x).split('.')[0]))\n                        \n        if patient_id == '3124' and scan_id == '5842':\n            # Remove corrupt DICOM file\n            file_names.remove('514.dcm')\n            \n        z_positions = []\n        patient_positions = []\n        scan = []\n        \n        for file_idx, file_name in enumerate(file_names, start=1):\n            \n            dicom = pydicom.dcmread(str(scan_directory / file_name))\n\n            try:\n                patient_position = dicom.PatientPosition\n            except AttributeError:\n                patient_position = 'FFS'\n            \n            patient_positions.append(patient_position)\n            \n            try:\n                z_position = float(dicom.ImagePositionPatient[-1])\n            except AttributeError:\n                z_position = file_idx * -1\n                \n            z_positions.append(z_position)\n            \n            image = dicom.pixel_array\n            image = adjust_pixel_values(\n                image=image, dicom=dicom,\n                bits_allocated='dataset', bits_stored='dataset',\n                rescale_slope='dataset', rescale_intercept='dataset',\n                window_centers=['dataset'], window_widths=['dataset'],\n                photometric_interpretation='dataset', max_pixel_value=1\n            )\n            image = np.squeeze(image, axis=-1)\n            scan.append(image)\n            \n        scan = np.array(scan)\n            \n        # Sort CT scan slices by head to feet\n        sorting_idx_z = np.argsort(z_positions)[::-1]\n        scan = scan[sorting_idx_z]\n        \n        patient_position = pd.Series(patient_positions).value_counts().index[0]\n        if patient_position is not None:\n            if patient_position == 'HFS':\n                # Flip x-axis if patient position is head first\n                scan = np.flip(scan, axis=2)\n\n        # Find partial slices by calculating sum of all zero vertical lines\n        if scan.shape[0] != 1:\n            scan_all_zero_vertical_line_transitions = np.diff(np.all(scan == 0, axis=1).sum(axis=1))\n            # Heuristically select high and low transitions on z-axis and drop them\n            slices_with_all_zero_vertical_lines = (scan_all_zero_vertical_line_transitions > 5) | (scan_all_zero_vertical_line_transitions < -5)\n            slices_with_all_zero_vertical_lines = np.append(slices_with_all_zero_vertical_lines, slices_with_all_zero_vertical_lines[-1])\n            scan = scan[~slices_with_all_zero_vertical_lines]\n            del scan_all_zero_vertical_line_transitions, slices_with_all_zero_vertical_lines\n        \n        # Crop the largest contour\n        largest_contour_bounding_boxes = np.array([get_largest_contour(image) for image in scan])\n        largest_contour_bounding_box = [\n            int(largest_contour_bounding_boxes[:, 0].min()),\n            int(largest_contour_bounding_boxes[:, 1].min()),\n            int(largest_contour_bounding_boxes[:, 2].max()),\n            int(largest_contour_bounding_boxes[:, 3].max()),\n        ]\n        scan = scan[\n            :,\n            largest_contour_bounding_box[1]:largest_contour_bounding_box[3] + 1,\n            largest_contour_bounding_box[0]:largest_contour_bounding_box[2] + 1,\n        ]\n        \n        # Crop non-zero slices along xz, yz and xy planes\n        mmin = np.array((scan > 0).nonzero()).min(axis=1)\n        mmax = np.array((scan > 0).nonzero()).max(axis=1)\n        scan = scan[\n            mmin[0]:mmax[0] + 1,\n            mmin[1]:mmax[1] + 1,\n            mmin[2]:mmax[2] + 1,\n        ]\n        \n        inputs = inference_transforms(torch.from_numpy(np.expand_dims(scan.copy(), axis=0)))\n        inputs = torch.unsqueeze(inputs, dim=0)\n        inputs /= 255.\n        inputs = inputs.to(device)\n        \n        n_models = 4\n        bowel_batch_predictions = torch.zeros(n_models, inputs.shape[0], 1)\n        extravasation_batch_predictions = torch.zeros(n_models, inputs.shape[0], 1)\n        kidney_batch_predictions = torch.zeros(n_models, inputs.shape[0], 3)\n        liver_batch_predictions = torch.zeros(n_models, inputs.shape[0], 3)\n        spleen_batch_predictions = torch.zeros(n_models, inputs.shape[0], 3)\n        \n        for model_idx, models in enumerate([mil_efficientnetb0_models, mil_densenet121_models, lstm_efficientnetb0_models, lstm_efficientnetv2t_models]):\n            for model in models.values():\n                with torch.no_grad():\n                    if amp:\n                        with torch.cuda.amp.autocast():\n                            bowel_outputs, extravasation_outputs, kidney_outputs, liver_outputs, spleen_outputs = model(inputs.half())\n                    else:\n                        bowel_outputs, extravasation_outputs, kidney_outputs, liver_outputs, spleen_outputs = model(inputs)\n\n                bowel_outputs = bowel_outputs.cpu()\n                extravasation_outputs = extravasation_outputs.cpu()\n                kidney_outputs = kidney_outputs.cpu()\n                liver_outputs = liver_outputs.cpu()\n                spleen_outputs = spleen_outputs.cpu()\n\n                if tta:\n\n                    tta_bowel_outputs = []\n                    tta_extravasation_outputs = []\n                    tta_kidney_outputs = []\n                    tta_liver_outputs = []\n                    tta_spleen_outputs = []\n\n                    for dimensions in tta_flip_dimensions:\n\n                        augmented_inputs = torch.flip(inputs, dims=dimensions).to(device)\n\n                        with torch.no_grad():\n                            augmented_bowel_outputs, augmented_extravasation_outputs, augmented_kidney_outputs, augmented_liver_outputs, augmented_spleen_outputs = model(augmented_inputs)\n\n                        tta_bowel_outputs.append(augmented_bowel_outputs.cpu())\n                        tta_extravasation_outputs.append(augmented_extravasation_outputs.cpu())\n                        tta_kidney_outputs.append(augmented_kidney_outputs.cpu())\n                        tta_liver_outputs.append(augmented_liver_outputs.cpu())\n                        tta_spleen_outputs.append(augmented_spleen_outputs.cpu())\n\n                    bowel_outputs = torch.stack(([bowel_outputs] + tta_bowel_outputs), dim=-1)\n                    extravasation_outputs = torch.stack(([extravasation_outputs] + tta_extravasation_outputs), dim=-1)\n                    kidney_outputs = torch.stack(([kidney_outputs] + tta_kidney_outputs), dim=-1)\n                    liver_outputs = torch.stack(([liver_outputs] + tta_liver_outputs), dim=-1)\n                    spleen_outputs = torch.stack(([spleen_outputs] + tta_spleen_outputs), dim=-1)\n\n                    bowel_outputs = torch.mean(bowel_outputs, dim=-1)\n                    extravasation_outputs = torch.mean(extravasation_outputs, dim=-1)\n                    kidney_outputs = torch.mean(kidney_outputs, dim=-1)\n                    liver_outputs = torch.mean(liver_outputs, dim=-1)\n                    spleen_outputs = torch.mean(spleen_outputs, dim=-1)\n                \n                bowel_batch_predictions[model_idx] += bowel_outputs / len(models)\n                extravasation_batch_predictions[model_idx] += extravasation_outputs / len(models)\n                kidney_batch_predictions[model_idx] += kidney_outputs / len(models)\n                liver_batch_predictions[model_idx] += liver_outputs / len(models)\n                spleen_batch_predictions[model_idx] += spleen_outputs / len(models)\n                \n        bowel_predictions += [bowel_batch_predictions]\n        extravasation_predictions += [extravasation_batch_predictions]\n        kidney_predictions += [kidney_batch_predictions]\n        liver_predictions += [liver_batch_predictions]\n        spleen_predictions += [spleen_batch_predictions]\n        \n        patient_ids_predictions.append(patient_id)\n        scan_ids_predictions.append(scan_id)\n\nbowel_predictions = torch.sigmoid(torch.stack(bowel_predictions, dim=0)).numpy()\nextravasation_predictions = torch.sigmoid(torch.stack(extravasation_predictions, dim=0)).numpy()\nkidney_predictions = torch.softmax(torch.stack(kidney_predictions, dim=0), dim=-1).numpy()\nliver_predictions = torch.softmax(torch.stack(liver_predictions, dim=0), dim=-1).numpy()\nspleen_predictions = torch.softmax(torch.stack(spleen_predictions, dim=0), dim=-1).numpy()\npatient_ids_predictions = np.array(patient_ids_predictions).reshape(-1, 1)\nscan_ids_predictions = np.array(scan_ids_predictions).reshape(-1, 1)","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:54:18.386939Z","iopub.execute_input":"2023-10-13T09:54:18.387337Z","iopub.status.idle":"2023-10-13T09:54:51.749738Z","shell.execute_reply.started":"2023-10-13T09:54:18.387309Z","shell.execute_reply":"2023-10-13T09:54:51.748736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_mil_efficientnetb0_predictions = pd.DataFrame(np.hstack([\n    patient_ids_predictions,\n    scan_ids_predictions,\n    bowel_predictions[:, 0, :, :].reshape(-1, 1),\n    extravasation_predictions[:, 0, :, :].reshape(-1, 1),\n    kidney_predictions[:, 0, :, :].reshape(-1, 3),\n    liver_predictions[:, 0, :, :].reshape(-1, 3),\n    spleen_predictions[:, 0, :, :].reshape(-1, 3)\n]), columns=[\n    'patient_id', 'scan_id',\n    'bowel_injury', 'extravasation_injury',\n    'kidney_healthy', 'kidney_low', 'kidney_high',\n    'liver_healthy', 'liver_low', 'liver_high',\n    'spleen_healthy', 'spleen_low', 'spleen_high',\n])\n\ndf_mil_efficientnetb0_predictions[df_mil_efficientnetb0_predictions.columns[:2]] = df_mil_efficientnetb0_predictions[df_mil_efficientnetb0_predictions.columns[:2]].astype(np.int64)\ndf_mil_efficientnetb0_predictions[df_mil_efficientnetb0_predictions.columns[2:]] = df_mil_efficientnetb0_predictions[df_mil_efficientnetb0_predictions.columns[2:]].astype(np.float32)\n\ndf_mil_efficientnetb0_predictions","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:55:02.411412Z","iopub.execute_input":"2023-10-13T09:55:02.411775Z","iopub.status.idle":"2023-10-13T09:55:02.449086Z","shell.execute_reply.started":"2023-10-13T09:55:02.411744Z","shell.execute_reply":"2023-10-13T09:55:02.44788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_mil_densenet121_predictions = pd.DataFrame(np.hstack([\n    patient_ids_predictions,\n    scan_ids_predictions,\n    bowel_predictions[:, 1, :, :].reshape(-1, 1),\n    extravasation_predictions[:, 1, :, :].reshape(-1, 1),\n    kidney_predictions[:, 1, :, :].reshape(-1, 3),\n    liver_predictions[:, 1, :, :].reshape(-1, 3),\n    spleen_predictions[:, 1, :, :].reshape(-1, 3)\n]), columns=[\n    'patient_id', 'scan_id',\n    'bowel_injury', 'extravasation_injury',\n    'kidney_healthy', 'kidney_low', 'kidney_high',\n    'liver_healthy', 'liver_low', 'liver_high',\n    'spleen_healthy', 'spleen_low', 'spleen_high',\n])\n\ndf_mil_densenet121_predictions[df_mil_densenet121_predictions.columns[:2]] = df_mil_densenet121_predictions[df_mil_densenet121_predictions.columns[:2]].astype(np.int64)\ndf_mil_densenet121_predictions[df_mil_densenet121_predictions.columns[2:]] = df_mil_densenet121_predictions[df_mil_densenet121_predictions.columns[2:]].astype(np.float32)\n\ndf_mil_densenet121_predictions","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:55:03.245916Z","iopub.execute_input":"2023-10-13T09:55:03.247012Z","iopub.status.idle":"2023-10-13T09:55:03.274753Z","shell.execute_reply.started":"2023-10-13T09:55:03.246971Z","shell.execute_reply":"2023-10-13T09:55:03.27363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_lstm_efficientnetb0_predictions = pd.DataFrame(np.hstack([\n    patient_ids_predictions,\n    scan_ids_predictions,\n    bowel_predictions[:, 2, :, :].reshape(-1, 1),\n    extravasation_predictions[:, 2, :, :].reshape(-1, 1),\n    kidney_predictions[:, 2, :, :].reshape(-1, 3),\n    liver_predictions[:, 2, :, :].reshape(-1, 3),\n    spleen_predictions[:, 2, :, :].reshape(-1, 3)\n]), columns=[\n    'patient_id', 'scan_id',\n    'bowel_injury', 'extravasation_injury',\n    'kidney_healthy', 'kidney_low', 'kidney_high',\n    'liver_healthy', 'liver_low', 'liver_high',\n    'spleen_healthy', 'spleen_low', 'spleen_high',\n])\n\ndf_lstm_efficientnetb0_predictions[df_lstm_efficientnetb0_predictions.columns[:2]] = df_lstm_efficientnetb0_predictions[df_lstm_efficientnetb0_predictions.columns[:2]].astype(np.int64)\ndf_lstm_efficientnetb0_predictions[df_lstm_efficientnetb0_predictions.columns[2:]] = df_lstm_efficientnetb0_predictions[df_lstm_efficientnetb0_predictions.columns[2:]].astype(np.float32)\n\ndf_lstm_efficientnetb0_predictions","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:55:05.286613Z","iopub.execute_input":"2023-10-13T09:55:05.286985Z","iopub.status.idle":"2023-10-13T09:55:05.315455Z","shell.execute_reply.started":"2023-10-13T09:55:05.286948Z","shell.execute_reply":"2023-10-13T09:55:05.314083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_lstm_efficientnetv2t_predictions = pd.DataFrame(np.hstack([\n    patient_ids_predictions,\n    scan_ids_predictions,\n    bowel_predictions[:, 3, :, :].reshape(-1, 1),\n    extravasation_predictions[:, 3, :, :].reshape(-1, 1),\n    kidney_predictions[:, 3, :, :].reshape(-1, 3),\n    liver_predictions[:, 3, :, :].reshape(-1, 3),\n    spleen_predictions[:, 3, :, :].reshape(-1, 3)\n]), columns=[\n    'patient_id', 'scan_id',\n    'bowel_injury', 'extravasation_injury',\n    'kidney_healthy', 'kidney_low', 'kidney_high',\n    'liver_healthy', 'liver_low', 'liver_high',\n    'spleen_healthy', 'spleen_low', 'spleen_high',\n])\n\ndf_lstm_efficientnetv2t_predictions[df_lstm_efficientnetv2t_predictions.columns[:2]] = df_lstm_efficientnetv2t_predictions[df_lstm_efficientnetv2t_predictions.columns[:2]].astype(np.int64)\ndf_lstm_efficientnetv2t_predictions[df_lstm_efficientnetv2t_predictions.columns[2:]] = df_lstm_efficientnetv2t_predictions[df_lstm_efficientnetv2t_predictions.columns[2:]].astype(np.float32)\n\ndf_lstm_efficientnetv2t_predictions","metadata":{"execution":{"iopub.status.busy":"2023-10-13T09:55:44.26081Z","iopub.execute_input":"2023-10-13T09:55:44.261391Z","iopub.status.idle":"2023-10-13T09:55:44.300784Z","shell.execute_reply.started":"2023-10-13T09:55:44.26135Z","shell.execute_reply":"2023-10-13T09:55:44.299605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 6. Post processing","metadata":{}},{"cell_type":"code","source":"df_predictions = pd.DataFrame(columns=[\n    'patient_id', 'scan_id',\n    'bowel_injury', 'extravasation_injury',\n    'kidney_healthy', 'kidney_low', 'kidney_high',\n    'liver_healthy', 'liver_low', 'liver_high',\n    'spleen_healthy', 'spleen_low', 'spleen_high',\n])\n\ndf_predictions['patient_id'] = patient_ids_predictions.reshape(-1).astype(int)\ndf_predictions['scan_id'] = scan_ids_predictions.reshape(-1).astype(int)\n\nmil_efficientnetb0_bowel_weight = 0.45\nmil_densenet121_bowel_weight = 0.25\nlstm_efficientnetb0_bowel_weight = 0.15\nlstm_efficientnetv2t_bowel_weight = 0.15\n\ndf_predictions['bowel_injury'] = (df_mil_efficientnetb0_predictions['bowel_injury'] * mil_efficientnetb0_bowel_weight) +\\\n                                 (df_mil_densenet121_predictions['bowel_injury'] * mil_densenet121_bowel_weight) +\\\n                                 (df_lstm_efficientnetb0_predictions['bowel_injury'] * lstm_efficientnetb0_bowel_weight) +\\\n                                 (df_lstm_efficientnetv2t_predictions['bowel_injury'] * lstm_efficientnetv2t_bowel_weight)\n\nmil_efficientnetb0_extravasation_weight = 0.3\nmil_densenet121_extravasation_weight = 0.3\nlstm_efficientnetb0_extravasation_weight = 0.3\nlstm_efficientnetv2t_extravasation_weight = 0.1\n\ndf_predictions['extravasation_injury'] = (df_mil_efficientnetb0_predictions['extravasation_injury'] * mil_efficientnetb0_extravasation_weight) +\\\n                                         (df_mil_densenet121_predictions['extravasation_injury'] * mil_densenet121_extravasation_weight) +\\\n                                         (df_lstm_efficientnetb0_predictions['extravasation_injury'] * lstm_efficientnetb0_extravasation_weight) +\\\n                                         (df_lstm_efficientnetv2t_predictions['extravasation_injury'] * lstm_efficientnetv2t_extravasation_weight)\n\nmil_efficientnetb0_kidney_weight = 0.25\nmil_densenet121_kidney_weight = 0.25\nlstm_efficientnetb0_kidney_weight = 0.25\nlstm_efficientnetv2t_kidney_weight = 0.25\n\ndf_predictions['kidney_healthy'] = (df_mil_efficientnetb0_predictions['kidney_healthy'] * mil_efficientnetb0_kidney_weight) +\\\n                                   (df_mil_densenet121_predictions['kidney_healthy'] * mil_densenet121_kidney_weight) +\\\n                                   (df_lstm_efficientnetb0_predictions['kidney_healthy'] * lstm_efficientnetb0_kidney_weight) +\\\n                                   (df_lstm_efficientnetv2t_predictions['kidney_healthy'] * lstm_efficientnetv2t_kidney_weight)\n\ndf_predictions['kidney_low'] = (df_mil_efficientnetb0_predictions['kidney_low'] * mil_efficientnetb0_kidney_weight) +\\\n                               (df_mil_densenet121_predictions['kidney_low'] * mil_densenet121_kidney_weight) +\\\n                               (df_lstm_efficientnetb0_predictions['kidney_low'] * lstm_efficientnetb0_kidney_weight) +\\\n                               (df_lstm_efficientnetv2t_predictions['kidney_low'] * lstm_efficientnetv2t_kidney_weight)\n\ndf_predictions['kidney_high'] = (df_mil_efficientnetb0_predictions['kidney_high'] * mil_efficientnetb0_kidney_weight) +\\\n                                (df_mil_densenet121_predictions['kidney_high'] * mil_densenet121_kidney_weight) +\\\n                                (df_lstm_efficientnetb0_predictions['kidney_high'] * lstm_efficientnetb0_kidney_weight) +\\\n                                (df_lstm_efficientnetv2t_predictions['kidney_high'] * lstm_efficientnetv2t_kidney_weight)\n\nmil_efficientnetb0_liver_weight = 0.25\nmil_densenet121_liver_weight = 0.25\nlstm_efficientnetb0_liver_weight = 0.25\nlstm_efficientnetv2t_liver_weight = 0.25\n\ndf_predictions['liver_healthy'] = (df_mil_efficientnetb0_predictions['liver_healthy'] * mil_efficientnetb0_liver_weight) +\\\n                                  (df_mil_densenet121_predictions['liver_healthy'] * mil_densenet121_liver_weight) +\\\n                                  (df_lstm_efficientnetb0_predictions['liver_healthy'] * lstm_efficientnetb0_liver_weight) +\\\n                                  (df_lstm_efficientnetv2t_predictions['liver_healthy'] * lstm_efficientnetv2t_liver_weight)\n\ndf_predictions['liver_low'] = (df_mil_efficientnetb0_predictions['liver_low'] * mil_efficientnetb0_liver_weight) +\\\n                              (df_mil_densenet121_predictions['liver_low'] * mil_densenet121_liver_weight) +\\\n                              (df_lstm_efficientnetb0_predictions['liver_low'] * lstm_efficientnetb0_liver_weight) +\\\n                              (df_lstm_efficientnetv2t_predictions['liver_low'] * lstm_efficientnetv2t_liver_weight)\n\ndf_predictions['liver_high'] = (df_mil_efficientnetb0_predictions['liver_high'] * mil_efficientnetb0_liver_weight) +\\\n                               (df_mil_densenet121_predictions['liver_high'] * mil_densenet121_liver_weight) +\\\n                               (df_lstm_efficientnetb0_predictions['liver_high'] * lstm_efficientnetb0_liver_weight) +\\\n                               (df_lstm_efficientnetv2t_predictions['liver_high'] * lstm_efficientnetv2t_liver_weight)\n\nmil_efficientnetb0_spleen_weight = 0.25\nmil_densenet121_spleen_weight = 0.25\nlstm_efficientnetb0_spleen_weight = 0.25\nlstm_efficientnetv2t_spleen_weight = 0.25\n\ndf_predictions['spleen_healthy'] = (df_mil_efficientnetb0_predictions['spleen_healthy'] * mil_efficientnetb0_spleen_weight) +\\\n                                   (df_mil_densenet121_predictions['spleen_healthy'] * mil_densenet121_spleen_weight) +\\\n                                   (df_lstm_efficientnetb0_predictions['spleen_healthy'] * lstm_efficientnetb0_spleen_weight) +\\\n                                   (df_lstm_efficientnetv2t_predictions['spleen_healthy'] * lstm_efficientnetv2t_spleen_weight)\n\ndf_predictions['spleen_low'] = (df_mil_efficientnetb0_predictions['spleen_low'] * mil_efficientnetb0_spleen_weight) +\\\n                               (df_mil_densenet121_predictions['spleen_low'] * mil_densenet121_spleen_weight) +\\\n                               (df_lstm_efficientnetb0_predictions['spleen_low'] * lstm_efficientnetb0_spleen_weight) +\\\n                               (df_lstm_efficientnetv2t_predictions['spleen_low'] * lstm_efficientnetv2t_spleen_weight)\n\ndf_predictions['spleen_high'] = (df_mil_efficientnetb0_predictions['spleen_high'] * mil_efficientnetb0_spleen_weight) +\\\n                                (df_mil_densenet121_predictions['spleen_high'] * mil_densenet121_spleen_weight) +\\\n                                (df_lstm_efficientnetb0_predictions['spleen_high'] * lstm_efficientnetb0_spleen_weight) +\\\n                                (df_lstm_efficientnetv2t_predictions['spleen_high'] * lstm_efficientnetv2t_spleen_weight)\n\ndf_predictions","metadata":{"execution":{"iopub.status.busy":"2023-10-13T10:01:16.442911Z","iopub.execute_input":"2023-10-13T10:01:16.443299Z","iopub.status.idle":"2023-10-13T10:01:16.494657Z","shell.execute_reply.started":"2023-10-13T10:01:16.443271Z","shell.execute_reply":"2023-10-13T10:01:16.493458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_predictions['bowel_healthy'] = 1 - df_predictions['bowel_injury']\ndf_predictions['bowel_injury'] *= 1.\n\ndf_predictions['extravasation_healthy'] = 1 - df_predictions['extravasation_injury']\ndf_predictions['extravasation_injury'] *= 1.4\n\ndf_predictions['kidney_low'] *= 1.1\ndf_predictions['kidney_high'] *= 1.1\n\ndf_predictions['liver_low'] *= 1.3\ndf_predictions['liver_high'] *= 1.3\n\ndf_predictions['spleen_low'] *= 1.75\ndf_predictions['spleen_high'] *= 1.75\n\ndf_predictions","metadata":{"execution":{"iopub.status.busy":"2023-10-13T10:01:32.350098Z","iopub.execute_input":"2023-10-13T10:01:32.350454Z","iopub.status.idle":"2023-10-13T10:01:32.374118Z","shell.execute_reply.started":"2023-10-13T10:01:32.350426Z","shell.execute_reply":"2023-10-13T10:01:32.372908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"prediction_columns = [\n    'bowel_healthy', 'bowel_injury',\n    'extravasation_healthy', 'extravasation_injury',\n    'kidney_healthy', 'kidney_low', 'kidney_high',\n    'liver_healthy', 'liver_low', 'liver_high',\n    'spleen_healthy', 'spleen_low', 'spleen_high'\n]\ndf_predictions = df_predictions.groupby('patient_id')[prediction_columns].max().reset_index()\n\ndf_predictions","metadata":{"execution":{"iopub.status.busy":"2023-10-13T10:01:33.087949Z","iopub.execute_input":"2023-10-13T10:01:33.088326Z","iopub.status.idle":"2023-10-13T10:01:33.117918Z","shell.execute_reply.started":"2023-10-13T10:01:33.088298Z","shell.execute_reply":"2023-10-13T10:01:33.116743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 7. Submission","metadata":{}},{"cell_type":"code","source":"df = df[['patient_id']].merge(df_predictions, on='patient_id', how='left')\ndf","metadata":{"execution":{"iopub.status.busy":"2023-10-13T10:01:34.230849Z","iopub.execute_input":"2023-10-13T10:01:34.231255Z","iopub.status.idle":"2023-10-13T10:01:34.257225Z","shell.execute_reply.started":"2023-10-13T10:01:34.231227Z","shell.execute_reply":"2023-10-13T10:01:34.255968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-10-12T07:45:13.915821Z","iopub.execute_input":"2023-10-12T07:45:13.916144Z","iopub.status.idle":"2023-10-12T07:45:13.923107Z","shell.execute_reply.started":"2023-10-12T07:45:13.916118Z","shell.execute_reply":"2023-10-12T07:45:13.922197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}