{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":52254,"databundleVersionId":6863140,"sourceType":"competition"},{"sourceId":6523471,"sourceType":"datasetVersion","datasetId":3771357},{"sourceId":6524344,"sourceType":"datasetVersion","datasetId":3771912},{"sourceId":7432254,"sourceType":"datasetVersion","datasetId":4325089},{"sourceId":7665407,"sourceType":"datasetVersion","datasetId":4470305}],"dockerImageVersionId":30626,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install monai\n!python -c \"import monai\" || pip install -q \"monai-weekly[nibabel]\"\n!python -c \"import matplotlib\" || pip install -q matplotlib\n%matplotlib inline\n!pip install einops","metadata":{"execution":{"iopub.status.busy":"2024-02-20T19:10:42.688426Z","iopub.execute_input":"2024-02-20T19:10:42.68887Z","iopub.status.idle":"2024-02-20T19:12:18.462937Z","shell.execute_reply.started":"2024-02-20T19:10:42.688834Z","shell.execute_reply":"2024-02-20T19:12:18.461682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nimport itertools\nfrom collections.abc import Sequence\nfrom monai.networks.layers import Conv, trunc_normal_\nclass WindowAttention(nn.Module):\n\n    def __init__(\n        self,\n        dim: int,\n        num_heads: int,\n        window_size: Sequence[int]=(7,7,7),\n        qkv_bias: bool = False,\n        attn_drop: float = 0.0,\n        proj_drop: float = 0.0,\n    ) -> None:\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=None):\n        b, n, c = x.shape\n        print(x.shape, b, n, c, \"...........................................\")\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        print(q.shape, k.shape, v.shape, \"q, k, v\")\n        q = q * self.scale\n        attn = q @ k.transpose(-2, -1)\n        print(attn.shape, \"attn\")\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        print(attn.shape, \"attn.shape#########################################################\")\n        x = (attn @ v).transpose(1, 2).reshape(b, n, c)\n        print(x.shape, \"@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@\")\n        x = self.proj(x)\n        print(x.shape, \"After proj\")\n        x = self.proj_drop(x)\n        return x\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-02-20T19:15:09.504628Z","iopub.execute_input":"2024-02-20T19:15:09.505356Z","iopub.status.idle":"2024-02-20T19:15:54.488377Z","shell.execute_reply.started":"2024-02-20T19:15:09.505315Z","shell.execute_reply":"2024-02-20T19:15:54.486999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\n# from models.TransBTS.IntmdSequential import IntermediateSequential \nclass IntermediateSequential(nn.Sequential):\n    def __init__(self, *args, return_intermediate=True):\n        super().__init__(*args)\n        self.return_intermediate = return_intermediate\n\n    def forward(self, input):\n        if not self.return_intermediate:\n            return super().forward(input)\n\n        intermediate_outputs = {}\n        output = input\n        for name, module in self.named_children():\n            output = intermediate_outputs[name] = module(output)\n\n        return output, intermediate_outputs\nclass SelfAttention(nn.Module):\n    def __init__(\n        self, dim, heads=8, qkv_bias=False, qk_scale=None, dropout_rate=0.0\n    ):\n        super().__init__()\n        self.num_heads = heads\n        head_dim = dim // heads\n        self.scale = qk_scale or head_dim ** -0.5\n\n        self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)\n        self.attn_drop = nn.Dropout(dropout_rate)\n        self.proj = nn.Linear(dim, dim)\n        self.proj_drop = nn.Dropout(dropout_rate)\n\n    def forward(self, x):\n        B, N, C = x.shape\n#         print(B, N, C, \"B, N, C\")\n        qkv = (\n            self.qkv(x)\n            .reshape(B, N, 3, self.num_heads, C // self.num_heads)\n            .permute(2, 0, 3, 1, 4)\n        )\n#         print(x.shape)\n#         print(self.qkv(x).shape, \"self.qkv.shape\")\n        \n        q, k, v = (\n            qkv[0],\n            qkv[1],\n            qkv[2],\n        )  # make torchscript happy (cannot use tensor as tuple)\n        \n#         print(q.shape, k.shape, v.shape, \"q.shape, k.shape, v.shape\")\n        attn = (q @ k.transpose(-2, -1)) * self.scale\n#         print(attn.shape, \"attn.shape\")\n        attn = attn.softmax(dim=-1)\n        attn = self.attn_drop(attn)\n\n        x = (attn @ v).transpose(1, 2).reshape(B, N, C)\n#         print(x.shape, \"after multiplication with attn and V\")\n        x = self.proj(x)\n#         print(x.shape, \"after proj\")\n        x = self.proj_drop(x)\n#         print(x.shape, \"after proj drop\")\n        return x\n\n\nclass Residual(nn.Module):\n    def __init__(self, fn):\n        super().__init__()\n        self.fn = fn\n\n    def forward(self, x):\n#         print(\"In residual\", x.shape)\n#         print(self.fn, \"self.fn\")\n        return self.fn(x) + x\n\n\nclass PreNorm(nn.Module):\n    def __init__(self, dim, fn):\n        super().__init__()\n        self.norm = nn.LayerNorm(dim)\n        self.fn = fn\n\n    def forward(self, x):\n#         print(\"in PreNorm\", self.fn)\n#         print()\n        return self.fn(self.norm(x))\n\n\nclass PreNormDrop(nn.Module):\n    def __init__(self, dim, dropout_rate, fn):\n        super().__init__()\n        self.norm = nn.LayerNorm(dim)\n        self.dropout = nn.Dropout(p=dropout_rate)\n        self.fn = fn\n\n    def forward(self, x):\n#         print(\"In PreNormDrop\")\n#         print(self.fn, \"self.fn\")\n#         print(x.shape, \"x.shape\")\n#         print(self.norm(x).shape, \"self.norm(x)\")\n        return self.dropout(self.fn(self.norm(x)))\n\n\nclass FeedForward(nn.Module):\n    def __init__(self, dim, hidden_dim, dropout_rate):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(dim, hidden_dim),\n            nn.GELU(),\n            nn.Dropout(p=dropout_rate),\n            nn.Linear(hidden_dim, dim),\n            nn.Dropout(p=dropout_rate),\n        )\n\n    def forward(self, x):\n#         print(\"In feedforward\", x.shape)\n#         print(\"self.net(x)\", self.net(x).shape)\n        return self.net(x)\n\n\nclass TransformerModel(nn.Module):\n    def __init__(\n        self,\n        dim,\n        depth,\n        heads,\n        mlp_dim,\n        dropout_rate=0.1,\n        attn_dropout_rate=0.1,\n    ):\n        super().__init__()\n        layers = []\n        for _ in range(depth):\n            layers.extend(\n                [\n                   \n                    Residual(\n                        PreNormDrop(\n                            dim,\n                            dropout_rate,\n                            \n                            WindowAttention(\n            dim,\n            window_size=(7,7,7),\n            num_heads=heads,\n            qkv_bias=True,\n            attn_drop=0.1,\n            proj_drop=0.1,\n        )\n#                             VSAWindowAttention(\n#                 in_chans=1, out_dim=14, num_heads=heads, window_size=(7,7,7), qkv_bias=True, qk_scale=None,\n#             attn_drop=attn_dropout_rate, proj_drop=0.1, img_size=(128//4, 128//4,128//4))\n        \n#                             SelfAttention(dim, heads=heads, dropout_rate=attn_dropout_rate),\n                        )\n                    ),\n                    Residual(\n                        PreNorm(dim, FeedForward(dim, mlp_dim, dropout_rate))\n                    ),\n                ]\n            )\n            # dim = dim / 2\n        self.net = IntermediateSequential(*layers)\n\n\n    def forward(self, x):\n        return self.net(x)","metadata":{"execution":{"iopub.status.busy":"2024-02-20T19:16:45.308872Z","iopub.execute_input":"2024-02-20T19:16:45.309283Z","iopub.status.idle":"2024-02-20T19:16:45.338215Z","shell.execute_reply.started":"2024-02-20T19:16:45.309251Z","shell.execute_reply":"2024-02-20T19:16:45.336845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom monai.utils import deprecated_arg, ensure_tuple_rep, optional_import\nfrom monai.utils.module import look_up_option\nimport numpy as np\nrearrange, _ = optional_import(\"einops\", name=\"rearrange\")\nRearrange, _ = optional_import(\"einops.layers.torch\", name=\"Rearrange\")\nimport torch\nimport torch.nn as nn\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        print(b, d, h, w, c, \"b, d, h, w, c\", \"In window partition\")\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        print(x.shape)\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    print(len(windows), windows[0].shape, \"windows[0].shape, windows[999].shape\")\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    print(use_window_size, \"use_window_size\", shift_size, \"shift_size\", x_size, \"x_size\", use_shift_size)\n    \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\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                    print(img_mask[:, d, h, w, :].shape, cnt, \"img_mask.shape, cnt\", d, h, w, \"d, h, w\")\n#                     print(img_mask)\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    print(img_mask.shape, window_size, \"img_mask.shape, window_size\")\n    mask_windows = window_partition(img_mask, window_size)\n    print(mask_windows.shape, len(mask_windows), \"len(mask_windows)::::::::::::::::::::::::::::::::::::::::::::\", mask_windows[0].shape, \"::::::::::::::::::::\")\n    mask_windows = mask_windows.squeeze(-1)\n    print(mask_windows.shape)\n    attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)\n    print(attn_mask.shape)\n    attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))\n#     print(attn_mask[125])\n    return attn_mask\n\n\nclass ConvolutionalVisionTransformer(nn.Module):\n    def __init__(self, in_channels, num_classes, patch_size=16, dim=8, num_layers=6, num_heads=8, dim_feedforward=2048, dropout=0.1):\n        super(ConvolutionalVisionTransformer, self).__init__()\n\n        # Patch embedding layer\n        self.patch_embedding = nn.Conv3d(in_channels, dim, kernel_size=patch_size, stride=patch_size)\n        self.dim=dim\n        # Calculate number of patches\n        self.num_patches = (int((128 - patch_size) / patch_size) + 1) ** 3\n\n        # Positional embedding (corrected shape)\n        self.positional_embedding = nn.Parameter(torch.zeros(1, self.num_patches, dim))  # Removed extra dimension\n\n\n        # Convolutional layers\n        \n        self.transformer2=TransformerModel(512,4,8,4096,0.1,0.1)\n#             num_layers=4,\n#             num_heads=8,\n#             hidden_dim=4096,\n#             dropout_rate=0.1,\n#             attn_dropout_rate=0.1)\n        self.classification_head = nn.Linear(512, num_classes)\n\n    def forward(self, x):\n        x = self.patch_embedding(x)  \n        print(x.shape)\n#         x = x.flatten(2).transpose(1, 2)\n        x = x + self.positional_embedding[:, :x.size(1)]\n        print(x.shape)\n        print(\"shape of x before passing into conv layers\",x.shape)\n#         x = torch.tensor([x[0][0], x[0][2], x[0][1]])\n#         print(x.shape)\n#         x = x.permute(0, 2, 1)\n        print(x.shape)\n#         x = x.view(x.shape[0], x.shape[1], x.size(2), 1, 1) \n        print(x.shape)\n        \n#         x = self.conv_layers(x)\n#         x = x.flatten(2).transpose(1, 2)\n#         x=x.squeeze(-1).squeeze(-1)\n        print(\"before transformer\",x.shape)\n        x_shape = x.size()\n        print(x_shape, \"In basic layer \")\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), (7,7,7), (3,3,3))\n            x = rearrange(x, \"b c d h w -> b d h w c\")\n            print(b, c, h, w, \"b, c, h, w values\", (7,7,7), (3,3,3), \"self.window_size, self.shift_size\")\n            \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            print(window_size[0], window_size[1], h, w, d, hp, wp, dp, \n                  \"window_size[0], window_size[1], h, w, d, hp, wp, dp\")\n            \n            attn_mask = compute_mask([dp, hp, wp], window_size, shift_size, x.device)\n            \n            x = x.view(b, d, h, w, -1)\n            x = rearrange(x, \"b d h w c -> b c d h w\")\n        print(\"Dfasdf\",x.shape)\n        x = x.flatten(2).transpose(1, 2)\n        print(\"Dfasdf\",x.shape)\n        x = x.permute(0, 2, 1)\n        print(\"Dfasdf\",x.shape)\n        x, _ = self.transformer2(x)  # Unpack the tuple\n        x = x.mean(dim=1)  # Now you can take the mean\n        logits = self.classification_head(x)\n        return logits\n\n\nmodel = ConvolutionalVisionTransformer(1, 14)\n\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)\n\n# Create sample input tensors\nbatch_images = torch.randn(32, 1, 128, 128,128).to(device)\n\n# # Forward pass\nclassification_outputs = model(batch_images)\nprint(classification_outputs.shape)  # Output shape: (32, 14)\nprint(\"what are these\",classification_outputs)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T19:19:47.108188Z","iopub.execute_input":"2024-02-20T19:19:47.108649Z","iopub.status.idle":"2024-02-20T19:19:48.314156Z","shell.execute_reply.started":"2024-02-20T19:19:47.108609Z","shell.execute_reply":"2024-02-20T19:19:48.312884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport nibabel as nib\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport numpy as np\nfrom sklearn import metrics\nfrom sklearn.metrics import precision_recall_fscore_support\nfrom sklearn.model_selection import train_test_split\nfrom scipy.ndimage import zoom\nfrom sklearn.metrics import confusion_matrix\nfrom sklearn.metrics import accuracy_score, precision_score\n# Create a function to move data to the device\ndef move_data_to_device(data, device):\n    return data.to(torch.float32).to(device)\n\nclass CustomDataset(Dataset):\n    def __init__(self, image_paths, labels, transform=None):\n        self.image_paths = image_paths\n        self.labels = labels\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        image_path = self.image_paths[idx]\n\n        # Load the 3D NIfTI image using nibabel\n        image = nib.load(image_path).get_fdata()\n\n        # Apply transformations if provided to the image\n        if self.transform:\n            image = self.transform(image)\n\n        label = torch.tensor(self.labels[idx], dtype=torch.float32)\n        \n        return image, label\n\nimport torch.nn.functional as F\n\n# Function to resize NIfTI data\ndef resize_nifti(nifti_data, target_shape):\n    factors = (target_shape[0] / nifti_data.shape[0],\n               target_shape[1] / nifti_data.shape[1],\n               target_shape[2] / nifti_data.shape[2])\n    resized_data = zoom(nifti_data, factors, order=3)  # Cubic interpolation (higher quality)\n    return resized_data\n\n# Paths and settings\ncsv_file = '/kaggle/input/unhealthy-csv-file/combined_data (6).csv'  # Update with the correct path\nbatch_size = 32\nnum_workers = 4  # Number of CPU cores to use for data loading\nnum_classes = 14  # Number of classes\ndesired_shape = (128, 128, 128)\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n# Define transformations if needed\ntransform = transforms.Compose([\n    transforms.ToTensor(),  # Convert to tensor\n    # Add more transformations if necessary\n])\n\n# Load the CSV file\ndata = pd.read_csv(csv_file).head(1280)\n\n# Remove the extra space from the column name\ndata.columns = data.columns.str.strip()\n\n# Assuming 'data' is your DataFrame\ndata_length = len(data)\nprint(\"Length of DataFrame:\", data_length)\n\n# Split the data into training, validation, and test sets\ntrain_data, temp_data = train_test_split(data, test_size=0.2, random_state=42)\nval_data, test_data = train_test_split(temp_data, test_size=0.5, random_state=42)\n\n# Set the display option to show all rows\npd.set_option('display.max_rows', None)\n\nindex_values = train_data.index.values\n\n# Reset the display option to its default value (if needed)\npd.reset_option('display.max_rows')\n\n# Extract file paths and labels from the data\ntrain_paths = train_data['file_path'].values\ntrain_labels = train_data[['bowel_healthy','bowel_injury','extravasation_healthy','extravasation_injury','kidney_healthy','kidney_low','kidney_high','liver_healthy','liver_low','liver_high','spleen_healthy','spleen_low','spleen_high','any_injury']].values\n# print(\"train paths\", train_paths)\n\n\nval_paths = val_data['file_path'].values\nval_labels = val_data[['bowel_healthy','bowel_injury','extravasation_healthy','extravasation_injury','kidney_healthy','kidney_low','kidney_high','liver_healthy','liver_low','liver_high','spleen_healthy','spleen_low','spleen_high','any_injury']].values\n# print(val_paths)\n\n\ntest_paths = test_data['file_path'].values\ntest_labels = test_data[['bowel_healthy','bowel_injury','extravasation_healthy','extravasation_injury','kidney_healthy','kidney_low','kidney_high','liver_healthy','liver_low','liver_high','spleen_healthy','spleen_low','spleen_high','any_injury']].values\n# print(test_paths)\n\n\n# Instantiate the datasets\ntrain_dataset = CustomDataset(train_paths, train_labels, transform=transform)\nprint('len of train_dataset', len(train_dataset))\nval_dataset = CustomDataset(val_paths, val_labels, transform=transform)\ntest_dataset = CustomDataset(test_paths, test_labels, transform=transform)\n\n# Instantiate the data loaders\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers, drop_last=True)\nprint('train_loader', len(train_loader))\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers)\ntest_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers)\nprint(train_loader)\n\n# Instantiate the model with the appropriate number of classes for classification\nin_channels = 1  # Input channels (e.g., for grayscale images or volumes)\nnum_classes_classification = 14  # Number of classes for classification\nmodel_class = ConvolutionalVisionTransformer(in_channels, num_classes_classification)\n\n# Count the number of parameters\ntotal_params_class = sum(p.numel() for p in model_class.parameters())\nprint(f\"Total Trainable Parameters for Classification: {total_params_class}\")\n\n# Define loss function and optimizer\nclass_criterion = nn.CrossEntropyLoss()  # Binary Cross-Entropy loss for classification\nclass_optimizer = optim.Adam(model_class.parameters(), lr=0.001)\n\n# Training loop\n# Training loop\nclass_labels = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen', 'any_injury']\n\n\n# Training loop\nnum_epochs = 100\nfor epoch in range(num_epochs):\n    model_class.train()\n    running_loss = 0.0\n    correct_train = 0\n    total_train = 0\n    batch_number=0\n    all_predicted_labels = []\n    all_batch_labels = []\n    for batch_images, batch_labels in train_loader:\n        batch_number=batch_number+1\n        print(batch_number)\n        # Move data to the GPU if available\n        batch_images = batch_images.to(torch.float32).to(device)\n        batch_labels = batch_labels.to(torch.float32).to(device)\n        print(batch_labels.shape)\n        print(batch_labels)\n\n        # Assuming batch_images has shape (batch_size, num_frames, num_channels, height, width)\n        batch_images = batch_images.unsqueeze(1)  # Add a singleton dimension for channels\n        \n        # Forward pass for classification\n        classification_outputs = model_class(batch_images)\n#         print(classification_outputs.shape, \"before sigmoid\")\n#         print(classification_outputs)\n        # Apply sigmoid activation to the classification outputs\n        \n        # Calculate binary cross-entropy loss for each class separately\n        class_loss = class_criterion(classification_outputs, batch_labels)\n        \n        class_optimizer.zero_grad()\n        class_loss.backward()\n        class_optimizer.step()\n        \n        classification_outputs = torch.sigmoid(classification_outputs)\n#         print(classification_outputs.shape, \"after sigmoid\")\n#         print(classification_outputs)\n        \n        # Calculate accuracy and precision\n        predicted_labels = (classification_outputs > 0.5).float()\n        print(\"predicted labels\",  predicted_labels)\n        all_predicted_labels.append(predicted_labels.cpu().numpy())\n        all_batch_labels.append(batch_labels.cpu().numpy())\n        y_true = batch_labels.flatten()\n        y_pred = predicted_labels.flatten()\n#         true_positives = (predicted_labels * batch_labels).sum(dim=0)\n#         false_positives = ((1 - batch_labels) * predicted_labels).sum(dim=0)\n#         false_negatives = (batch_labels * (1 - predicted_labels)).sum(dim=0)\n#         true_negatives = ((1 - batch_labels) * (1 - predicted_labels)).sum(dim=0)\n#         accuracy = (true_positives + true_negatives) / (true_positives + true_negatives + false_positives + false_negatives)\n#         precision = true_positives / (true_positives + false_positives)\n        \n        accuracy = metrics.accuracy_score(y_true, y_pred)\n        print(\"Batch Classification Loss:\", class_loss.item())\n#         print(\"precision\",precision)\n        print(\"Batch accuracy\",accuracy)\n    \n            # Flatten arrays for binary classification metrics\n    all_batch_labels = np.concatenate(all_batch_labels, axis=0)\n    all_predicted_labels = np.concatenate(all_predicted_labels, axis=0)\n    y_true = all_batch_labels.flatten()\n    y_pred = all_predicted_labels.flatten()\n   \n\n    precision = metrics.precision_score(y_true, y_pred, average='binary')\n    accuracy = metrics.accuracy_score(y_true, y_pred)\n    print(\"Epoch\", epoch)\n    print(f'Precision: {precision:.4f}')\n    print(f'Accuracy: {accuracy:.4f}')\n    \n    if epoch % 50 == 0:\n        torch.save({\n            'epoch': epoch,\n            'model_state_dict': model_class.state_dict(),\n            'optimizer_state_dict': class_optimizer.state_dict(),\n            'loss': class_loss.item()\n            # Add any other information you want to save\n        }, f'/kaggle/working/denseNet_epoch_{epoch}.pth')\n\n#     Validation loop    \nwith torch.no_grad():\n        model_class.eval()\n        all_predicted_labels = []\n        all_batch_labels = []\n        for batch_images, batch_labels in val_loader:\n            batch_images = batch_images.to(torch.float32).to(device)\n            batch_labels = batch_labels.to(torch.float32).to(device)\n\n            batch_images = batch_images.unsqueeze(1)\n\n            classification_outputs = model_class(batch_images)\n            classification_outputs = torch.sigmoid(classification_outputs)\n\n            predicted_labels = (classification_outputs > 0.5).float()\n\n            all_predicted_labels.append(predicted_labels.cpu().numpy())\n            all_batch_labels.append(batch_labels.cpu().numpy())\n        \n        # Flatten arrays for binary classification metrics\n        all_batch_labels = np.concatenate(all_batch_labels, axis=0)\n        all_predicted_labels = np.concatenate(all_predicted_labels, axis=0)\n\n        y_true = all_batch_labels.flatten()\n        y_pred = all_predicted_labels.flatten()\n\n        # Precision\n        precision = metrics.precision_score(y_true, y_pred, average='binary')\n\n\n        # Accuracy\n        accuracy = metrics.accuracy_score(y_true, y_pred)\n\n        \n        print(\"Validation Results: \")\n        print(f'Precision: {precision:.4f}')\n        print(f'Accuracy: {accuracy:.4f}')\n    # Testing loop\nwith torch.no_grad():\n        model_class.eval()\n        all_predicted_labels = []\n        all_batch_labels = []\n        for batch_images, batch_labels in test_loader:\n            batch_images = batch_images.to(torch.float32).to(device)\n            batch_labels = batch_labels.to(torch.float32).to(device)\n\n            batch_images = batch_images.unsqueeze(1)\n\n            classification_outputs = model_class(batch_images)\n            classification_outputs = torch.sigmoid(classification_outputs)\n\n            predicted_labels = (classification_outputs > 0.5).float()\n\n            all_predicted_labels.append(predicted_labels.cpu().numpy())\n            all_batch_labels.append(batch_labels.cpu().numpy())\n        \n        # Flatten arrays for binary classification metrics\n        all_batch_labels = np.concatenate(all_batch_labels, axis=0)\n        all_predicted_labels = np.concatenate(all_predicted_labels, axis=0)\n\n        y_true = all_batch_labels.flatten()\n        y_pred = all_predicted_labels.flatten()\n        print(\"y_true shape \",y_true.shape)\n        print(\"y_pred_shape\" ,y_pred.shape)\n        # Precision, Recall, F1 Score\n        precision = metrics.precision_score(y_true, y_pred, average='binary')\n        recall = metrics.recall_score(y_true, y_pred, average='binary')\n        f1_score = metrics.f1_score(y_true, y_pred, average='binary')\n\n        # Accuracy\n        accuracy = metrics.accuracy_score(y_true, y_pred)\n\n        # AUC\n        fpr, tpr, thresholds = metrics.roc_curve(y_true, y_pred)\n        auc = metrics.auc(fpr, tpr)\n        \n        print(\"Testing Results: \")\n        print(f'Precision: {precision:.4f}')\n        print(f'Recall: {recall:.4f}')\n        print(f'F1 Score: {f1_score:.4f}')\n        print(f'Accuracy: {accuracy:.4f}')\n        print(f'AUC: {auc:.4f}')\n        \n        # Reshape predictions for multilabel classification metrics\n        y_true_multilabel = all_batch_labels.T\n        y_pred_multilabel = all_predicted_labels.T\n        print(y_true_multilabel.shape)\n        print(y_pred_multilabel.shape)\n        y_true_1 = y_true_multilabel[:, -1]\n        y_pred_1 = y_pred_multilabel[:, -1]\n# Compute precision, recall, F1 score, and support for each class\n        precision, recall, f1_score, support = precision_recall_fscore_support(y_true_multilabel, y_pred_multilabel, average=None)\n\n# Print scores for each class\n        for i in range(num_classes):\n            class_accuracy = accuracy_score(y_true_multilabel[i, :], y_pred_multilabel[i, :])\n            class_precision = precision_score(y_true_multilabel[i, :], y_pred_multilabel[i, :])\n\n            print(f\"   Class {i + 1}:\")\n            print(f\"      Accuracy: {class_accuracy:.4f}\")\n            print(f\"      Precision: {class_precision:.4f}\")\n            print()","metadata":{"execution":{"iopub.status.busy":"2024-01-01T14:47:52.164274Z","iopub.execute_input":"2024-01-01T14:47:52.164623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # --------------------------------------------------------\n# # Swin Transformer V2\n# # Copyright (c) 2022 Microsoft\n# # Licensed under The MIT License [see LICENSE for details]\n# # Written by Ze Liu\n# # --------------------------------------------------------\n\n# import torch\n# import torch.nn as nn\n# import torch.nn.functional as F\n# import torch.utils.checkpoint as checkpoint\n# from timm.models.layers import DropPath, to_2tuple, trunc_normal_\n# import numpy as np\n\n\n# class Mlp(nn.Module):\n#     def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):\n#         super().__init__()\n#         out_features = out_features or in_features\n#         hidden_features = hidden_features or in_features\n#         self.fc1 = nn.Linear(in_features, hidden_features)\n#         self.act = act_layer()\n#         self.fc2 = nn.Linear(hidden_features, out_features)\n#         self.drop = nn.Dropout(drop)\n\n#     def forward(self, x):\n#         x = self.fc1(x)\n#         x = self.act(x)\n#         x = self.drop(x)\n#         x = self.fc2(x)\n#         x = self.drop(x)\n#         return x\n\n\n# def window_partition(x, window_size):\n#     \"\"\"\n#     Args:\n#         x: (B, H, W, C)\n#         window_size (int): window size\n\n#     Returns:\n#         windows: (num_windows*B, window_size, window_size, C)\n#     \"\"\"\n#     B, H, W, C = x.shape\n#     x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)\n#     windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)\n#     return windows\n\n\n# def window_reverse(windows, window_size, H, W):\n#     \"\"\"\n#     Args:\n#         windows: (num_windows*B, window_size, window_size, C)\n#         window_size (int): Window size\n#         H (int): Height of image\n#         W (int): Width of image\n\n#     Returns:\n#         x: (B, H, W, C)\n#     \"\"\"\n#     B = int(windows.shape[0] / (H * W / window_size / window_size))\n#     x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1)\n#     x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)\n#     return x\n\n\n# class WindowAttention(nn.Module):\n#     r\"\"\" Window based multi-head self attention (W-MSA) module with relative position bias.\n#     It supports both of shifted and non-shifted window.\n\n#     Args:\n#         dim (int): Number of input channels.\n#         window_size (tuple[int]): The height and width of the window.\n#         num_heads (int): Number of attention heads.\n#         qkv_bias (bool, optional):  If True, add a learnable bias to query, key, value. Default: True\n#         attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0\n#         proj_drop (float, optional): Dropout ratio of output. Default: 0.0\n#         pretrained_window_size (tuple[int]): The height and width of the window in pre-training.\n#     \"\"\"\n\n#     def __init__(self, dim, window_size, num_heads, qkv_bias=True, attn_drop=0., proj_drop=0.,\n#                  pretrained_window_size=[0, 0]):\n\n#         super().__init__()\n#         self.dim = dim\n#         self.window_size = window_size  # Wh, Ww\n#         self.pretrained_window_size = pretrained_window_size\n#         self.num_heads = num_heads\n\n#         self.logit_scale = nn.Parameter(torch.log(10 * torch.ones((num_heads, 1, 1))), requires_grad=True)\n\n#         # mlp to generate continuous relative position bias\n#         self.cpb_mlp = nn.Sequential(nn.Linear(2, 512, bias=True),\n#                                      nn.ReLU(inplace=True),\n#                                      nn.Linear(512, num_heads, bias=False))\n\n#         # get relative_coords_table\n#         relative_coords_h = torch.arange(-(self.window_size[0] - 1), self.window_size[0], dtype=torch.float32)\n#         relative_coords_w = torch.arange(-(self.window_size[1] - 1), self.window_size[1], dtype=torch.float32)\n#         relative_coords_table = torch.stack(\n#             torch.meshgrid([relative_coords_h,\n#                             relative_coords_w])).permute(1, 2, 0).contiguous().unsqueeze(0)  # 1, 2*Wh-1, 2*Ww-1, 2\n#         if pretrained_window_size[0] > 0:\n#             relative_coords_table[:, :, :, 0] /= (pretrained_window_size[0] - 1)\n#             relative_coords_table[:, :, :, 1] /= (pretrained_window_size[1] - 1)\n#         else:\n#             relative_coords_table[:, :, :, 0] /= (self.window_size[0] - 1)\n#             relative_coords_table[:, :, :, 1] /= (self.window_size[1] - 1)\n#         relative_coords_table *= 8  # normalize to -8, 8\n#         relative_coords_table = torch.sign(relative_coords_table) * torch.log2(\n#             torch.abs(relative_coords_table) + 1.0) / np.log2(8)\n\n#         self.register_buffer(\"relative_coords_table\", relative_coords_table)\n\n#         # get pair-wise relative position index for each token inside the window\n#         coords_h = torch.arange(self.window_size[0])\n#         coords_w = torch.arange(self.window_size[1])\n#         coords = torch.stack(torch.meshgrid([coords_h, coords_w]))  # 2, Wh, Ww\n#         coords_flatten = torch.flatten(coords, 1)  # 2, Wh*Ww\n#         relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :]  # 2, Wh*Ww, Wh*Ww\n#         relative_coords = relative_coords.permute(1, 2, 0).contiguous()  # Wh*Ww, Wh*Ww, 2\n#         relative_coords[:, :, 0] += self.window_size[0] - 1  # shift to start from 0\n#         relative_coords[:, :, 1] += self.window_size[1] - 1\n#         relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1\n#         relative_position_index = relative_coords.sum(-1)  # Wh*Ww, Wh*Ww\n#         self.register_buffer(\"relative_position_index\", relative_position_index)\n\n#         self.qkv = nn.Linear(dim, dim * 3, bias=False)\n#         if qkv_bias:\n#             self.q_bias = nn.Parameter(torch.zeros(dim))\n#             self.v_bias = nn.Parameter(torch.zeros(dim))\n#         else:\n#             self.q_bias = None\n#             self.v_bias = None\n#         self.attn_drop = nn.Dropout(attn_drop)\n#         self.proj = nn.Linear(dim, dim)\n#         self.proj_drop = nn.Dropout(proj_drop)\n#         self.softmax = nn.Softmax(dim=-1)\n\n#     def forward(self, x, mask=None):\n#         \"\"\"\n#         Args:\n#             x: input features with shape of (num_windows*B, N, C)\n#             mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None\n#         \"\"\"\n#         B_, N, C = x.shape\n#         qkv_bias = None\n#         if self.q_bias is not None:\n#             qkv_bias = torch.cat((self.q_bias, torch.zeros_like(self.v_bias, requires_grad=False), self.v_bias))\n#         qkv = F.linear(input=x, weight=self.qkv.weight, bias=qkv_bias)\n#         qkv = qkv.reshape(B_, N, 3, self.num_heads, -1).permute(2, 0, 3, 1, 4)\n#         q, k, v = qkv[0], qkv[1], qkv[2]  # make torchscript happy (cannot use tensor as tuple)\n\n#         # cosine attention\n#         attn = (F.normalize(q, dim=-1) @ F.normalize(k, dim=-1).transpose(-2, -1))\n#         logit_scale = torch.clamp(self.logit_scale, max=torch.log(torch.tensor(1. / 0.01))).exp()\n#         attn = attn * logit_scale\n\n#         relative_position_bias_table = self.cpb_mlp(self.relative_coords_table).view(-1, self.num_heads)\n#         relative_position_bias = relative_position_bias_table[self.relative_position_index.view(-1)].view(\n#             self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1)  # Wh*Ww,Wh*Ww,nH\n#         relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous()  # nH, Wh*Ww, Wh*Ww\n#         relative_position_bias = 16 * torch.sigmoid(relative_position_bias)\n#         attn = attn + relative_position_bias.unsqueeze(0)\n\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)\n\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#     def extra_repr(self) -> str:\n#         return f'dim={self.dim}, window_size={self.window_size}, ' \\\n#                f'pretrained_window_size={self.pretrained_window_size}, num_heads={self.num_heads}'\n\n#     def flops(self, N):\n#         # calculate flops for 1 window with token length of N\n#         flops = 0\n#         # qkv = self.qkv(x)\n#         flops += N * self.dim * 3 * self.dim\n#         # attn = (q @ k.transpose(-2, -1))\n#         flops += self.num_heads * N * (self.dim // self.num_heads) * N\n#         #  x = (attn @ v)\n#         flops += self.num_heads * N * N * (self.dim // self.num_heads)\n#         # x = self.proj(x)\n#         flops += N * self.dim * self.dim\n#         return flops\n\n\n# class SwinTransformerBlock(nn.Module):\n#     r\"\"\" Swin Transformer Block.\n\n#     Args:\n#         dim (int): Number of input channels.\n#         input_resolution (tuple[int]): Input resulotion.\n#         num_heads (int): Number of attention heads.\n#         window_size (int): Window size.\n#         shift_size (int): Shift size for SW-MSA.\n#         mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.\n#         qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True\n#         drop (float, optional): Dropout rate. Default: 0.0\n#         attn_drop (float, optional): Attention dropout rate. Default: 0.0\n#         drop_path (float, optional): Stochastic depth rate. Default: 0.0\n#         act_layer (nn.Module, optional): Activation layer. Default: nn.GELU\n#         norm_layer (nn.Module, optional): Normalization layer.  Default: nn.LayerNorm\n#         pretrained_window_size (int): Window size in pre-training.\n#     \"\"\"\n\n#     def __init__(self, dim, input_resolution, num_heads, window_size=7, shift_size=0,\n#                  mlp_ratio=4., qkv_bias=True, drop=0., attn_drop=0., drop_path=0.,\n#                  act_layer=nn.GELU, norm_layer=nn.LayerNorm, pretrained_window_size=0):\n#         super().__init__()\n#         self.dim = dim\n#         self.input_resolution = input_resolution\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#         if min(self.input_resolution) <= self.window_size:\n#             # if window size is larger than input resolution, we don't partition windows\n#             self.shift_size = 0\n#             self.window_size = min(self.input_resolution)\n#         assert 0 <= self.shift_size < self.window_size, \"shift_size must in 0-window_size\"\n\n#         self.norm1 = norm_layer(dim)\n#         self.attn = WindowAttention(\n#             dim, window_size=to_2tuple(self.window_size), num_heads=num_heads,\n#             qkv_bias=qkv_bias, attn_drop=attn_drop, proj_drop=drop,\n#             pretrained_window_size=to_2tuple(pretrained_window_size))\n\n#         self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n#         self.norm2 = norm_layer(dim)\n#         mlp_hidden_dim = int(dim * mlp_ratio)\n#         self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)\n\n#         if self.shift_size > 0:\n#             # calculate attention mask for SW-MSA\n#             H, W = self.input_resolution\n#             img_mask = torch.zeros((1, H, W, 1))  # 1 H W 1\n#             h_slices = (slice(0, -self.window_size),\n#                         slice(-self.window_size, -self.shift_size),\n#                         slice(-self.shift_size, None))\n#             w_slices = (slice(0, -self.window_size),\n#                         slice(-self.window_size, -self.shift_size),\n#                         slice(-self.shift_size, None))\n#             cnt = 0\n#             for h in h_slices:\n#                 for w in w_slices:\n#                     img_mask[:, h, w, :] = cnt\n#                     cnt += 1\n\n#             mask_windows = window_partition(img_mask, self.window_size)  # nW, window_size, window_size, 1\n#             mask_windows = mask_windows.view(-1, self.window_size * self.window_size)\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#         else:\n#             attn_mask = None\n\n#         self.register_buffer(\"attn_mask\", attn_mask)\n\n#     def forward(self, x):\n#         H, W = self.input_resolution\n#         B, L, C = x.shape\n#         assert L == H * W, \"input feature has wrong size\"\n\n#         shortcut = x\n#         x = x.view(B, H, W, C)\n\n#         # cyclic shift\n#         if self.shift_size > 0:\n#             shifted_x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2))\n#         else:\n#             shifted_x = x\n\n#         # partition windows\n#         x_windows = window_partition(shifted_x, self.window_size)  # nW*B, window_size, window_size, C\n#         x_windows = x_windows.view(-1, self.window_size * self.window_size, C)  # nW*B, window_size*window_size, C\n\n#         # W-MSA/SW-MSA\n#         attn_windows = self.attn(x_windows, mask=self.attn_mask)  # nW*B, window_size*window_size, C\n\n#         # merge windows\n#         attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C)\n#         shifted_x = window_reverse(attn_windows, self.window_size, H, W)  # B H' W' C\n\n#         # reverse cyclic shift\n#         if self.shift_size > 0:\n#             x = torch.roll(shifted_x, shifts=(self.shift_size, self.shift_size), dims=(1, 2))\n#         else:\n#             x = shifted_x\n#         x = x.view(B, H * W, C)\n#         x = shortcut + self.drop_path(self.norm1(x))\n\n#         # FFN\n#         x = x + self.drop_path(self.norm2(self.mlp(x)))\n\n#         return x\n\n#     def extra_repr(self) -> str:\n#         return f\"dim={self.dim}, input_resolution={self.input_resolution}, num_heads={self.num_heads}, \" \\\n#                f\"window_size={self.window_size}, shift_size={self.shift_size}, mlp_ratio={self.mlp_ratio}\"\n\n#     def flops(self):\n#         flops = 0\n#         H, W = self.input_resolution\n#         # norm1\n#         flops += self.dim * H * W\n#         # W-MSA/SW-MSA\n#         nW = H * W / self.window_size / self.window_size\n#         flops += nW * self.attn.flops(self.window_size * self.window_size)\n#         # mlp\n#         flops += 2 * H * W * self.dim * self.dim * self.mlp_ratio\n#         # norm2\n#         flops += self.dim * H * W\n#         return flops\n\n\n# class PatchMerging(nn.Module):\n#     r\"\"\" Patch Merging Layer.\n\n#     Args:\n#         input_resolution (tuple[int]): Resolution of input feature.\n#         dim (int): Number of input channels.\n#         norm_layer (nn.Module, optional): Normalization layer.  Default: nn.LayerNorm\n#     \"\"\"\n\n#     def __init__(self, input_resolution, dim, norm_layer=nn.LayerNorm):\n#         super().__init__()\n#         self.input_resolution = input_resolution\n#         self.dim = dim\n#         self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)\n#         self.norm = norm_layer(2 * dim)\n\n#     def forward(self, x):\n#         \"\"\"\n#         x: B, H*W, C\n#         \"\"\"\n#         H, W = self.input_resolution\n#         B, L, C = x.shape\n#         assert L == H * W, \"input feature has wrong size\"\n#         assert H % 2 == 0 and W % 2 == 0, f\"x size ({H}*{W}) are not even.\"\n\n#         x = x.view(B, H, W, C)\n\n#         x0 = x[:, 0::2, 0::2, :]  # B H/2 W/2 C\n#         x1 = x[:, 1::2, 0::2, :]  # B H/2 W/2 C\n#         x2 = x[:, 0::2, 1::2, :]  # B H/2 W/2 C\n#         x3 = x[:, 1::2, 1::2, :]  # B H/2 W/2 C\n#         x = torch.cat([x0, x1, x2, x3], -1)  # B H/2 W/2 4*C\n#         x = x.view(B, -1, 4 * C)  # B H/2*W/2 4*C\n\n#         x = self.reduction(x)\n#         x = self.norm(x)\n\n#         return x\n\n#     def extra_repr(self) -> str:\n#         return f\"input_resolution={self.input_resolution}, dim={self.dim}\"\n\n#     def flops(self):\n#         H, W = self.input_resolution\n#         flops = (H // 2) * (W // 2) * 4 * self.dim * 2 * self.dim\n#         flops += H * W * self.dim // 2\n#         return flops\n\n\n# class BasicLayer(nn.Module):\n#     \"\"\" A basic Swin Transformer layer for one stage.\n\n#     Args:\n#         dim (int): Number of input channels.\n#         input_resolution (tuple[int]): Input resolution.\n#         depth (int): Number of blocks.\n#         num_heads (int): Number of attention heads.\n#         window_size (int): Local window size.\n#         mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.\n#         qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True\n#         drop (float, optional): Dropout rate. Default: 0.0\n#         attn_drop (float, optional): Attention dropout rate. Default: 0.0\n#         drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0\n#         norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm\n#         downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None\n#         use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.\n#         pretrained_window_size (int): Local window size in pre-training.\n#     \"\"\"\n\n#     def __init__(self, dim, input_resolution, depth, num_heads, window_size,\n#                  mlp_ratio=4., qkv_bias=True, drop=0., attn_drop=0.,\n#                  drop_path=0., norm_layer=nn.LayerNorm, downsample=None, use_checkpoint=False,\n#                  pretrained_window_size=0):\n\n#         super().__init__()\n#         self.dim = dim\n#         self.input_resolution = input_resolution\n#         self.depth = depth\n#         self.use_checkpoint = use_checkpoint\n\n#         # build blocks\n#         self.blocks = nn.ModuleList([\n#             SwinTransformerBlock(dim=dim, input_resolution=input_resolution,\n#                                  num_heads=num_heads, window_size=window_size,\n#                                  shift_size=0 if (i % 2 == 0) else window_size // 2,\n#                                  mlp_ratio=mlp_ratio,\n#                                  qkv_bias=qkv_bias,\n#                                  drop=drop, attn_drop=attn_drop,\n#                                  drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path,\n#                                  norm_layer=norm_layer,\n#                                  pretrained_window_size=pretrained_window_size)\n#             for i in range(depth)])\n\n#         # patch merging layer\n#         if downsample is not None:\n#             self.downsample = downsample(input_resolution, dim=dim, norm_layer=norm_layer)\n#         else:\n#             self.downsample = None\n\n#     def forward(self, x):\n#         for blk in self.blocks:\n#             if self.use_checkpoint:\n#                 x = checkpoint.checkpoint(blk, x)\n#             else:\n#                 x = blk(x)\n#         if self.downsample is not None:\n#             x = self.downsample(x)\n#         return x\n\n#     def extra_repr(self) -> str:\n#         return f\"dim={self.dim}, input_resolution={self.input_resolution}, depth={self.depth}\"\n\n#     def flops(self):\n#         flops = 0\n#         for blk in self.blocks:\n#             flops += blk.flops()\n#         if self.downsample is not None:\n#             flops += self.downsample.flops()\n#         return flops\n\n#     def _init_respostnorm(self):\n#         for blk in self.blocks:\n#             nn.init.constant_(blk.norm1.bias, 0)\n#             nn.init.constant_(blk.norm1.weight, 0)\n#             nn.init.constant_(blk.norm2.bias, 0)\n#             nn.init.constant_(blk.norm2.weight, 0)\n\n\n# class PatchEmbed(nn.Module):\n#     r\"\"\" Image to Patch Embedding\n\n#     Args:\n#         img_size (int): Image size.  Default: 224.\n#         patch_size (int): Patch token size. Default: 4.\n#         in_chans (int): Number of input image channels. Default: 3.\n#         embed_dim (int): Number of linear projection output channels. Default: 96.\n#         norm_layer (nn.Module, optional): Normalization layer. Default: None\n#     \"\"\"\n\n#     def __init__(self, img_size=224, patch_size=4, in_chans=3, embed_dim=96, norm_layer=None):\n#         super().__init__()\n#         img_size = to_2tuple(img_size)\n#         patch_size = to_2tuple(patch_size)\n#         patches_resolution = [img_size[0] // patch_size[0], img_size[1] // patch_size[1]]\n#         self.img_size = img_size\n#         self.patch_size = patch_size\n#         self.patches_resolution = patches_resolution\n#         self.num_patches = patches_resolution[0] * patches_resolution[1]\n\n#         self.in_chans = in_chans\n#         self.embed_dim = embed_dim\n\n#         self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)\n#         if norm_layer is not None:\n#             self.norm = norm_layer(embed_dim)\n#         else:\n#             self.norm = None\n\n#     def forward(self, x):\n#         B, C, H, W = x.shape\n#         # FIXME look at relaxing size constraints\n#         assert H == self.img_size[0] and W == self.img_size[1], \\\n#             f\"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]}).\"\n#         x = self.proj(x).flatten(2).transpose(1, 2)  # B Ph*Pw C\n#         if self.norm is not None:\n#             x = self.norm(x)\n#         return x\n\n#     def flops(self):\n#         Ho, Wo = self.patches_resolution\n#         flops = Ho * Wo * self.embed_dim * self.in_chans * (self.patch_size[0] * self.patch_size[1])\n#         if self.norm is not None:\n#             flops += Ho * Wo * self.embed_dim\n#         return flops\n\n\n# class SwinTransformerV2(nn.Module):\n#     r\"\"\" Swin Transformer\n#         A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows`  -\n#           https://arxiv.org/pdf/2103.14030\n\n#     Args:\n#         img_size (int | tuple(int)): Input image size. Default 224\n#         patch_size (int | tuple(int)): Patch size. Default: 4\n#         in_chans (int): Number of input image channels. Default: 3\n#         num_classes (int): Number of classes for classification head. Default: 1000\n#         embed_dim (int): Patch embedding dimension. Default: 96\n#         depths (tuple(int)): Depth of each Swin Transformer layer.\n#         num_heads (tuple(int)): Number of attention heads in different layers.\n#         window_size (int): Window size. Default: 7\n#         mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4\n#         qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True\n#         drop_rate (float): Dropout rate. Default: 0\n#         attn_drop_rate (float): Attention dropout rate. Default: 0\n#         drop_path_rate (float): Stochastic depth rate. Default: 0.1\n#         norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.\n#         ape (bool): If True, add absolute position embedding to the patch embedding. Default: False\n#         patch_norm (bool): If True, add normalization after patch embedding. Default: True\n#         use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False\n#         pretrained_window_sizes (tuple(int)): Pretrained window sizes of each layer.\n#     \"\"\"\n\n#     def __init__(self, img_size=224, patch_size=4, in_chans=3, num_classes=1000,\n#                  embed_dim=96, depths=[2, 2, 6, 2], num_heads=[3, 6, 12, 24],\n#                  window_size=7, mlp_ratio=4., qkv_bias=True,\n#                  drop_rate=0., attn_drop_rate=0., drop_path_rate=0.1,\n#                  norm_layer=nn.LayerNorm, ape=False, patch_norm=True,\n#                  use_checkpoint=False, pretrained_window_sizes=[0, 0, 0, 0], **kwargs):\n#         super().__init__()\n\n#         self.num_classes = num_classes\n#         self.num_layers = len(depths)\n#         self.embed_dim = embed_dim\n#         self.ape = ape\n#         self.patch_norm = patch_norm\n#         self.num_features = int(embed_dim * 2 ** (self.num_layers - 1))\n#         self.mlp_ratio = mlp_ratio\n\n#         # split image into non-overlapping patches\n#         self.patch_embed = PatchEmbed(\n#             img_size=img_size, patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dim,\n#             norm_layer=norm_layer if self.patch_norm else None)\n#         num_patches = self.patch_embed.num_patches\n#         patches_resolution = self.patch_embed.patches_resolution\n#         self.patches_resolution = patches_resolution\n\n#         # absolute position embedding\n#         if self.ape:\n#             self.absolute_pos_embed = nn.Parameter(torch.zeros(1, num_patches, embed_dim))\n#             trunc_normal_(self.absolute_pos_embed, std=.02)\n\n#         self.pos_drop = nn.Dropout(p=drop_rate)\n\n#         # stochastic depth\n#         dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))]  # stochastic depth decay rule\n\n#         # build layers\n#         self.layers = nn.ModuleList()\n#         for i_layer in range(self.num_layers):\n#             layer = BasicLayer(dim=int(embed_dim * 2 ** i_layer),\n#                                input_resolution=(patches_resolution[0] // (2 ** i_layer),\n#                                                  patches_resolution[1] // (2 ** i_layer)),\n#                                depth=depths[i_layer],\n#                                num_heads=num_heads[i_layer],\n#                                window_size=window_size,\n#                                mlp_ratio=self.mlp_ratio,\n#                                qkv_bias=qkv_bias,\n#                                drop=drop_rate, attn_drop=attn_drop_rate,\n#                                drop_path=dpr[sum(depths[:i_layer]):sum(depths[:i_layer + 1])],\n#                                norm_layer=norm_layer,\n#                                downsample=PatchMerging if (i_layer < self.num_layers - 1) else None,\n#                                use_checkpoint=use_checkpoint,\n#                                pretrained_window_size=pretrained_window_sizes[i_layer])\n#             self.layers.append(layer)\n\n#         self.norm = norm_layer(self.num_features)\n#         self.avgpool = nn.AdaptiveAvgPool1d(1)\n#         self.head = nn.Linear(self.num_features, num_classes) if num_classes > 0 else nn.Identity()\n\n#         self.apply(self._init_weights)\n#         for bly in self.layers:\n#             bly._init_respostnorm()\n\n#     def _init_weights(self, m):\n#         if isinstance(m, nn.Linear):\n#             trunc_normal_(m.weight, std=.02)\n#             if isinstance(m, nn.Linear) and m.bias is not None:\n#                 nn.init.constant_(m.bias, 0)\n#         elif isinstance(m, nn.LayerNorm):\n#             nn.init.constant_(m.bias, 0)\n#             nn.init.constant_(m.weight, 1.0)\n\n#     @torch.jit.ignore\n#     def no_weight_decay(self):\n#         return {'absolute_pos_embed'}\n\n#     @torch.jit.ignore\n#     def no_weight_decay_keywords(self):\n#         return {\"cpb_mlp\", \"logit_scale\", 'relative_position_bias_table'}\n\n#     def forward_features(self, x):\n#         x = self.patch_embed(x)\n#         if self.ape:\n#             x = x + self.absolute_pos_embed\n#         x = self.pos_drop(x)\n\n#         for layer in self.layers:\n#             x = layer(x)\n\n#         x = self.norm(x)  # B L C\n#         x = self.avgpool(x.transpose(1, 2))  # B C 1\n#         x = torch.flatten(x, 1)\n#         return x\n\n#     def forward(self, x):\n#         x = self.forward_features(x)\n#         x = self.head(x)\n#         return x\n\n#     def flops(self):\n#         flops = 0\n#         flops += self.patch_embed.flops()\n#         for i, layer in enumerate(self.layers):\n#             flops += layer.flops()\n#         flops += self.num_features * self.patches_resolution[0] * self.patches_resolution[1] // (2 ** self.num_layers)\n#         flops += self.num_features * self.num_classes\n#         return flops","metadata":{"execution":{"iopub.status.busy":"2024-01-02T07:15:30.498401Z","iopub.execute_input":"2024-01-02T07:15:30.49907Z","iopub.status.idle":"2024-01-02T07:15:38.315792Z","shell.execute_reply.started":"2024-01-02T07:15:30.499036Z","shell.execute_reply":"2024-01-02T07:15:38.314401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # --------------------------------------------------------\n# # Swin Transformer\n# # Copyright (c) 2021 Microsoft\n# # Licensed under The MIT License [see LICENSE for details]\n# # Written by Ze Liu\n# # --------------------------------------------------------\n\n# # from .swin_transformer import SwinTransformer\n# # from .swin_transformer_v2 import SwinTransformerV2\n# # from .swin_transformer_moe import SwinTransformerMoE\n# # from .swin_mlp import SwinMLP\n# # from .simmim import build_simmim\n\n\n# def build_model(is_pretrain=False):\n#     model_type = \"swinv2\"\n\n#     # accelerate layernorm\n#     if False:\n#         try:\n#             import apex as amp\n#             layernorm = amp.normalization.FusedLayerNorm\n#         except:\n#             layernorm = None\n#             print(\"To use FusedLayerNorm, please install apex.\")\n#     else:\n#         import torch.nn as nn\n#         layernorm = nn.LayerNorm\n\n#     if is_pretrain:\n#         model = build_simmim(config)\n#         return model\n\n#     if model_type == 'swin':\n#         model = SwinTransformer(img_size=config.DATA.IMG_SIZE,\n#                                 patch_size=config.MODEL.SWIN.PATCH_SIZE,\n#                                 in_chans=config.MODEL.SWIN.IN_CHANS,\n#                                 num_classes=config.MODEL.NUM_CLASSES,\n#                                 embed_dim=config.MODEL.SWIN.EMBED_DIM,\n#                                 depths=config.MODEL.SWIN.DEPTHS,\n#                                 num_heads=config.MODEL.SWIN.NUM_HEADS,\n#                                 window_size=config.MODEL.SWIN.WINDOW_SIZE,\n#                                 mlp_ratio=config.MODEL.SWIN.MLP_RATIO,\n#                                 qkv_bias=config.MODEL.SWIN.QKV_BIAS,\n#                                 qk_scale=config.MODEL.SWIN.QK_SCALE,\n#                                 drop_rate=config.MODEL.DROP_RATE,\n#                                 drop_path_rate=config.MODEL.DROP_PATH_RATE,\n#                                 ape=config.MODEL.SWIN.APE,\n#                                 norm_layer=layernorm,\n#                                 patch_norm=config.MODEL.SWIN.PATCH_NORM,\n#                                 use_checkpoint=config.TRAIN.USE_CHECKPOINT,\n#                                 fused_window_process=config.FUSED_WINDOW_PROCESS)\n#     elif model_type == 'swinv2':\n#         model = SwinTransformerV2(img_size=224,\n#                                   patch_size=4,\n#                                   in_chans=3,\n#                                   num_classes=6,\n#                                   embed_dim=96,\n#                                   depths=[2, 2, 6, 2],\n#                                   num_heads=[3, 6, 12, 24],\n#                                   window_size=7,\n#                                   mlp_ratio=4,\n#                                   qkv_bias=True,\n#                                   drop_rate=0,\n#                                   drop_path_rate=0.1,\n#                                   ape=False,\n#                                   patch_norm=True,\n#                                   use_checkpoint=False,\n#                                   pretrained_window_sizes=[0, 0, 0, 0])\n#     elif model_type == 'swin_moe':\n#         model = SwinTransformerMoE(img_size=config.DATA.IMG_SIZE,\n#                                    patch_size=config.MODEL.SWIN_MOE.PATCH_SIZE,\n#                                    in_chans=config.MODEL.SWIN_MOE.IN_CHANS,\n#                                    num_classes=config.MODEL.NUM_CLASSES,\n#                                    embed_dim=config.MODEL.SWIN_MOE.EMBED_DIM,\n#                                    depths=config.MODEL.SWIN_MOE.DEPTHS,\n#                                    num_heads=config.MODEL.SWIN_MOE.NUM_HEADS,\n#                                    window_size=config.MODEL.SWIN_MOE.WINDOW_SIZE,\n#                                    mlp_ratio=config.MODEL.SWIN_MOE.MLP_RATIO,\n#                                    qkv_bias=config.MODEL.SWIN_MOE.QKV_BIAS,\n#                                    qk_scale=config.MODEL.SWIN_MOE.QK_SCALE,\n#                                    drop_rate=config.MODEL.DROP_RATE,\n#                                    drop_path_rate=config.MODEL.DROP_PATH_RATE,\n#                                    ape=config.MODEL.SWIN_MOE.APE,\n#                                    patch_norm=config.MODEL.SWIN_MOE.PATCH_NORM,\n#                                    mlp_fc2_bias=config.MODEL.SWIN_MOE.MLP_FC2_BIAS,\n#                                    init_std=config.MODEL.SWIN_MOE.INIT_STD,\n#                                    use_checkpoint=config.TRAIN.USE_CHECKPOINT,\n#                                    pretrained_window_sizes=config.MODEL.SWIN_MOE.PRETRAINED_WINDOW_SIZES,\n#                                    moe_blocks=config.MODEL.SWIN_MOE.MOE_BLOCKS,\n#                                    num_local_experts=config.MODEL.SWIN_MOE.NUM_LOCAL_EXPERTS,\n#                                    top_value=config.MODEL.SWIN_MOE.TOP_VALUE,\n#                                    capacity_factor=config.MODEL.SWIN_MOE.CAPACITY_FACTOR,\n#                                    cosine_router=config.MODEL.SWIN_MOE.COSINE_ROUTER,\n#                                    normalize_gate=config.MODEL.SWIN_MOE.NORMALIZE_GATE,\n#                                    use_bpr=config.MODEL.SWIN_MOE.USE_BPR,\n#                                    is_gshard_loss=config.MODEL.SWIN_MOE.IS_GSHARD_LOSS,\n#                                    gate_noise=config.MODEL.SWIN_MOE.GATE_NOISE,\n#                                    cosine_router_dim=config.MODEL.SWIN_MOE.COSINE_ROUTER_DIM,\n#                                    cosine_router_init_t=config.MODEL.SWIN_MOE.COSINE_ROUTER_INIT_T,\n#                                    moe_drop=config.MODEL.SWIN_MOE.MOE_DROP,\n#                                    aux_loss_weight=config.MODEL.SWIN_MOE.AUX_LOSS_WEIGHT)\n#     elif model_type == 'swin_mlp':\n#         model = SwinMLP(img_size=config.DATA.IMG_SIZE,\n#                         patch_size=config.MODEL.SWIN_MLP.PATCH_SIZE,\n#                         in_chans=config.MODEL.SWIN_MLP.IN_CHANS,\n#                         num_classes=config.MODEL.NUM_CLASSES,\n#                         embed_dim=config.MODEL.SWIN_MLP.EMBED_DIM,\n#                         depths=config.MODEL.SWIN_MLP.DEPTHS,\n#                         num_heads=config.MODEL.SWIN_MLP.NUM_HEADS,\n#                         window_size=config.MODEL.SWIN_MLP.WINDOW_SIZE,\n#                         mlp_ratio=config.MODEL.SWIN_MLP.MLP_RATIO,\n#                         drop_rate=config.MODEL.DROP_RATE,\n#                         drop_path_rate=config.MODEL.DROP_PATH_RATE,\n#                         ape=config.MODEL.SWIN_MLP.APE,\n#                         patch_norm=config.MODEL.SWIN_MLP.PATCH_NORM,\n#                         use_checkpoint=config.TRAIN.USE_CHECKPOINT)\n#     else:\n#         raise NotImplementedError(f\"Unkown model: {model_type}\")\n\n#     return model","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model=build_model()\n# print(model)","metadata":{},"execution_count":null,"outputs":[]}]}