{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":36363,"databundleVersionId":4050810,"sourceType":"competition"},{"sourceId":52254,"databundleVersionId":6863140,"sourceType":"competition"},{"sourceId":2768643,"sourceType":"datasetVersion","datasetId":1688447},{"sourceId":4407134,"sourceType":"datasetVersion","datasetId":2556595},{"sourceId":4454029,"sourceType":"datasetVersion","datasetId":2607864},{"sourceId":7980113,"sourceType":"datasetVersion","datasetId":4696701},{"sourceId":7980151,"sourceType":"datasetVersion","datasetId":4696727},{"sourceId":7980250,"sourceType":"datasetVersion","datasetId":4696803},{"sourceId":7980314,"sourceType":"datasetVersion","datasetId":4696793},{"sourceId":7980348,"sourceType":"datasetVersion","datasetId":4696880},{"sourceId":7980767,"sourceType":"datasetVersion","datasetId":4697191},{"sourceId":7981043,"sourceType":"datasetVersion","datasetId":4697395},{"sourceId":7981093,"sourceType":"datasetVersion","datasetId":4697438},{"sourceId":7981182,"sourceType":"datasetVersion","datasetId":4697501},{"sourceId":7981226,"sourceType":"datasetVersion","datasetId":4697537},{"sourceId":104036025,"sourceType":"kernelVersion"}],"dockerImageVersionId":30262,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install /kaggle/input/rsna-2022-whl/{pydicom-2.3.0-py3-none-any.whl,pylibjpeg-1.4.0-py3-none-any.whl,python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl}\n!pip install /kaggle/input/rsna-2022-whl/{torch-1.12.1-cp37-cp37m-manylinux1_x86_64.whl,torchvision-0.13.1-cp37-cp37m-manylinux1_x86_64.whl}","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-03-31T09:25:49.943272Z","iopub.execute_input":"2024-03-31T09:25:49.94409Z","iopub.status.idle":"2024-03-31T09:27:23.149506Z","shell.execute_reply.started":"2024-03-31T09:25:49.943963Z","shell.execute_reply":"2024-03-31T09:27:23.148335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install /kaggle/input/rsna-weights/timm-0.5.4-py3-none-any.whl\n!pip install /kaggle/input/rsna-weights/tifffile-2022.8.8-py3-none-any.whl\n!pip install /kaggle/input/rsna-weights/einops-0.5.0-py3-none-any.whl\n\n","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:27:23.152142Z","iopub.execute_input":"2024-03-31T09:27:23.152549Z","iopub.status.idle":"2024-03-31T09:27:52.031529Z","shell.execute_reply.started":"2024-03-31T09:27:23.152508Z","shell.execute_reply":"2024-03-31T09:27:52.03016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nos.listdir('/kaggle/input/mmdetection-2-17-offline')\n\n!pip install /kaggle/input/mmdetection-2-17-offline/mmcv_full-1.3.14-cp37-cp37m-linux_x86_64.whl --no-deps\n!pip install /kaggle/input/mmdetection-2-17-offline/pycocotools-2.0.2-cp37-cp37m-linux_x86_64.whl --no-deps\n!pip install /kaggle/input/mmdetection-2-17-offline/terminaltables-3.1.0-py3-none-any.whl --no-deps\n!pip install /kaggle/input/mmdetection-2-17-offline/pytest_runner-5.3.1-py3-none-any.whl --no-deps\n!pip install /kaggle/input/mmdetection-2-17-offline/mmpycocotools-12.0.3-cp37-cp37m-linux_x86_64.whl --no-deps\n!pip install /kaggle/input/mmdetection-2-17-offline/terminal-0.4.0-py3-none-any.whl --no-deps\n!pip install /kaggle/input/mmdetection-2-17-offline/mmdet-2.17.0-py3-none-any.whl --no-deps\n!pip install /kaggle/input/mmdetection-2-17-offline/addict-2.4.0-py3-none-any.whl --no-deps\n!pip install /kaggle/input/mmdetection-2-17-offline/yapf-0.31.0-py2.py3-none-any.whl --no-deps","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:27:52.033267Z","iopub.execute_input":"2024-03-31T09:27:52.033643Z","iopub.status.idle":"2024-03-31T09:28:16.973412Z","shell.execute_reply.started":"2024-03-31T09:27:52.033605Z","shell.execute_reply":"2024-03-31T09:28:16.971791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install einops\n!pip install monai\nimport time","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:16.976368Z","iopub.execute_input":"2024-03-31T09:28:16.976758Z","iopub.status.idle":"2024-03-31T09:28:43.295782Z","shell.execute_reply.started":"2024-03-31T09:28:16.976707Z","shell.execute_reply":"2024-03-31T09:28:43.294588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/rsnazoopublic')\n","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:43.297371Z","iopub.execute_input":"2024-03-31T09:28:43.297707Z","iopub.status.idle":"2024-03-31T09:28:43.303415Z","shell.execute_reply.started":"2024-03-31T09:28:43.297671Z","shell.execute_reply":"2024-03-31T09:28:43.302267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport random\nimport re\nfrom dataclasses import dataclass\nfrom typing import Dict\nfrom typing import List\n\nimport albumentations\nimport cv2\nimport numpy as np\nimport pydicom\nimport tifffile\nimport torch\nimport torch.hub\nfrom albumentations import ReplayCompose\nfrom skimage import measure\nfrom torch.functional import Tensor\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:43.305118Z","iopub.execute_input":"2024-03-31T09:28:43.305496Z","iopub.status.idle":"2024-03-31T09:28:47.210939Z","shell.execute_reply.started":"2024-03-31T09:28:43.30546Z","shell.execute_reply":"2024-03-31T09:28:47.209482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Copyright (c) MONAI Consortium\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#     http://www.apache.org/licenses/LICENSE-2.0\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom __future__ import annotations\n\nimport itertools\nfrom collections.abc import Sequence\n\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.utils.checkpoint as checkpoint\nfrom torch.nn import LayerNorm\nfrom typing_extensions import Final\n\nfrom monai.networks.blocks import MLPBlock as Mlp\nfrom monai.networks.blocks import PatchEmbed, UnetOutBlock, UnetrBasicBlock, UnetrUpBlock\nfrom monai.networks.layers import DropPath, trunc_normal_\nfrom monai.utils import ensure_tuple_rep, look_up_option, optional_import\nfrom monai.utils.deprecate_utils import deprecated_arg\n\nrearrange, _ = optional_import(\"einops\", name=\"rearrange\")\n\n__all__ = [\n    \"SwinUNETR\",\n    \"window_partition\",\n    \"window_reverse\",\n    \"WindowAttention\",\n    \"SwinTransformerBlock\",\n    \"PatchMerging\",\n    \"PatchMergingV2\",\n    \"MERGING_MODE\",\n    \"BasicLayer\",\n    \"SwinTransformer\",\n]\n\n\n\n# [docs]\nclass SwinUNETR(nn.Module):\n    \"\"\"\n    Swin UNETR based on: \"Hatamizadeh et al.,\n    Swin UNETR: Swin Transformers for Semantic Segmentation of Brain Tumors in MRI Images\n    <https://arxiv.org/abs/2201.01266>\"\n    \"\"\"\n\n    patch_size: Final[int] = 2\n\n\n# [docs]\n#     @deprecated_arg(\n#         name=\"img_size\",\n#         since=\"1.3\",\n#         removed=\"1.5\",\n#         msg_suffix=\"The img_size argument is not required anymore and \"\n#         \"checks on the input size are run during forward().\",\n#     )\n    def __init__(\n        self,\n        img_size: Sequence[int] | int,\n        in_channels: int,\n        out_channels: int,\n        depths: Sequence[int] = (2, 2, 2, 2),\n        num_heads: Sequence[int] = (3, 6, 12, 24),\n        feature_size: int = 24,\n        norm_name: tuple | str = \"instance\",\n        drop_rate: float = 0.0,\n        attn_drop_rate: float = 0.0,\n        dropout_path_rate: float = 0.0,\n        normalize: bool = True,\n        use_checkpoint: bool = False,\n        spatial_dims: int = 3,\n        downsample=\"merging\",\n        use_v2=False,\n    ) -> None:\n        \"\"\"\n        Args:\n            img_size: spatial dimension of input image.\n                This argument is only used for checking that the input image size is divisible by the patch size.\n                The tensor passed to forward() can have a dynamic shape as long as its spatial dimensions are divisible by 2**5.\n                It will be removed in an upcoming version.\n            in_channels: dimension of input channels.\n            out_channels: dimension of output channels.\n            feature_size: dimension of network feature size.\n            depths: number of layers in each stage.\n            num_heads: number of attention heads.\n            norm_name: feature normalization type and arguments.\n            drop_rate: dropout rate.\n            attn_drop_rate: attention dropout rate.\n            dropout_path_rate: drop path rate.\n            normalize: normalize output intermediate features in each stage.\n            use_checkpoint: use gradient checkpointing for reduced memory usage.\n            spatial_dims: number of spatial dims.\n            downsample: module used for downsampling, available options are `\"mergingv2\"`, `\"merging\"` and a\n                user-specified `nn.Module` following the API defined in :py:class:`monai.networks.nets.PatchMerging`.\n                The default is currently `\"merging\"` (the original version defined in v0.9.0).\n            use_v2: using swinunetr_v2, which adds a residual convolution block at the beggining of each swin stage.\n\n        Examples::\n\n            # for 3D single channel input with size (96,96,96), 4-channel output and feature size of 48.\n            >>> net = SwinUNETR(img_size=(96,96,96), in_channels=1, out_channels=4, feature_size=48)\n\n            # for 3D 4-channel input with size (128,128,128), 3-channel output and (2,4,2,2) layers in each stage.\n            >>> net = SwinUNETR(img_size=(128,128,128), in_channels=4, out_channels=3, depths=(2,4,2,2))\n\n            # for 2D single channel input with size (96,96), 2-channel output and gradient checkpointing.\n            >>> net = SwinUNETR(img_size=(96,96), in_channels=3, out_channels=2, use_checkpoint=True, spatial_dims=2)\n\n        \"\"\"\n\n        super().__init__()\n\n        img_size = ensure_tuple_rep(img_size, spatial_dims)\n        patch_sizes = ensure_tuple_rep(self.patch_size, spatial_dims)\n        window_size = ensure_tuple_rep(7, spatial_dims)\n\n        if spatial_dims not in (2, 3):\n            raise ValueError(\"spatial dimension should be 2 or 3.\")\n\n        self._check_input_size(img_size)\n\n        if not (0 <= drop_rate <= 1):\n            raise ValueError(\"dropout rate should be between 0 and 1.\")\n\n        if not (0 <= attn_drop_rate <= 1):\n            raise ValueError(\"attention dropout rate should be between 0 and 1.\")\n\n        if not (0 <= dropout_path_rate <= 1):\n            raise ValueError(\"drop path rate should be between 0 and 1.\")\n\n        if feature_size % 12 != 0:\n            raise ValueError(\"feature_size should be divisible by 12.\")\n\n        self.normalize = normalize\n\n        self.swinViT = SwinTransformer(\n            in_chans=in_channels,\n            embed_dim=feature_size,\n            window_size=window_size,\n            patch_size=patch_sizes,\n            depths=depths,\n            num_heads=num_heads,\n            mlp_ratio=4.0,\n            qkv_bias=True,\n            drop_rate=drop_rate,\n            attn_drop_rate=attn_drop_rate,\n            drop_path_rate=dropout_path_rate,\n            norm_layer=nn.LayerNorm,\n            use_checkpoint=use_checkpoint,\n            spatial_dims=spatial_dims,\n            downsample=look_up_option(downsample, MERGING_MODE) if isinstance(downsample, str) else downsample,\n            use_v2=use_v2,\n        )\n\n        self.encoder1 = UnetrBasicBlock(\n            spatial_dims=spatial_dims,\n            in_channels=in_channels,\n            out_channels=feature_size,\n            kernel_size=3,\n            stride=1,\n            norm_name=norm_name,\n            res_block=True,\n        )\n\n        self.encoder2 = UnetrBasicBlock(\n            spatial_dims=spatial_dims,\n            in_channels=feature_size,\n            out_channels=feature_size,\n            kernel_size=3,\n            stride=1,\n            norm_name=norm_name,\n            res_block=True,\n        )\n\n        self.encoder3 = UnetrBasicBlock(\n            spatial_dims=spatial_dims,\n            in_channels=2 * feature_size,\n            out_channels=2 * feature_size,\n            kernel_size=3,\n            stride=1,\n            norm_name=norm_name,\n            res_block=True,\n        )\n\n        self.encoder4 = UnetrBasicBlock(\n            spatial_dims=spatial_dims,\n            in_channels=4 * feature_size,\n            out_channels=4 * feature_size,\n            kernel_size=3,\n            stride=1,\n            norm_name=norm_name,\n            res_block=True,\n        )\n\n        self.encoder10 = UnetrBasicBlock(\n            spatial_dims=spatial_dims,\n            in_channels=16 * feature_size,\n            out_channels=16 * feature_size,\n            kernel_size=3,\n            stride=1,\n            norm_name=norm_name,\n            res_block=True,\n        )\n\n        self.decoder5 = UnetrUpBlock(\n            spatial_dims=spatial_dims,\n            in_channels=16 * feature_size,\n            out_channels=8 * feature_size,\n            kernel_size=3,\n            upsample_kernel_size=2,\n            norm_name=norm_name,\n            res_block=True,\n        )\n\n        self.decoder4 = UnetrUpBlock(\n            spatial_dims=spatial_dims,\n            in_channels=feature_size * 8,\n            out_channels=feature_size * 4,\n            kernel_size=3,\n            upsample_kernel_size=2,\n            norm_name=norm_name,\n            res_block=True,\n        )\n\n        self.decoder3 = UnetrUpBlock(\n            spatial_dims=spatial_dims,\n            in_channels=feature_size * 4,\n            out_channels=feature_size * 2,\n            kernel_size=3,\n            upsample_kernel_size=2,\n            norm_name=norm_name,\n            res_block=True,\n        )\n        self.decoder2 = UnetrUpBlock(\n            spatial_dims=spatial_dims,\n            in_channels=feature_size * 2,\n            out_channels=feature_size,\n            kernel_size=3,\n            upsample_kernel_size=2,\n            norm_name=norm_name,\n            res_block=True,\n        )\n\n        self.decoder1 = UnetrUpBlock(\n            spatial_dims=spatial_dims,\n            in_channels=feature_size,\n            out_channels=feature_size,\n            kernel_size=3,\n            upsample_kernel_size=2,\n            norm_name=norm_name,\n            res_block=True,\n        )\n\n        self.out = UnetOutBlock(spatial_dims=spatial_dims, in_channels=feature_size, out_channels=out_channels)\n\n\n\n    def load_from(self, weights):\n        with torch.no_grad():\n            self.swinViT.patch_embed.proj.weight.copy_(weights[\"state_dict\"][\"module.patch_embed.proj.weight\"])\n            self.swinViT.patch_embed.proj.bias.copy_(weights[\"state_dict\"][\"module.patch_embed.proj.bias\"])\n            for bname, block in self.swinViT.layers1[0].blocks.named_children():\n                block.load_from(weights, n_block=bname, layer=\"layers1\")\n            self.swinViT.layers1[0].downsample.reduction.weight.copy_(\n                weights[\"state_dict\"][\"module.layers1.0.downsample.reduction.weight\"]\n            )\n            self.swinViT.layers1[0].downsample.norm.weight.copy_(\n                weights[\"state_dict\"][\"module.layers1.0.downsample.norm.weight\"]\n            )\n            self.swinViT.layers1[0].downsample.norm.bias.copy_(\n                weights[\"state_dict\"][\"module.layers1.0.downsample.norm.bias\"]\n            )\n            for bname, block in self.swinViT.layers2[0].blocks.named_children():\n                block.load_from(weights, n_block=bname, layer=\"layers2\")\n            self.swinViT.layers2[0].downsample.reduction.weight.copy_(\n                weights[\"state_dict\"][\"module.layers2.0.downsample.reduction.weight\"]\n            )\n            self.swinViT.layers2[0].downsample.norm.weight.copy_(\n                weights[\"state_dict\"][\"module.layers2.0.downsample.norm.weight\"]\n            )\n            self.swinViT.layers2[0].downsample.norm.bias.copy_(\n                weights[\"state_dict\"][\"module.layers2.0.downsample.norm.bias\"]\n            )\n            for bname, block in self.swinViT.layers3[0].blocks.named_children():\n                block.load_from(weights, n_block=bname, layer=\"layers3\")\n            self.swinViT.layers3[0].downsample.reduction.weight.copy_(\n                weights[\"state_dict\"][\"module.layers3.0.downsample.reduction.weight\"]\n            )\n            self.swinViT.layers3[0].downsample.norm.weight.copy_(\n                weights[\"state_dict\"][\"module.layers3.0.downsample.norm.weight\"]\n            )\n            self.swinViT.layers3[0].downsample.norm.bias.copy_(\n                weights[\"state_dict\"][\"module.layers3.0.downsample.norm.bias\"]\n            )\n            for bname, block in self.swinViT.layers4[0].blocks.named_children():\n                block.load_from(weights, n_block=bname, layer=\"layers4\")\n            self.swinViT.layers4[0].downsample.reduction.weight.copy_(\n                weights[\"state_dict\"][\"module.layers4.0.downsample.reduction.weight\"]\n            )\n            self.swinViT.layers4[0].downsample.norm.weight.copy_(\n                weights[\"state_dict\"][\"module.layers4.0.downsample.norm.weight\"]\n            )\n            self.swinViT.layers4[0].downsample.norm.bias.copy_(\n                weights[\"state_dict\"][\"module.layers4.0.downsample.norm.bias\"]\n            )\n\n    @torch.jit.unused\n    def _check_input_size(self, spatial_shape):\n        img_size = np.array(spatial_shape)\n        remainder = (img_size % np.power(self.patch_size, 5)) > 0\n        if remainder.any():\n            wrong_dims = (np.where(remainder)[0] + 2).tolist()\n            raise ValueError(\n                f\"spatial dimensions {wrong_dims} of input image (spatial shape: {spatial_shape})\"\n                f\" must be divisible by {self.patch_size}**5.\"\n            )\n\n\n# [docs]\n    def forward(self, x_in):\n        if not torch.jit.is_scripting():\n            self._check_input_size(x_in.shape[2:])\n        hidden_states_out = self.swinViT(x_in, self.normalize)\n        enc0 = self.encoder1(x_in)\n        enc1 = self.encoder2(hidden_states_out[0])\n        enc2 = self.encoder3(hidden_states_out[1])\n        enc3 = self.encoder4(hidden_states_out[2])\n        dec4 = self.encoder10(hidden_states_out[4])\n        dec3 = self.decoder5(dec4, hidden_states_out[3])\n        dec2 = self.decoder4(dec3, enc3)\n        dec1 = self.decoder3(dec2, enc2)\n        dec0 = self.decoder2(dec1, enc1)\n        out = self.decoder1(dec0, enc0)\n        logits = self.out(out)\n        return logits\n\n\n\n\n\ndef window_partition(x, window_size):\n    \"\"\"window partition operation based on: \"Liu et al.,\n    Swin Transformer: Hierarchical Vision Transformer using Shifted Windows\n    <https://arxiv.org/abs/2103.14030>\"\n    https://github.com/microsoft/Swin-Transformer\n\n     Args:\n        x: input tensor.\n        window_size: local window size.\n    \"\"\"\n    x_shape = x.size()\n    if len(x_shape) == 5:\n        b, d, h, w, c = x_shape\n        x = x.view(\n            b,\n            d // window_size[0],\n            window_size[0],\n            h // window_size[1],\n            window_size[1],\n            w // window_size[2],\n            window_size[2],\n            c,\n        )\n        windows = (\n            x.permute(0, 1, 3, 5, 2, 4, 6, 7).contiguous().view(-1, window_size[0] * window_size[1] * window_size[2], c)\n        )\n    elif len(x_shape) == 4:\n        b, h, w, c = x.shape\n        x = x.view(b, h // window_size[0], window_size[0], w // window_size[1], window_size[1], c)\n        windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size[0] * window_size[1], c)\n    return windows\n\n\ndef window_reverse(windows, window_size, dims):\n    \"\"\"window reverse operation based on: \"Liu et al.,\n    Swin Transformer: Hierarchical Vision Transformer using Shifted Windows\n    <https://arxiv.org/abs/2103.14030>\"\n    https://github.com/microsoft/Swin-Transformer\n\n     Args:\n        windows: windows tensor.\n        window_size: local window size.\n        dims: dimension values.\n    \"\"\"\n    if len(dims) == 4:\n        b, d, h, w = dims\n        x = windows.view(\n            b,\n            d // window_size[0],\n            h // window_size[1],\n            w // window_size[2],\n            window_size[0],\n            window_size[1],\n            window_size[2],\n            -1,\n        )\n        x = x.permute(0, 1, 4, 2, 5, 3, 6, 7).contiguous().view(b, d, h, w, -1)\n\n    elif len(dims) == 3:\n        b, h, w = dims\n        x = windows.view(b, h // window_size[0], w // window_size[1], window_size[0], window_size[1], -1)\n        x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(b, h, w, -1)\n    return x\n\n\ndef get_window_size(x_size, window_size, shift_size=None):\n    \"\"\"Computing window size based on: \"Liu et al.,\n    Swin Transformer: Hierarchical Vision Transformer using Shifted Windows\n    <https://arxiv.org/abs/2103.14030>\"\n    https://github.com/microsoft/Swin-Transformer\n\n     Args:\n        x_size: input size.\n        window_size: local window size.\n        shift_size: window shifting size.\n    \"\"\"\n\n    use_window_size = list(window_size)\n    if shift_size is not None:\n        use_shift_size = list(shift_size)\n    for i in range(len(x_size)):\n        if x_size[i] <= window_size[i]:\n            use_window_size[i] = x_size[i]\n            if shift_size is not None:\n                use_shift_size[i] = 0\n\n    if shift_size is None:\n        return tuple(use_window_size)\n    else:\n        return tuple(use_window_size), tuple(use_shift_size)\n\n\nclass WindowAttention(nn.Module):\n    \"\"\"\n    Window based multi-head self attention module with relative position bias based on: \"Liu et al.,\n    Swin Transformer: Hierarchical Vision Transformer using Shifted Windows\n    <https://arxiv.org/abs/2103.14030>\"\n    https://github.com/microsoft/Swin-Transformer\n    \"\"\"\n\n    def __init__(\n        self,\n        dim: int,\n        num_heads: int,\n        window_size: Sequence[int],\n        qkv_bias: bool = False,\n        attn_drop: float = 0.0,\n        proj_drop: float = 0.0,\n    ) -> None:\n        \"\"\"\n        Args:\n            dim: number of feature channels.\n            num_heads: number of attention heads.\n            window_size: local window size.\n            qkv_bias: add a learnable bias to query, key, value.\n            attn_drop: attention dropout rate.\n            proj_drop: dropout rate of output.\n        \"\"\"\n\n        super().__init__()\n        self.dim = dim\n        self.window_size = window_size\n        self.num_heads = num_heads\n        head_dim = dim // num_heads\n        self.scale = head_dim**-0.5\n        mesh_args = torch.meshgrid.__kwdefaults__\n\n        if len(self.window_size) == 3:\n            self.relative_position_bias_table = nn.Parameter(\n                torch.zeros(\n                    (2 * self.window_size[0] - 1) * (2 * self.window_size[1] - 1) * (2 * self.window_size[2] - 1),\n                    num_heads,\n                )\n            )\n            coords_d = torch.arange(self.window_size[0])\n            coords_h = torch.arange(self.window_size[1])\n            coords_w = torch.arange(self.window_size[2])\n            if mesh_args is not None:\n                coords = torch.stack(torch.meshgrid(coords_d, coords_h, coords_w, indexing=\"ij\"))\n            else:\n                coords = torch.stack(torch.meshgrid(coords_d, coords_h, coords_w))\n            coords_flatten = torch.flatten(coords, 1)\n            relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :]\n            relative_coords = relative_coords.permute(1, 2, 0).contiguous()\n            relative_coords[:, :, 0] += self.window_size[0] - 1\n            relative_coords[:, :, 1] += self.window_size[1] - 1\n            relative_coords[:, :, 2] += self.window_size[2] - 1\n            relative_coords[:, :, 0] *= (2 * self.window_size[1] - 1) * (2 * self.window_size[2] - 1)\n            relative_coords[:, :, 1] *= 2 * self.window_size[2] - 1\n        elif len(self.window_size) == 2:\n            self.relative_position_bias_table = nn.Parameter(\n                torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads)\n            )\n            coords_h = torch.arange(self.window_size[0])\n            coords_w = torch.arange(self.window_size[1])\n            if mesh_args is not None:\n                coords = torch.stack(torch.meshgrid(coords_h, coords_w, indexing=\"ij\"))\n            else:\n                coords = torch.stack(torch.meshgrid(coords_h, coords_w))\n            coords_flatten = torch.flatten(coords, 1)\n            relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :]\n            relative_coords = relative_coords.permute(1, 2, 0).contiguous()\n            relative_coords[:, :, 0] += self.window_size[0] - 1\n            relative_coords[:, :, 1] += self.window_size[1] - 1\n            relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1\n\n        relative_position_index = relative_coords.sum(-1)\n        self.register_buffer(\"relative_position_index\", relative_position_index)\n        self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)\n        self.attn_drop = nn.Dropout(attn_drop)\n        self.proj = nn.Linear(dim, dim)\n        self.proj_drop = nn.Dropout(proj_drop)\n        trunc_normal_(self.relative_position_bias_table, std=0.02)\n        self.softmax = nn.Softmax(dim=-1)\n\n    def forward(self, x, mask):\n        b, n, c = x.shape\n        qkv = self.qkv(x).reshape(b, n, 3, self.num_heads, c // self.num_heads).permute(2, 0, 3, 1, 4)\n        q, k, v = qkv[0], qkv[1], qkv[2]\n        q = q * self.scale\n        attn = q @ k.transpose(-2, -1)\n        relative_position_bias = self.relative_position_bias_table[\n            self.relative_position_index.clone()[:n, :n].reshape(-1)\n        ].reshape(n, n, -1)\n        relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous()\n        attn = attn + relative_position_bias.unsqueeze(0)\n        if mask is not None:\n            nw = mask.shape[0]\n            attn = attn.view(b // nw, nw, self.num_heads, n, n) + mask.unsqueeze(1).unsqueeze(0)\n            attn = attn.view(-1, self.num_heads, n, n)\n            attn = self.softmax(attn)\n        else:\n            attn = self.softmax(attn)\n\n        attn = self.attn_drop(attn).to(v.dtype)\n        x = (attn @ v).transpose(1, 2).reshape(b, n, c)\n        x = self.proj(x)\n        x = self.proj_drop(x)\n        return x\n\n\nclass SwinTransformerBlock(nn.Module):\n    \"\"\"\n    Swin Transformer block based on: \"Liu et al.,\n    Swin Transformer: Hierarchical Vision Transformer using Shifted Windows\n    <https://arxiv.org/abs/2103.14030>\"\n    https://github.com/microsoft/Swin-Transformer\n    \"\"\"\n\n    def __init__(\n        self,\n        dim: int,\n        num_heads: int,\n        window_size: Sequence[int],\n        shift_size: Sequence[int],\n        mlp_ratio: float = 4.0,\n        qkv_bias: bool = True,\n        drop: float = 0.0,\n        attn_drop: float = 0.0,\n        drop_path: float = 0.0,\n        act_layer: str = \"GELU\",\n        norm_layer: type[LayerNorm] = nn.LayerNorm,\n        use_checkpoint: bool = False,\n    ) -> None:\n        \"\"\"\n        Args:\n            dim: number of feature channels.\n            num_heads: number of attention heads.\n            window_size: local window size.\n            shift_size: window shift size.\n            mlp_ratio: ratio of mlp hidden dim to embedding dim.\n            qkv_bias: add a learnable bias to query, key, value.\n            drop: dropout rate.\n            attn_drop: attention dropout rate.\n            drop_path: stochastic depth rate.\n            act_layer: activation layer.\n            norm_layer: normalization layer.\n            use_checkpoint: use gradient checkpointing for reduced memory usage.\n        \"\"\"\n\n        super().__init__()\n        self.dim = dim\n        self.num_heads = num_heads\n        self.window_size = window_size\n        self.shift_size = shift_size\n        self.mlp_ratio = mlp_ratio\n        self.use_checkpoint = use_checkpoint\n        self.norm1 = norm_layer(dim)\n        self.attn = WindowAttention(\n            dim,\n            window_size=self.window_size,\n            num_heads=num_heads,\n            qkv_bias=qkv_bias,\n            attn_drop=attn_drop,\n            proj_drop=drop,\n        )\n\n        self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()\n        self.norm2 = norm_layer(dim)\n        mlp_hidden_dim = int(dim * mlp_ratio)\n        self.mlp = Mlp(hidden_size=dim, mlp_dim=mlp_hidden_dim, act=act_layer, dropout_rate=drop, dropout_mode=\"swin\")\n\n    def forward_part1(self, x, mask_matrix):\n        x_shape = x.size()\n        x = self.norm1(x)\n        if len(x_shape) == 5:\n            b, d, h, w, c = x.shape\n            window_size, shift_size = get_window_size((d, h, w), self.window_size, self.shift_size)\n            pad_l = pad_t = pad_d0 = 0\n            pad_d1 = (window_size[0] - d % window_size[0]) % window_size[0]\n            pad_b = (window_size[1] - h % window_size[1]) % window_size[1]\n            pad_r = (window_size[2] - w % window_size[2]) % window_size[2]\n            x = F.pad(x, (0, 0, pad_l, pad_r, pad_t, pad_b, pad_d0, pad_d1))\n            _, dp, hp, wp, _ = x.shape\n            dims = [b, dp, hp, wp]\n\n        elif len(x_shape) == 4:\n            b, h, w, c = x.shape\n            window_size, shift_size = get_window_size((h, w), self.window_size, self.shift_size)\n            pad_l = pad_t = 0\n            pad_b = (window_size[0] - h % window_size[0]) % window_size[0]\n            pad_r = (window_size[1] - w % window_size[1]) % window_size[1]\n            x = F.pad(x, (0, 0, pad_l, pad_r, pad_t, pad_b))\n            _, hp, wp, _ = x.shape\n            dims = [b, hp, wp]\n\n        if any(i > 0 for i in shift_size):\n            if len(x_shape) == 5:\n                shifted_x = torch.roll(x, shifts=(-shift_size[0], -shift_size[1], -shift_size[2]), dims=(1, 2, 3))\n            elif len(x_shape) == 4:\n                shifted_x = torch.roll(x, shifts=(-shift_size[0], -shift_size[1]), dims=(1, 2))\n            attn_mask = mask_matrix\n        else:\n            shifted_x = x\n            attn_mask = None\n        x_windows = window_partition(shifted_x, window_size)\n        attn_windows = self.attn(x_windows, mask=attn_mask)\n        attn_windows = attn_windows.view(-1, *(window_size + (c,)))\n        shifted_x = window_reverse(attn_windows, window_size, dims)\n        if any(i > 0 for i in shift_size):\n            if len(x_shape) == 5:\n                x = torch.roll(shifted_x, shifts=(shift_size[0], shift_size[1], shift_size[2]), dims=(1, 2, 3))\n            elif len(x_shape) == 4:\n                x = torch.roll(shifted_x, shifts=(shift_size[0], shift_size[1]), dims=(1, 2))\n        else:\n            x = shifted_x\n\n        if len(x_shape) == 5:\n            if pad_d1 > 0 or pad_r > 0 or pad_b > 0:\n                x = x[:, :d, :h, :w, :].contiguous()\n        elif len(x_shape) == 4:\n            if pad_r > 0 or pad_b > 0:\n                x = x[:, :h, :w, :].contiguous()\n\n        return x\n\n    def forward_part2(self, x):\n        return self.drop_path(self.mlp(self.norm2(x)))\n\n    def load_from(self, weights, n_block, layer):\n        root = f\"module.{layer}.0.blocks.{n_block}.\"\n        block_names = [\n            \"norm1.weight\",\n            \"norm1.bias\",\n            \"attn.relative_position_bias_table\",\n            \"attn.relative_position_index\",\n            \"attn.qkv.weight\",\n            \"attn.qkv.bias\",\n            \"attn.proj.weight\",\n            \"attn.proj.bias\",\n            \"norm2.weight\",\n            \"norm2.bias\",\n            \"mlp.fc1.weight\",\n            \"mlp.fc1.bias\",\n            \"mlp.fc2.weight\",\n            \"mlp.fc2.bias\",\n        ]\n        with torch.no_grad():\n            self.norm1.weight.copy_(weights[\"state_dict\"][root + block_names[0]])\n            self.norm1.bias.copy_(weights[\"state_dict\"][root + block_names[1]])\n            self.attn.relative_position_bias_table.copy_(weights[\"state_dict\"][root + block_names[2]])\n            self.attn.relative_position_index.copy_(weights[\"state_dict\"][root + block_names[3]])\n            self.attn.qkv.weight.copy_(weights[\"state_dict\"][root + block_names[4]])\n            self.attn.qkv.bias.copy_(weights[\"state_dict\"][root + block_names[5]])\n            self.attn.proj.weight.copy_(weights[\"state_dict\"][root + block_names[6]])\n            self.attn.proj.bias.copy_(weights[\"state_dict\"][root + block_names[7]])\n            self.norm2.weight.copy_(weights[\"state_dict\"][root + block_names[8]])\n            self.norm2.bias.copy_(weights[\"state_dict\"][root + block_names[9]])\n            self.mlp.linear1.weight.copy_(weights[\"state_dict\"][root + block_names[10]])\n            self.mlp.linear1.bias.copy_(weights[\"state_dict\"][root + block_names[11]])\n            self.mlp.linear2.weight.copy_(weights[\"state_dict\"][root + block_names[12]])\n            self.mlp.linear2.bias.copy_(weights[\"state_dict\"][root + block_names[13]])\n\n    def forward(self, x, mask_matrix):\n        shortcut = x\n        if self.use_checkpoint:\n            x = checkpoint.checkpoint(self.forward_part1, x, mask_matrix, use_reentrant=False)\n        else:\n            x = self.forward_part1(x, mask_matrix)\n        x = shortcut + self.drop_path(x)\n        if self.use_checkpoint:\n            x = x + checkpoint.checkpoint(self.forward_part2, x, use_reentrant=False)\n        else:\n            x = x + self.forward_part2(x)\n        return x\n\n\nclass PatchMergingV2(nn.Module):\n    \"\"\"\n    Patch merging layer based on: \"Liu et al.,\n    Swin Transformer: Hierarchical Vision Transformer using Shifted Windows\n    <https://arxiv.org/abs/2103.14030>\"\n    https://github.com/microsoft/Swin-Transformer\n    \"\"\"\n\n    def __init__(self, dim: int, norm_layer: type[LayerNorm] = nn.LayerNorm, spatial_dims: int = 3) -> None:\n        \"\"\"\n        Args:\n            dim: number of feature channels.\n            norm_layer: normalization layer.\n            spatial_dims: number of spatial dims.\n        \"\"\"\n\n        super().__init__()\n        self.dim = dim\n        if spatial_dims == 3:\n            self.reduction = nn.Linear(8 * dim, 2 * dim, bias=False)\n            self.norm = norm_layer(8 * dim)\n        elif spatial_dims == 2:\n            self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)\n            self.norm = norm_layer(4 * dim)\n\n    def forward(self, x):\n        x_shape = x.size()\n        if len(x_shape) == 5:\n            b, d, h, w, c = x_shape\n            pad_input = (h % 2 == 1) or (w % 2 == 1) or (d % 2 == 1)\n            if pad_input:\n                x = F.pad(x, (0, 0, 0, w % 2, 0, h % 2, 0, d % 2))\n            x = torch.cat(\n                [x[:, i::2, j::2, k::2, :] for i, j, k in itertools.product(range(2), range(2), range(2))], -1\n            )\n\n        elif len(x_shape) == 4:\n            b, h, w, c = x_shape\n            pad_input = (h % 2 == 1) or (w % 2 == 1)\n            if pad_input:\n                x = F.pad(x, (0, 0, 0, w % 2, 0, h % 2))\n            x = torch.cat([x[:, j::2, i::2, :] for i, j in itertools.product(range(2), range(2))], -1)\n\n        x = self.norm(x)\n        x = self.reduction(x)\n        return x\n\n\nclass PatchMerging(PatchMergingV2):\n    \"\"\"The `PatchMerging` module previously defined in v0.9.0.\"\"\"\n\n    def forward(self, x):\n        x_shape = x.size()\n        if len(x_shape) == 4:\n            return super().forward(x)\n        if len(x_shape) != 5:\n            raise ValueError(f\"expecting 5D x, got {x.shape}.\")\n        b, d, h, w, c = x_shape\n        pad_input = (h % 2 == 1) or (w % 2 == 1) or (d % 2 == 1)\n        if pad_input:\n            x = F.pad(x, (0, 0, 0, w % 2, 0, h % 2, 0, d % 2))\n        x0 = x[:, 0::2, 0::2, 0::2, :]\n        x1 = x[:, 1::2, 0::2, 0::2, :]\n        x2 = x[:, 0::2, 1::2, 0::2, :]\n        x3 = x[:, 0::2, 0::2, 1::2, :]\n        x4 = x[:, 1::2, 0::2, 1::2, :]\n        x5 = x[:, 0::2, 1::2, 0::2, :]\n        x6 = x[:, 0::2, 0::2, 1::2, :]\n        x7 = x[:, 1::2, 1::2, 1::2, :]\n        x = torch.cat([x0, x1, x2, x3, x4, x5, x6, x7], -1)\n        x = self.norm(x)\n        x = self.reduction(x)\n        return x\n\n\nMERGING_MODE = {\"merging\": PatchMerging, \"mergingv2\": PatchMergingV2}\n\n\ndef compute_mask(dims, window_size, shift_size, device):\n    \"\"\"Computing region masks based on: \"Liu et al.,\n    Swin Transformer: Hierarchical Vision Transformer using Shifted Windows\n    <https://arxiv.org/abs/2103.14030>\"\n    https://github.com/microsoft/Swin-Transformer\n\n     Args:\n        dims: dimension values.\n        window_size: local window size.\n        shift_size: shift size.\n        device: device.\n    \"\"\"\n\n    cnt = 0\n\n    if len(dims) == 3:\n        d, h, w = dims\n        img_mask = torch.zeros((1, d, h, w, 1), device=device)\n        for d in slice(-window_size[0]), slice(-window_size[0], -shift_size[0]), slice(-shift_size[0], None):\n            for h in slice(-window_size[1]), slice(-window_size[1], -shift_size[1]), slice(-shift_size[1], None):\n                for w in slice(-window_size[2]), slice(-window_size[2], -shift_size[2]), slice(-shift_size[2], None):\n                    img_mask[:, d, h, w, :] = cnt\n                    cnt += 1\n\n    elif len(dims) == 2:\n        h, w = dims\n        img_mask = torch.zeros((1, h, w, 1), device=device)\n        for h in slice(-window_size[0]), slice(-window_size[0], -shift_size[0]), slice(-shift_size[0], None):\n            for w in slice(-window_size[1]), slice(-window_size[1], -shift_size[1]), slice(-shift_size[1], None):\n                img_mask[:, h, w, :] = cnt\n                cnt += 1\n\n    mask_windows = window_partition(img_mask, window_size)\n    mask_windows = mask_windows.squeeze(-1)\n    attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)\n    attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))\n\n    return attn_mask\n\n\nclass BasicLayer(nn.Module):\n    \"\"\"\n    Basic Swin Transformer layer in one stage based on: \"Liu et al.,\n    Swin Transformer: Hierarchical Vision Transformer using Shifted Windows\n    <https://arxiv.org/abs/2103.14030>\"\n    https://github.com/microsoft/Swin-Transformer\n    \"\"\"\n\n    def __init__(\n        self,\n        dim: int,\n        depth: int,\n        num_heads: int,\n        window_size: Sequence[int],\n        drop_path: list,\n        mlp_ratio: float = 4.0,\n        qkv_bias: bool = False,\n        drop: float = 0.0,\n        attn_drop: float = 0.0,\n        norm_layer: type[LayerNorm] = nn.LayerNorm,\n        downsample: nn.Module | None = None,\n        use_checkpoint: bool = False,\n    ) -> None:\n        \"\"\"\n        Args:\n            dim: number of feature channels.\n            depth: number of layers in each stage.\n            num_heads: number of attention heads.\n            window_size: local window size.\n            drop_path: stochastic depth rate.\n            mlp_ratio: ratio of mlp hidden dim to embedding dim.\n            qkv_bias: add a learnable bias to query, key, value.\n            drop: dropout rate.\n            attn_drop: attention dropout rate.\n            norm_layer: normalization layer.\n            downsample: an optional downsampling layer at the end of the layer.\n            use_checkpoint: use gradient checkpointing for reduced memory usage.\n        \"\"\"\n\n        super().__init__()\n        self.window_size = window_size\n        self.shift_size = tuple(i // 2 for i in window_size)\n        self.no_shift = tuple(0 for i in window_size)\n        self.depth = depth\n        self.use_checkpoint = use_checkpoint\n        self.blocks = nn.ModuleList(\n            [\n                SwinTransformerBlock(\n                    dim=dim,\n                    num_heads=num_heads,\n                    window_size=self.window_size,\n                    shift_size=self.no_shift if (i % 2 == 0) else self.shift_size,\n                    mlp_ratio=mlp_ratio,\n                    qkv_bias=qkv_bias,\n                    drop=drop,\n                    attn_drop=attn_drop,\n                    drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path,\n                    norm_layer=norm_layer,\n                    use_checkpoint=use_checkpoint,\n                )\n                for i in range(depth)\n            ]\n        )\n        self.downsample = downsample\n        if callable(self.downsample):\n            self.downsample = downsample(dim=dim, norm_layer=norm_layer, spatial_dims=len(self.window_size))\n\n    def forward(self, x):\n        x_shape = x.size()\n        if len(x_shape) == 5:\n            b, c, d, h, w = x_shape\n            window_size, shift_size = get_window_size((d, h, w), self.window_size, self.shift_size)\n            x = rearrange(x, \"b c d h w -> b d h w c\")\n            dp = int(np.ceil(d / window_size[0])) * window_size[0]\n            hp = int(np.ceil(h / window_size[1])) * window_size[1]\n            wp = int(np.ceil(w / window_size[2])) * window_size[2]\n            attn_mask = compute_mask([dp, hp, wp], window_size, shift_size, x.device)\n            for blk in self.blocks:\n                x = blk(x, attn_mask)\n            x = x.view(b, d, h, w, -1)\n            if self.downsample is not None:\n                x = self.downsample(x)\n            x = rearrange(x, \"b d h w c -> b c d h w\")\n\n        elif len(x_shape) == 4:\n            b, c, h, w = x_shape\n            window_size, shift_size = get_window_size((h, w), self.window_size, self.shift_size)\n            x = rearrange(x, \"b c h w -> b h w c\")\n            hp = int(np.ceil(h / window_size[0])) * window_size[0]\n            wp = int(np.ceil(w / window_size[1])) * window_size[1]\n            attn_mask = compute_mask([hp, wp], window_size, shift_size, x.device)\n            for blk in self.blocks:\n                x = blk(x, attn_mask)\n            x = x.view(b, h, w, -1)\n            if self.downsample is not None:\n                x = self.downsample(x)\n            x = rearrange(x, \"b h w c -> b c h w\")\n        return x\n\n\nclass SwinTransformer(nn.Module):\n    \"\"\"\n    Swin Transformer based on: \"Liu et al.,\n    Swin Transformer: Hierarchical Vision Transformer using Shifted Windows\n    <https://arxiv.org/abs/2103.14030>\"\n    https://github.com/microsoft/Swin-Transformer\n    \"\"\"\n\n    def __init__(\n        self,\n        in_chans: int,\n        embed_dim: int,\n        window_size: Sequence[int],\n        patch_size: Sequence[int],\n        depths: Sequence[int],\n        num_heads: Sequence[int],\n        mlp_ratio: float = 4.0,\n        qkv_bias: bool = True,\n        drop_rate: float = 0.0,\n        attn_drop_rate: float = 0.0,\n        drop_path_rate: float = 0.0,\n        norm_layer: type[LayerNorm] = nn.LayerNorm,\n        patch_norm: bool = False,\n        use_checkpoint: bool = False,\n        spatial_dims: int = 3,\n        downsample=\"merging\",\n        use_v2=False,\n    ) -> None:\n        \"\"\"\n        Args:\n            in_chans: dimension of input channels.\n            embed_dim: number of linear projection output channels.\n            window_size: local window size.\n            patch_size: patch size.\n            depths: number of layers in each stage.\n            num_heads: number of attention heads.\n            mlp_ratio: ratio of mlp hidden dim to embedding dim.\n            qkv_bias: add a learnable bias to query, key, value.\n            drop_rate: dropout rate.\n            attn_drop_rate: attention dropout rate.\n            drop_path_rate: stochastic depth rate.\n            norm_layer: normalization layer.\n            patch_norm: add normalization after patch embedding.\n            use_checkpoint: use gradient checkpointing for reduced memory usage.\n            spatial_dims: spatial dimension.\n            downsample: module used for downsampling, available options are `\"mergingv2\"`, `\"merging\"` and a\n                user-specified `nn.Module` following the API defined in :py:class:`monai.networks.nets.PatchMerging`.\n                The default is currently `\"merging\"` (the original version defined in v0.9.0).\n            use_v2: using swinunetr_v2, which adds a residual convolution block at the beginning of each swin stage.\n        \"\"\"\n\n        super().__init__()\n        self.num_layers = len(depths)\n        self.embed_dim = embed_dim\n        self.patch_norm = patch_norm\n        self.window_size = window_size\n        self.patch_size = patch_size\n        self.patch_embed = PatchEmbed(\n            patch_size=self.patch_size,\n            in_chans=in_chans,\n            embed_dim=embed_dim,\n            norm_layer=norm_layer if self.patch_norm else None,  # type: ignore\n            spatial_dims=spatial_dims,\n        )\n        self.pos_drop = nn.Dropout(p=drop_rate)\n        dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))]\n        self.use_v2 = use_v2\n        self.layers1 = nn.ModuleList()\n        self.layers2 = nn.ModuleList()\n        self.layers3 = nn.ModuleList()\n        self.layers4 = nn.ModuleList()\n        if self.use_v2:\n            self.layers1c = nn.ModuleList()\n            self.layers2c = nn.ModuleList()\n            self.layers3c = nn.ModuleList()\n            self.layers4c = nn.ModuleList()\n        down_sample_mod = look_up_option(downsample, MERGING_MODE) if isinstance(downsample, str) else downsample\n        for i_layer in range(self.num_layers):\n            layer = BasicLayer(\n                dim=int(embed_dim * 2**i_layer),\n                depth=depths[i_layer],\n                num_heads=num_heads[i_layer],\n                window_size=self.window_size,\n                drop_path=dpr[sum(depths[:i_layer]) : sum(depths[: i_layer + 1])],\n                mlp_ratio=mlp_ratio,\n                qkv_bias=qkv_bias,\n                drop=drop_rate,\n                attn_drop=attn_drop_rate,\n                norm_layer=norm_layer,\n                downsample=down_sample_mod,\n                use_checkpoint=use_checkpoint,\n            )\n            if i_layer == 0:\n                self.layers1.append(layer)\n            elif i_layer == 1:\n                self.layers2.append(layer)\n            elif i_layer == 2:\n                self.layers3.append(layer)\n            elif i_layer == 3:\n                self.layers4.append(layer)\n            if self.use_v2:\n                layerc = UnetrBasicBlock(\n                    spatial_dims=3,\n                    in_channels=embed_dim * 2**i_layer,\n                    out_channels=embed_dim * 2**i_layer,\n                    kernel_size=3,\n                    stride=1,\n                    norm_name=\"instance\",\n                    res_block=True,\n                )\n                if i_layer == 0:\n                    self.layers1c.append(layerc)\n                elif i_layer == 1:\n                    self.layers2c.append(layerc)\n                elif i_layer == 2:\n                    self.layers3c.append(layerc)\n                elif i_layer == 3:\n                    self.layers4c.append(layerc)\n\n        self.num_features = int(embed_dim * 2 ** (self.num_layers - 1))\n\n    def proj_out(self, x, normalize=False):\n        if normalize:\n            x_shape = x.size()\n            if len(x_shape) == 5:\n                n, ch, d, h, w = x_shape\n                x = rearrange(x, \"n c d h w -> n d h w c\")\n                x = F.layer_norm(x, [ch])\n                x = rearrange(x, \"n d h w c -> n c d h w\")\n            elif len(x_shape) == 4:\n                n, ch, h, w = x_shape\n                x = rearrange(x, \"n c h w -> n h w c\")\n                x = F.layer_norm(x, [ch])\n                x = rearrange(x, \"n h w c -> n c h w\")\n        return x\n\n    def forward(self, x, normalize=True):\n        x0 = self.patch_embed(x)\n        x0 = self.pos_drop(x0)\n        x0_out = self.proj_out(x0, normalize)\n        if self.use_v2:\n            x0 = self.layers1c[0](x0.contiguous())\n        x1 = self.layers1[0](x0.contiguous())\n        x1_out = self.proj_out(x1, normalize)\n        if self.use_v2:\n            x1 = self.layers2c[0](x1.contiguous())\n        x2 = self.layers2[0](x1.contiguous())\n        x2_out = self.proj_out(x2, normalize)\n        if self.use_v2:\n            x2 = self.layers3c[0](x2.contiguous())\n        x3 = self.layers3[0](x2.contiguous())\n        x3_out = self.proj_out(x3, normalize)\n        if self.use_v2:\n            x3 = self.layers4c[0](x3.contiguous())\n        x4 = self.layers4[0](x3.contiguous())\n        x4_out = self.proj_out(x4, normalize)\n        return [x0_out, x1_out, x2_out, x3_out, x4_out]\n\n\ndef filter_swinunetr(key, value):\n    \"\"\"\n    A filter function used to filter the pretrained weights from [1], then the weights can be loaded into MONAI SwinUNETR Model.\n    This function is typically used with `monai.networks.copy_model_state`\n    [1] \"Valanarasu JM et al., Disruptive Autoencoders: Leveraging Low-level features for 3D Medical Image Pre-training\n    <https://arxiv.org/abs/2307.16896>\"\n\n    Args:\n        key: the key in the source state dict used for the update.\n        value: the value in the source state dict used for the update.\n\n    Examples::\n\n        import torch\n        from monai.apps import download_url\n        from monai.networks.utils import copy_model_state\n        from monai.networks.nets.swin_unetr import SwinUNETR, filter_swinunetr\n\n        model = SwinUNETR(img_size=(96, 96, 96), in_channels=1, out_channels=3, feature_size=48)\n        resource = (\n            \"https://github.com/Project-MONAI/MONAI-extra-test-data/releases/download/0.8.1/ssl_pretrained_weights.pth\"\n        )\n        ssl_weights_path = \"./ssl_pretrained_weights.pth\"\n        download_url(resource, ssl_weights_path)\n        ssl_weights = torch.load(ssl_weights_path)[\"model\"]\n\n        dst_dict, loaded, not_loaded = copy_model_state(model, ssl_weights, filter_func=filter_swinunetr)\n\n    \"\"\"\n    if key in [\n        \"encoder.mask_token\",\n        \"encoder.norm.weight\",\n        \"encoder.norm.bias\",\n        \"out.conv.conv.weight\",\n        \"out.conv.conv.bias\",\n    ]:\n        return None\n\n    if key[:8] == \"encoder.\":\n        if key[8:19] == \"patch_embed\":\n            new_key = \"swinViT.\" + key[8:]\n        else:\n            new_key = \"swinViT.\" + key[8:18] + key[20:]\n\n        return new_key, value\n    else:\n        return None","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:47.213346Z","iopub.execute_input":"2024-03-31T09:28:47.213865Z","iopub.status.idle":"2024-03-31T09:28:53.799082Z","shell.execute_reply.started":"2024-03-31T09:28:47.213834Z","shell.execute_reply":"2024-03-31T09:28:53.798166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SwinTransformerWithClassification(nn.Module):\n    def __init__(self, in_chans, embed_dim, window_size, patch_size, depths, num_heads, num_classes):\n        super(SwinTransformerWithClassification, self).__init__()\n        self.features = SwinTransformer(\n            in_chans=in_chans,\n            embed_dim=embed_dim,\n            window_size=window_size,\n            patch_size=patch_size,\n            depths=depths,\n            num_heads=num_heads\n        )\n        self.avgpool = nn.AdaptiveAvgPool3d((1, 1, 1))  # Adaptive average pooling to get a fixed-size output\n        self.fc = nn.Linear(embed_dim * 16, num_classes)  # Fully connected layer for classification\n\n    def forward(self, x):\n        x = self.features(x)\n        print(x[-1].shape)\n        x = self.avgpool(x[-1])  # Apply average pooling to the last feature map\n        print(x.shape)\n        x = x.view(x.size(0), -1)  # Flatten the feature map\n#         print(x.shape)\n        x = self.fc(x)  # Classification head\n        x = torch.sigmoid(x)  # Apply sigmoid activation for multi-label classification\n        return x\nmodel=SwinTransformerWithClassification(in_chans=3,embed_dim=48,window_size=(7,7,7),patch_size=(2,2,2), depths=(2,2,2,2),num_heads=(3,6,12,24),num_classes=8)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:53.800619Z","iopub.execute_input":"2024-03-31T09:28:53.80117Z","iopub.status.idle":"2024-03-31T09:28:53.918354Z","shell.execute_reply.started":"2024-03-31T09:28:53.801131Z","shell.execute_reply":"2024-03-31T09:28:53.917415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from monai.networks.nets.vit import ViT\nclass ViTWithClassification(nn.Module):\n    def __init__(self, in_channels, img_size, patch_size, num_classes):\n        super(ViTWithClassification, self).__init__()\n        self.features = ViT(in_channels, img_size, patch_size)\n        self.avgpool = nn.AdaptiveAvgPool1d((1))  # Adaptive average pooling to get a fixed-size output\n        self.fc = nn.Linear(768, num_classes)  # Fully connected layer for classification\n\n    def forward(self, x):\n        _, intermediate = self.features(x)  # Get the intermediate representations from the features module\n        print(intermediate[-1].shape)  # Print the shape of the last intermediate representation\n        x = intermediate[-1].transpose(1, 2)\n        x = self.avgpool(x)  # Apply average pooling to the last intermediate representation\n        print(x.shape)  # Print the shape after average pooling\n        x = torch.flatten(x, 1)  # Flatten the feature map\n        print(x.shape)  # Print the shape after flattening\n        x = self.fc(x)  # Classification head\n        x = torch.sigmoid(x)  # Apply sigmoid activation for multi-label classification\n        return x\n    \nmodel=ViTWithClassification(in_channels=3,img_size=(40, 256, 256), patch_size=(8,8,8), num_classes=8)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:53.919785Z","iopub.execute_input":"2024-03-31T09:28:53.92024Z","iopub.status.idle":"2024-03-31T09:28:54.671544Z","shell.execute_reply.started":"2024-03-31T09:28:53.920195Z","shell.execute_reply":"2024-03-31T09:28:54.670442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DeformableTransformerWithClassification(nn.Module):\n    def __init__(self, in_channels, img_size, patch_size, num_classes):\n        super(DeformableTransformerWithClassification, self).__init__()\n        self.features = U_ResTran3D(norm_cfg='BN', activation_cfg='ReLU', img_size=128, num_classes=3, weight_std=False)\n        self.avgpool = nn.AdaptiveAvgPool1d((1))  # Adaptive average pooling to get a fixed-size output\n        self.fc = nn.Linear(384, num_classes)  # Fully connected layer for classification\n\n    def forward(self, x):\n        x = self.features(x)  # Get the intermediate representations from the features module\n        print(x.shape, \"........\")  # Print the shape of the last intermediate representation\n        x = x.transpose(1, 2)\n        x = self.avgpool(x)  # Apply average pooling to the last intermediate representation\n        print(x.shape)  # Print the shape after average pooling\n        x = torch.flatten(x, 1)  # Flatten the feature map\n        print(x.shape)  # Print the shape after flattening\n        x = self.fc(x)  # Classification head\n        x = torch.sigmoid(x)  # Apply sigmoid activation for multi-label classification\n        return x\n    \n# model=DeformableTransformerWithClassification(in_channels=3,img_size=(40, 256, 256), patch_size=(8,8,8), num_classes=8)","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:54.674981Z","iopub.execute_input":"2024-03-31T09:28:54.675296Z","iopub.status.idle":"2024-03-31T09:28:54.685058Z","shell.execute_reply.started":"2024-03-31T09:28:54.675268Z","shell.execute_reply":"2024-03-31T09:28:54.684007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nfrom monai.networks.nets import densenet, resnet, senet, efficientnet\nclass MonaiModelWithClassification(nn.Module):\n    def __init__(self):\n        super(MonaiModelWithClassification, self).__init__()\n#         self.features = densenet.DenseNet(spatial_dims = 3, in_channels = 3, out_channels = 8)\n\n#         self.features = resnet.ResNet(spatial_dims=3, num_classes = 8, block = \"basic\", layers = [3, 4, 23, 3], block_inplanes = [64, 64, 64, 64])\n\n#         self.features = senet.SEResNext101(spatial_dims=3, num_classes = 8, in_channels = 3)\n\n#         self.features = senet.SEResNet152(spatial_dims=3, num_classes = 8, in_channels = 3)\n   \n        self.features = efficientnet.EfficientNetBN(\"efficientnet-b0\", spatial_dims=3, num_classes = 8)\n \n    def forward(self, x):\n        x = self.features(x)\n        x = torch.sigmoid(x)  # Apply sigmoid activation for multi-label classification\n        return x\n    \nmodel = MonaiModelWithClassification()\n# input=torch.rand(1,3,40,256,256)\n# output=model(input)\n# print(output.shape)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:54.686166Z","iopub.execute_input":"2024-03-31T09:28:54.686431Z","iopub.status.idle":"2024-03-31T09:28:54.840336Z","shell.execute_reply.started":"2024-03-31T09:28:54.686406Z","shell.execute_reply":"2024-03-31T09:28:54.839376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n@dataclass\nclass BatchSlice:\n    i_from: int\n    i_to: int\n    i_start: int\n\n\ndef get_slices(batch: Tensor, dim=1, window: int = 16, overlap: int = 8) -> List[BatchSlice]:\n    num_imgs = batch.size(dim)\n    if num_imgs <= window:\n        return [BatchSlice(0, num_imgs, 0)]\n    stride = window - overlap\n    result = []\n    current_idx = 0\n    while True:\n        next_idx = current_idx + window\n\n        if next_idx >= num_imgs:\n            current_idx = num_imgs - window\n            offset = overlap // 2 if current_idx > 0 else 0\n            next_idx = num_imgs\n            result.append(BatchSlice(current_idx, next_idx, offset))\n            break\n        else:\n            offset = overlap // 2 if current_idx > 0 else 0\n            result.append(BatchSlice(current_idx, next_idx, offset))\n        current_idx += stride\n    return result","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:54.841567Z","iopub.execute_input":"2024-03-31T09:28:54.8419Z","iopub.status.idle":"2024-03-31T09:28:54.853187Z","shell.execute_reply.started":"2024-03-31T09:28:54.841858Z","shell.execute_reply":"2024-03-31T09:28:54.851859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\ndef _read_labels():\n    labels_df = pd.read_csv(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\n    labels_dict = {}\n    for index, row in labels_df.iterrows():\n        cube_id = row['StudyInstanceUID']\n        overall_patient = row['patient_overall']\n        c1, c2, c3, c4, c5, c6, c7 = row['C1'], row['C2'], row['C3'], row['C4'], row['C5'], row['C6'], row['C7']\n        labels_dict[cube_id] = [overall_patient, c1, c2, c3, c4, c5, c6, c7]\n    return labels_dict\n\nlabels_dict = _read_labels()\n# print(labels_dict)","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:54.854751Z","iopub.execute_input":"2024-03-31T09:28:54.855815Z","iopub.status.idle":"2024-03-31T09:28:55.117414Z","shell.execute_reply.started":"2024-03-31T09:28:54.855775Z","shell.execute_reply":"2024-03-31T09:28:55.116385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\ndef _read_labels_abdominal_trauma():\n    labels_df = pd.read_csv(\"/kaggle/input/rsna-2023-abdominal-trauma-detection/train.csv\")\n    labels_dict = {}\n    for index, row in labels_df.iterrows():\n        cube_id = row['patient_id']\n        bowel_healthy, extravasation_injury, kidney_low, kidney_high, liver_low, liver_high, spleen_low, spleen_high = row['bowel_healthy'], row['extravasation_injury'], row['kidney_low'], row['kidney_high'], row['liver_low'], row['liver_high'], row['spleen_low'], row['spleen_high']\n        labels_dict[cube_id] = [bowel_healthy, extravasation_injury, kidney_low, kidney_high, liver_low, liver_high, spleen_low, spleen_high]\n    return labels_dict\n\nlabels_dict = _read_labels_abdominal_trauma()\n# print(labels_dict)","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:55.119187Z","iopub.execute_input":"2024-03-31T09:28:55.119659Z","iopub.status.idle":"2024-03-31T09:28:55.486865Z","shell.execute_reply.started":"2024-03-31T09:28:55.11961Z","shell.execute_reply":"2024-03-31T09:28:55.485946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import shutil  # Add this import statement\n\n# def process_segmentations(models: List[torch.nn.Module], test_dataset_dir: str, out_dir: str,  cases: List[str]):\n#     os.makedirs(out_dir, exist_ok=True)\n\n#     test_dataset = DatasetSeg(dataset_dir=test_dataset_dir, cases=cases)\n#     sampler = None\n#     oof_loader = DataLoader(\n#         test_dataset, batch_size=1, sampler=sampler, shuffle=False, num_workers=1, pin_memory=False\n#     )\n#     pred_dir = out_dir\n#     os.makedirs(pred_dir, exist_ok=True)\n#     for sample in tqdm(oof_loader):\n#         image = sample[\"image\"]\n#         print(image.shape)\n#         h = int(sample[\"h\"][0])\n#         cube_id = sample[\"cube_id\"][0]\n#         print(cube_id)\n#         imgs = image.cpu().float()\n#         case_preds = np.zeros((imgs.shape[2], 256, 256), dtype=np.float32)\n\n#         with torch.no_grad():\n#             slices = get_slices(imgs, dim=2, window=256, overlap=128)\n#             for slice in slices:\n#                 batch = imgs[:, :, slice.i_from:slice.i_to].cuda().float()\n#                 with torch.cuda.amp.autocast(enabled=True):\n#                     preds = None\n#                     for model in models:\n#                         if preds is None:\n#                             preds = torch.softmax(model(batch)[\"mask\"], dim=1)[0]\n#                         else:\n#                             preds += torch.softmax(model(batch)[\"mask\"], dim=1)[0]\n#                     preds = torch.argmax(preds, dim=0)\n#                 preds = preds.cpu().numpy()\n\n#                 for pred_idx in range(slice.i_start, preds.shape[0]):\n#                     idx = slice.i_from + pred_idx\n#                     y_pred = preds[pred_idx]\n#                     case_preds[idx] = y_pred[:, :]\n#                 torch.cuda.empty_cache()\n#         print(case_preds.shape, \"case_preds.shape\")\n#         case_preds = np.array(case_preds)[:h]\n#         case_preds = case_preds.astype(np.uint8)\n#         cube_id = cube_id.replace('/', '_')\n#         print(cube_id)\n#         tifffile.imwrite(os.path.join(pred_dir, f\"{cube_id}.tif\"), case_preds)\n\n        \n# import zoo\n# config_seg  = {\n#   \"network\": zoo.ResNet3dCSN2P1D,\n#   \"encoder_params\": {\n#     \"encoder\": \"r50ir\"\n#   }\n# }\n\n# def load_checkpoint(model, checkpoint_path, strict=False, verbose=True):\n#     if verbose:\n#         print(\"=> loading checkpoint '{}'\".format(checkpoint_path))\n#     checkpoint = torch.load(checkpoint_path, map_location='cpu')\n#     if 'state_dict' in checkpoint:\n#         state_dict = checkpoint['state_dict']\n#         state_dict = {re.sub(\"^module.\", \"\", k): w for k, w in state_dict.items()}\n#         orig_state_dict = model.state_dict()\n#         mismatched_keys = []\n#         for k, v in state_dict.items():\n#             ori_size = orig_state_dict[k].size() if k in orig_state_dict else None\n#             if v.size() != ori_size:\n#                 if verbose:\n#                     print(\"SKIPPING!!! Shape of {} changed from {} to {}\".format(k, v.size(), ori_size))\n#                 mismatched_keys.append(k)\n#         for k in mismatched_keys:\n#             del state_dict[k]\n#         model.load_state_dict(state_dict, strict=strict)\n#         del state_dict\n#         del orig_state_dict\n#         print(\"=> loaded checkpoint '{}' (epoch {})\"\n#               .format(checkpoint_path, checkpoint['epoch']))\n#     else:\n#         model.load_state_dict(checkpoint)\n#     del checkpoint\n\n# def load_model(conf: Dict, checkpoint: str):\n#     model = conf[\"network\"](**conf[\"encoder_params\"])\n#     model = model.cuda()\n#     load_checkpoint(model, checkpoint)\n#     return model.eval()\n\n# test_dataset_dir = \"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/test_images/\"\n# cases = os.listdir(test_dataset_dir)\n# seg_model = load_model(config_seg, \"/kaggle/input/rsna-weights/256_ResNet3dCSN2P1D_r50ir_0_dice\")\n# ds = DatasetSeg(test_dataset_dir, cases)\n\n# train_preds_dir = \"/kaggle/working/train_seg_preds\"\n# train_dataset_dir = \"/kaggle/input/rsna-2023-abdominal-trauma-detection/train_images\"\n# # Function to recursively get all subfolders within a directory\n\n# # Initialize a list to store subfolders\n# subfolders = []\n\n# # Function to recursively get all subfolders within a directory\n# def get_subfolders(directory):\n# #     subfolders = []\n#     for item in os.listdir(directory):\n#         item_path = os.path.join(directory, item)\n#         if os.path.isdir(item_path):\n#             if (item_path.count('/') == 6):\n# #                 print(item_path)\n# #                 print(item_path.split('/', 5)[-1])\n#                 subfolders.append(item_path.split('/', 5)[-1])\n# #                 print(item_path.split('/', 5)[-1])\n#             get_subfolders(item_path)\n       \n#     return subfolders\n\n# cases = get_subfolders(train_dataset_dir)\n\n# cases1 = cases[0:500]\n# print(cases1)\n# # process_segmentations([seg_model], train_dataset_dir,out_dir=\"/kaggle/working/train_seg_preds_1\", cases=cases1)\n\n# import zipfile\n# # Zip the folder\n# zip_file_name = \"/kaggle/working/masks_1.zip\"\n# with zipfile.ZipFile(zip_file_name, 'w', zipfile.ZIP_DEFLATED) as zipf:\n#     for root, _, files in os.walk(\"/kaggle/working/train_seg_preds_1\"):\n#         for file in files:\n#             zipf.write(os.path.join(root, file), arcname=file)\n            \n# # Delete the folder after zipping\n# shutil.rmtree(\"/kaggle/working/train_seg_preds_1\")\n    \n# cases2 = cases[500:1000]\n# print(cases2)\n# # process_segmentations([seg_model], train_dataset_dir,out_dir=\"/kaggle/working/train_seg_preds_2\", cases=cases2)\n\n# import zipfile\n# # Zip the folder\n# zip_file_name = \"/kaggle/working/masks_2.zip\"\n# with zipfile.ZipFile(zip_file_name, 'w', zipfile.ZIP_DEFLATED) as zipf:\n#     for root, _, files in os.walk(\"/kaggle/working/train_seg_preds_2\"):\n#         for file in files:\n#             zipf.write(os.path.join(root, file), arcname=file)\n            \n# shutil.rmtree(\"/kaggle/working/train_seg_preds_2\")\n# cases3 = cases[1000:1500]\n# print(cases3)\n# # process_segmentations([seg_model], train_dataset_dir,out_dir=\"/kaggle/working/train_seg_preds_3\", cases=cases3)\n\n# import zipfile\n# # Zip the folder\n# zip_file_name = \"/kaggle/working/masks_3.zip\"\n# with zipfile.ZipFile(zip_file_name, 'w', zipfile.ZIP_DEFLATED) as zipf:\n#     for root, _, files in os.walk(\"/kaggle/working/train_seg_preds_3\"):\n#         for file in files:\n#             zipf.write(os.path.join(root, file), arcname=file)\n            \n# shutil.rmtree(\"/kaggle/working/train_seg_preds_3\")            \n# cases4 = cases[1500:2000]\n# print(cases4)\n# # process_segmentations([seg_model], train_dataset_dir,out_dir=\"/kaggle/working/train_seg_preds_4\", cases=cases4)\n\n# import zipfile\n# # Zip the folder\n# zip_file_name = \"/kaggle/working/masks_4.zip\"\n# with zipfile.ZipFile(zip_file_name, 'w', zipfile.ZIP_DEFLATED) as zipf:\n#     for root, _, files in os.walk(\"/kaggle/working/train_seg_preds_4\"):\n#         for file in files:\n#             zipf.write(os.path.join(root, file), arcname=file)\n            \n# shutil.rmtree(\"/kaggle/working/train_seg_preds_4\")\n# cases5 = cases[2000:2500]\n# print(cases5)\n# # process_segmentations([seg_model], train_dataset_dir,out_dir=\"/kaggle/working/train_seg_preds_5\", cases=cases5)\n\n# import zipfile\n# # Zip the folder\n# zip_file_name = \"/kaggle/working/masks_5.zip\"\n# with zipfile.ZipFile(zip_file_name, 'w', zipfile.ZIP_DEFLATED) as zipf:\n#     for root, _, files in os.walk(\"/kaggle/working/train_seg_preds_5\"):\n#         for file in files:\n#             zipf.write(os.path.join(root, file), arcname=file)\n            \n# shutil.rmtree(\"/kaggle/working/train_seg_preds_5\")            \n# cases6 = cases[2500:3000]\n# print(cases6)\n# # process_segmentations([seg_model], train_dataset_dir,out_dir=\"/kaggle/working/train_seg_preds_6\", cases=cases6)\n\n# import zipfile\n# # Zip the folder\n# zip_file_name = \"/kaggle/working/masks_6.zip\"\n# with zipfile.ZipFile(zip_file_name, 'w', zipfile.ZIP_DEFLATED) as zipf:\n#     for root, _, files in os.walk(\"/kaggle/working/train_seg_preds_6\"):\n#         for file in files:\n#             zipf.write(os.path.join(root, file), arcname=file)\n            \n# shutil.rmtree(\"/kaggle/working/train_seg_preds_6\")\n# cases7 = cases[3000:3500]\n# print(cases7)\n# # process_segmentations([seg_model], train_dataset_dir,out_dir=\"/kaggle/working/train_seg_preds_7\", cases=cases7)\n\n# import zipfile\n# # Zip the folder\n# zip_file_name = \"/kaggle/working/masks_7.zip\"\n# with zipfile.ZipFile(zip_file_name, 'w', zipfile.ZIP_DEFLATED) as zipf:\n#     for root, _, files in os.walk(\"/kaggle/working/train_seg_preds_7\"):\n#         for file in files:\n#             zipf.write(os.path.join(root, file), arcname=file)\n            \n# shutil.rmtree(\"/kaggle/working/train_seg_preds_7\")            \n# cases8 = cases[3500:4000]\n# print(cases8)\n# # process_segmentations([seg_model], train_dataset_dir,out_dir=\"/kaggle/working/train_seg_preds_8\", cases=cases8)\n\n# import zipfile\n# # Zip the folder\n# zip_file_name = \"/kaggle/working/masks_8.zip\"\n# with zipfile.ZipFile(zip_file_name, 'w', zipfile.ZIP_DEFLATED) as zipf:\n#     for root, _, files in os.walk(\"/kaggle/working/train_seg_preds_8\"):\n#         for file in files:\n#             zipf.write(os.path.join(root, file), arcname=file)\n\n            \n# shutil.rmtree(\"/kaggle/working/train_seg_preds_8\")\n# cases9 = cases[4000:4500]\n# print(cases9)\n# # process_segmentations([seg_model], train_dataset_dir,out_dir=\"/kaggle/working/train_seg_preds_9\", cases=cases9)\n\n# import zipfile\n# # Zip the folder\n# zip_file_name = \"/kaggle/working/masks_9.zip\"\n# with zipfile.ZipFile(zip_file_name, 'w', zipfile.ZIP_DEFLATED) as zipf:\n#     for root, _, files in os.walk(\"/kaggle/working/train_seg_preds_9\"):\n#         for file in files:\n#             zipf.write(os.path.join(root, file), arcname=file)\n\n# shutil.rmtree(\"/kaggle/working/train_seg_preds_9\")\n# cases10 = cases[4500:]\n# print(cases10)\n# # process_segmentations([seg_model], train_dataset_dir,out_dir=\"/kaggle/working/train_seg_preds_10\", cases=cases10)\n\n# import zipfile\n# # Zip the folder\n# zip_file_name = \"/kaggle/working/masks_10.zip\"\n# with zipfile.ZipFile(zip_file_name, 'w', zipfile.ZIP_DEFLATED) as zipf:\n#     for root, _, files in os.walk(\"/kaggle/working/train_seg_preds_10\"):\n#         for file in files:\n#             zipf.write(os.path.join(root, file), arcname=file)\n            \n# shutil.rmtree(\"/kaggle/working/train_seg_preds_10\")","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:55.488624Z","iopub.execute_input":"2024-03-31T09:28:55.489093Z","iopub.status.idle":"2024-03-31T09:28:55.506187Z","shell.execute_reply.started":"2024-03-31T09:28:55.489048Z","shell.execute_reply":"2024-03-31T09:28:55.504971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import random_split\nimport numpy as np\nfrom torch.utils.data import Subset \n\n# Function to recursively get all subfolders within a directory\ndef get_subfolders_2(directory):\n#     subfolders = []\n    for item in os.listdir(directory):\n        item_path = os.path.join(directory, item)\n        if os.path.isdir(item_path):\n            if (item_path.count('/') == 6):\n#                 print(item_path)\n#                 print(item_path.split('/', 5)[-1])\n                x = item_path.split('/', 5)[-1]\n                x = x.replace('/', '_')\n#                 print(x)\n                subfolders_2.append(x)\n#                 print(item_path.split('/', 5)[-1])\n            get_subfolders_2(item_path)\n       \n    return subfolders_2\n\ndataset_dir = \"/kaggle/input/rsna-2023-abdominal-trauma-detection/train_images\"\n# cases = os.listdir(dataset_dir)\nsubfolders_2 = []\n\ncases = get_subfolders_2(dataset_dir)\nprint(cases)","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:32:10.283581Z","iopub.execute_input":"2024-03-31T09:32:10.284808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def combine_scan(scan_dir: str, size=512, fix_monochrome: bool = True) -> np.ndarray:\n    num_files = len(os.listdir(scan_dir))\n    images = []\n    offset = 0\n    first = None\n    last = None\n    files = []\n    for i in range(num_files):\n        dpath = os.path.join(scan_dir, f\"{i + offset}.dcm\")\n        if i == 0:\n            while not os.path.exists(dpath):\n                offset += 1\n                dpath = os.path.join(scan_dir, f\"{i + offset}.dcm\")\n        files.append(dpath)\n\n    for dpath in files[::2]:\n        ds = pydicom.dcmread(dpath)\n        if not first:\n            first = ds\n        last = ds\n\n        data = ds.pixel_array\n        data = cv2.resize(data, (size, size))\n        if fix_monochrome and ds.PhotometricInterpretation == \"MONOCHROME1\":\n            data = np.amax(data) - data\n        images.append(data)\n\n    if first and last:\n        if last.ImagePositionPatient[2] > first.ImagePositionPatient[2]:\n            images = images[::-1]\n    return np.array(images)\n\nclass DatasetSeg(Dataset):\n    def __init__(\n            self,\n            dataset_dir: str,\n            cases: List[str],\n    ):\n        self.dataset_dir = dataset_dir\n        self.cases = cases\n\n    def __getitem__(self, i):\n        cube_id = self.cases[i]\n        image_cube = combine_scan(os.path.join(self.dataset_dir, cube_id.replace(\"_\", \"/\")), size=256)\n        image_mean = image_cube.mean()\n        image_std = image_cube.std()\n        h = image_cube.shape[0]\n\n        images = image_cube\n        if h % 32 > 0:\n            tmp = np.zeros(((h // 32 + 1) * 32, 256, 256))\n            tmp[:h] = images\n            images = tmp\n        images = (images - image_mean) / image_std\n        images = np.expand_dims(images, 0)\n        sample = {}\n        sample['image'] = torch.from_numpy(images).float()\n        sample['cube_id'] = cube_id\n        sample['h'] = h\n        return sample\n\n    def __len__(self):\n        return len(self.cases)\n\n\ncrop_augs =  albumentations.ReplayCompose([\n            albumentations.LongestMaxSize(256),\n            albumentations.PadIfNeeded(256, 256, border_mode=cv2.BORDER_CONSTANT),\n        ])\n\nclass DatasetCrops(Dataset):\n    def __init__(\n            self,\n            dataset_dir: str,\n            cases: List[str],\n            transforms=crop_augs,\n            slice_size=40,\n    ):\n        self.dataset_dir = dataset_dir\n        self.transforms = transforms\n        self.slice_size = slice_size\n        self.cases = cases\n\n    def __getitem__(self, i):\n        cube_id = self.cases[i]\n#         mask_cube = tifffile.imread(os.path.join(\"seg_preds\", f\"{cube_id}.tif\"))\n        mask_cube_path = os.path.join(\"/kaggle/input/mask-1-trauma/masks_1\", f\"{cube_id}.tif\")\n        if os.path.exists(mask_cube_path):\n            mask_cube = tifffile.imread(mask_cube_path)\n        else:\n            # If the file is not found in the first directory, try the second directory\n            mask_cube_path = os.path.join(\"/kaggle/input/masks-2-trauma/masks_2\", f\"{cube_id}.tif\")\n            if os.path.exists(mask_cube_path):\n                mask_cube = tifffile.imread(mask_cube_path)\n            else:\n                mask_cube_path = os.path.join(\"/kaggle/input/mask-3-trauma/masks_3\", f\"{cube_id}.tif\")\n                if os.path.exists(mask_cube_path):\n                    mask_cube = tifffile.imread(mask_cube_path)\n                else:\n                    mask_cube_path = os.path.join(\"/kaggle/input/mask-4-trauma/masks_4\", f\"{cube_id}.tif\")\n                    if os.path.exists(mask_cube_path):\n                        mask_cube = tifffile.imread(mask_cube_path)\n                    else:\n                        mask_cube_path = os.path.join(\"/kaggle/input/mask-5-trauma/masks_5\", f\"{cube_id}.tif\")\n                        if os.path.exists(mask_cube_path):\n                            mask_cube = tifffile.imread(mask_cube_path)\n                        else:\n                            mask_cube_path = os.path.join(\"/kaggle/input/masks-6-trauma/masks_6\", f\"{cube_id}.tif\")\n                            if os.path.exists(mask_cube_path):\n                                mask_cube = tifffile.imread(mask_cube_path)\n                            else: \n                                mask_cube_path = os.path.join(\"/kaggle/input/mask-7-trauma/masks_7\", f\"{cube_id}.tif\")\n                                if os.path.exists(mask_cube_path):\n                                    mask_cube = tifffile.imread(mask_cube_path)\n                                else: \n                                    mask_cube_path = os.path.join(\"/kaggle/input/mask-8-trauma/masks_8\", f\"{cube_id}.tif\")\n                                    if os.path.exists(mask_cube_path):\n                                        mask_cube = tifffile.imread(mask_cube_path)\n                                    else: \n                                        mask_cube_path = os.path.join(\"/kaggle/input/mask-9-trauma/masks_9\", f\"{cube_id}.tif\")\n                                        if os.path.exists(mask_cube_path):\n                                            mask_cube = tifffile.imread(mask_cube_path)\n                                        else: \n                                            mask_cube_path = os.path.join(\"/kaggle/input/mask-10-trauma/masks_10\", f\"{cube_id}.tif\")\n                                            if os.path.exists(mask_cube_path):\n                                                mask_cube = tifffile.imread(mask_cube_path)\n                                            else: \n                                                mask_cube = torch.rand(256, 256, 256)\n                                                mask_cube = mask_cube.cpu().numpy().astype(np.int)\n                                                print(\"Error: mask_cube.tif not found in both directories.\")\n\n        image_cube = combine_scan(os.path.join(self.dataset_dir, cube_id.replace(\"_\", \"/\")) ,size=512)\n        boxes = {}\n        for rprop in measure.regionprops(mask_cube):\n            boxes[rprop.label] = rprop.bbox, rprop.area\n\n        image_mean = image_cube.mean()\n        image_std = image_cube.std()\n        slice_size = self.slice_size\n        all_images = []\n#         labels = np.zeros((8,))\n#         print(cube_id)\n        x = cube_id.split(\"_\")[0].strip()\n#         print(\".......\", x, \".................................................\")\n        labels = labels_dict[int(x)]\n        print(labels)\n        for li in range(1, 8):\n            if li not in boxes:\n                all_images.append(np.zeros((3, self.slice_size, 256, 256)))\n            else:\n                bbox, area = boxes[li]\n                z1, z2 = bbox[0], bbox[3]\n                y1, y2 = max(bbox[1] - 16, 0), min(bbox[4] + 16, 256)\n                x1, x2 = max(bbox[2] - 16, 0), min(bbox[5] + 16, 256)\n                # if z2 - z1 < slice_size:\n                #     z1 = random.randint(max(z2 - slice_size, 0), z1)\n                #     z2 = z1 + slice_size\n                # todo: verify\n                if z2 - z1 < slice_size:\n                    diff = (slice_size - z2 + z1) // 2\n                    z1 = max(0, z1 - diff)\n                    z2 = z1 + slice_size\n                images = image_cube[z1:z2, y1 * 2:y2 * 2, x1 * 2:x2 * 2].copy()\n                masks = mask_cube[z1:z2, y1:y2, x1:x2].copy()\n                slice_size = self.slice_size\n\n                replay = None\n                image_crops = []\n                mask_crops = []\n                for i in range(images.shape[0]):\n                    image = images[i]\n                    mask = masks[i]\n                    h, w, = mask.shape\n                    mask = cv2.resize(mask, (w * 2, h * 2), interpolation=cv2.INTER_NEAREST)\n                    if replay is None:\n                        sample = self.transforms(image=image, mask=mask)\n                        replay = sample[\"replay\"]\n                    else:\n                        sample = ReplayCompose.replay(replay, image=image, mask=mask)\n                    image_ = sample[\"image\"]\n                    image_crops.append(image_)\n                    mask_crops.append(sample[\"mask\"])\n                images = np.array(image_crops).astype(np.float32)\n                masks = np.array(mask_crops).astype(np.float32)\n                images = np.expand_dims(images, -1)\n                masks = np.expand_dims(masks, -1)\n                images = (images - image_mean) / image_std\n\n                images = np.concatenate([images, images, masks], axis=-1)\n                h = images.shape[0]\n                if h > slice_size:\n                    images = images[: slice_size]\n                    all_images.append(np.moveaxis(images, -1, 0))\n                    images = images[-slice_size:]\n                    all_images.append(np.moveaxis(images, -1, 0))\n                else:\n                    if h != slice_size:\n                        tmp = np.zeros((slice_size, *images.shape[1:]))\n                        tmp[:h] = images\n                        images = tmp\n                    all_images.append(np.moveaxis(images, -1, 0))\n\n        sample = {}\n        sample['image'] = torch.from_numpy(np.array(all_images)).float()\n        sample['label'] = torch.from_numpy(np.array(labels)).float()\n        sample['cube_id'] = cube_id\n        return sample\n\n    def __len__(self):\n        return len(self.cases)","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:56.094222Z","iopub.status.idle":"2024-03-31T09:28:56.094606Z","shell.execute_reply.started":"2024-03-31T09:28:56.094424Z","shell.execute_reply":"2024-03-31T09:28:56.094442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"validation_split = 0.1\ntest_split = 0.1\n\n# cases = cases.replace('/', '_')\n\n# Create dataset\ndataset = DatasetCrops(dataset_dir=dataset_dir, cases=cases)\n\n# Compute sizes\ntrain_size = int(len(dataset) * (1 - validation_split - test_split))\nval_test_size = len(dataset) - train_size\nval_size = int(val_test_size / 2)\ntest_size = val_test_size - val_size\n\n# Random split into train, validation, and test datasets\ntrain_dataset, val_test_dataset = random_split(dataset, [train_size, val_test_size])\nval_dataset, test_dataset = random_split(val_test_dataset, [val_size, test_size])\n\n# train_dataset = Subset(train_dataset, range(5))\n# val_dataset = Subset(val_dataset, range(5))\n# test_dataset = Subset(test_dataset, range(5))\n\n# Create data loaders\ntrain_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=1, shuffle=False)\ntest_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\nprint(len(train_loader), len(val_loader), len(test_loader)) \n\n# train_loader = train_loader[:100]\n","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:56.096948Z","iopub.status.idle":"2024-03-31T09:28:56.097346Z","shell.execute_reply.started":"2024-03-31T09:28:56.097155Z","shell.execute_reply":"2024-03-31T09:28:56.097174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import nn\nimport torch.optim as optim\nfrom sklearn.metrics import accuracy_score\nfrom torch.utils.data import random_split\n\nlearning_rate = 0.001  # Define the learning rate\n\ndef train_model(model, imgs):\n        preds = []\n        with torch.no_grad():\n            for i in range(len(imgs)):\n                with torch.cuda.amp.autocast():\n#                     output = model(imgs[i:i + 1])[\"cls\"][0]\n                    output = model(imgs[i:i + 1])\n                pred_slice = torch.sigmoid(output.float()).cpu().numpy().astype(np.float32)\n#                 with torch.cuda.amp.autocast():\n#                     output = model(torch.flip(imgs[i:i + 1], dims=(-1,)))\n                pred_slice += torch.sigmoid(output.float()).cpu().numpy().astype(np.float32)\n                pred_slice /= 2\n                preds.append(pred_slice)\n                break\n        preds = np.max(np.array(preds), axis=0)\n        preds[np.isnan(preds)] = 0.01\n        return preds\n\ndef train_classification(model, cases: List, num_epochs: int = 1):\n#     dataset = DatasetCrops(dataset_dir=dataset_dir, cases=cases)\n#     train_size = int(len(dataset) * (1 - validation_split))\n#     val_size = len(dataset) - train_size\n#     train_dataset, val_dataset = random_split(dataset, [train_size, val_size])\n#     train_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)\n#     val_loader = DataLoader(val_dataset, batch_size=1, shuffle=False)\n    \n    \n    # Define loss criterion for multi-label classification\n    loss_criterion = nn.BCEWithLogitsLoss()  # You can adjust the loss function as needed\n    best_metric = -1\n    best_metric_epoch = -1\n    best_metrics_epochs_and_time = [[], [], []]\n    total_start=time.time()\n    for epoch in range(num_epochs):\n        print(f\"Epoch [{epoch + 1}/{num_epochs}]\")\n        \n        # Training phase\n        train_loss = 0.0\n        train_preds = []\n        train_labels = []\n        \n        for sample in tqdm(train_loader, desc=\"Training\"):\n            imgs = sample[\"image\"].cuda().float()[0]\n            cube_id = sample[\"cube_id\"][0]\n            with torch.no_grad():\n                preds = []\n#                 for model in models:\n                model.train()  # Set the model to training mode\n                x = train_model(model, imgs)\n                preds.append(x.squeeze())\n                preds = np.average(np.array(preds), axis=0)\n                preds = np.clip(preds, 0.01, 0.99)\n                \n#                 print(sample['label'].shape)\n#                 print(torch.tensor(preds).shape)\n                loss = loss_criterion(torch.tensor(preds).squeeze(), sample['label'].squeeze())  # Calculate loss\n                loss.requires_grad = True\n                print(loss)\n                # Backpropagation and optimization (assuming models are trainable)\n#                 for model in models:\n#                     model.train()  # Set the model to training mode\n                model.zero_grad()  # Zero the gradients\n                loss.backward()  # Backpropagate the gradients\n                optimizer = optim.Adam(model.parameters(), lr=learning_rate)  # Define optimizer\n                optimizer.step()  # Update model parameters\n                \n                # Compute metrics\n                train_loss += loss.item() * imgs.size(0)\n                train_preds.extend(preds)\n                train_labels.extend(sample['label'].squeeze().cpu().numpy())\n        \n#         print(train_labels.shape, train_preds.shape)\n        train_loss /= len(train_loader.dataset)\n        train_accuracy = accuracy_score(np.array(train_labels), np.array(train_preds) >= 0.5)\n        print(f\"Train Loss: {train_loss:.4f} | Train Accuracy: {train_accuracy:.4f}\")\n        \n        # Validation phase\n        val_loss = 0.0\n        val_preds = []\n        val_labels = []\n        model.eval()  # Set the model to evaluation mode\n        with torch.no_grad():\n            for val_sample in tqdm(val_loader, desc=\"Validation\"):\n                val_imgs = val_sample[\"image\"].cuda().float()[0]\n                val_preds_batch = []\n#                 for val_model in models:\n                val_x = train_model(model, val_imgs)\n                val_preds_batch.append(val_x.squeeze())\n                val_preds_batch = np.average(np.array(val_preds_batch), axis=0)\n                val_preds_batch = np.clip(val_preds_batch, 0.01, 0.99)\n                val_loss += loss_criterion(torch.tensor(val_preds_batch), val_sample['label'].squeeze()).item() * val_imgs.size(0)\n                val_preds.extend(val_preds_batch)\n                val_labels.extend(val_sample['label'].squeeze().cpu().numpy())\n        \n        val_loss /= len(val_loader.dataset)\n        val_accuracy = accuracy_score(np.array(val_labels), np.array(val_preds) >= 0.5)\n        print(f\"Validation Loss: {val_loss:.4f} | Validation Accuracy: {val_accuracy:.4f}\")\n\n        \n        if val_accuracy > best_metric:\n                best_metric = val_accuracy\n                best_metric_epoch = epoch + 1\n                best_metrics_epochs_and_time[0].append(best_metric)\n                best_metrics_epochs_and_time[1].append(best_metric_epoch)\n                best_metrics_epochs_and_time[2].append(time.time() - total_start)\n        torch.save(\n            model.state_dict(),\n            os.path.join(\"/kaggle/working/\", f\"swin_trans_best_metric_model_{epoch+1}.pth\"),\n        )\n        print(\"saved new best metric model\")\n        chk_file_name = \"swin_trans_epoch_\" + str(epoch+1) + \"_model.pth\"    \n        torch.save(\n            model.state_dict(),\n            os.path.join(\"/kaggle/working/\", chk_file_name),\n        )\n","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:56.098935Z","iopub.status.idle":"2024-03-31T09:28:56.099297Z","shell.execute_reply.started":"2024-03-31T09:28:56.099119Z","shell.execute_reply":"2024-03-31T09:28:56.099136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"output = []\n\nstart_epoch = 0\nlatest_checkpoint_path = \"/kaggle/input/chk-points-swin-transformer/swin_trans_best_metric_model_1.pth\"\nif os.path.exists(latest_checkpoint_path):\n    checkpoint = torch.load(latest_checkpoint_path)\n    model.load_state_dict(checkpoint)\n    start_epoch = 2\n    print(\"checkpoint loaded\")\nelse:\n    print(\"No checkpoint. Starting from scratch\")\n# print(model)\ndevice=torch.device(\"cuda:0\")\nmodel = model.to(device)\ntrain_classification(model, cases, 1)","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:56.100757Z","iopub.status.idle":"2024-03-31T09:28:56.101193Z","shell.execute_reply.started":"2024-03-31T09:28:56.100987Z","shell.execute_reply":"2024-03-31T09:28:56.101017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import List\nimport torch\nimport numpy as np\nfrom sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, roc_auc_score\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\n\ndef predict_classification(models: List[nn.Module]):\n#     test_dataset = DatasetCrops(dataset_dir=test_dataset_dir, cases=cases)\n#     dataloader = DataLoader(\n#         test_dataset, batch_size=1, sampler=None, shuffle=False, num_workers=1, pin_memory=False\n#     )\n    \n    def predict_model(model, imgs):\n        preds = []\n        with torch.no_grad():\n            for i in range(len(imgs)):\n                with torch.cuda.amp.autocast():\n                    output = model(imgs[i:i + 1])\n                pred_slice = torch.sigmoid(output.float()).cpu().numpy().astype(np.float32)\n                with torch.cuda.amp.autocast():\n                    output = model(torch.flip(imgs[i:i + 1], dims=(-1,)))\n                pred_slice += torch.sigmoid(output.float()).cpu().numpy().astype(np.float32)\n                pred_slice /= 2\n                preds.append(pred_slice)\n        preds = np.max(np.array(preds), axis=0)\n        preds[np.isnan(preds)] = 0.01\n        return preds\n        \n    output_list = []  # Initialize the output list\n    all_labels = []\n    all_predictions = []\n    print(\".......................................................................\", len(test_loader))\n    for sample in tqdm(test_loader):\n        imgs = sample[\"image\"].cuda().float()[0]\n        cube_id = sample[\"cube_id\"][0]\n        labels = sample[\"label\"].cpu().numpy()\n        # Option 1: Remove outer list\n        labels = labels[0]\n        print(\"len(labels), labels\", len(labels), labels)\n        all_labels.extend(labels)\n        \n        with torch.no_grad():\n            preds = []\n            \n            preds.append(predict_model(model, imgs))  # Pass imgs to predict_model\n            preds = np.average(np.array(preds), axis=0)\n            preds = np.clip(preds, 0.01, 0.99)\n            print(preds)\n            all_predictions.extend(preds.squeeze())\n#             output_list.append([cube_id, preds])\n    \n    # Calculate metrics\n    \n    accuracy = accuracy_score(all_labels, (np.array(all_predictions) >= 0.5).astype(int))\n    precision = precision_score(all_labels, (np.array(all_predictions) >= 0.5).astype(int))\n    recall = recall_score(all_labels, (np.array(all_predictions) >= 0.5).astype(int))\n    f1 = f1_score(all_labels, (np.array(all_predictions) >= 0.5).astype(int))\n    print(\"accuracy, precision, recall, f1\", accuracy, precision, recall, f1)\n   # auc = roc_auc_score(all_labels, (np.array(all_predictions)))\n    \n    return output_list, accuracy, precision, recall, f1\n","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:56.102382Z","iopub.status.idle":"2024-03-31T09:28:56.102738Z","shell.execute_reply.started":"2024-03-31T09:28:56.102561Z","shell.execute_reply":"2024-03-31T09:28:56.102578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"output = []\n# cases = os.listdir(train_dataset_dir)\npredict_classification(model)","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:56.104748Z","iopub.status.idle":"2024-03-31T09:28:56.105256Z","shell.execute_reply.started":"2024-03-31T09:28:56.105004Z","shell.execute_reply":"2024-03-31T09:28:56.105027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\nshutil.rmtree('/kaggle/working/seg_preds')","metadata":{"execution":{"iopub.status.busy":"2024-03-31T09:28:56.106345Z","iopub.status.idle":"2024-03-31T09:28:56.106842Z","shell.execute_reply.started":"2024-03-31T09:28:56.106574Z","shell.execute_reply":"2024-03-31T09:28:56.1066Z"},"trusted":true},"execution_count":null,"outputs":[]}]}