{"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-20T18:39:39.97258Z","iopub.execute_input":"2024-02-20T18:39:39.972863Z","iopub.status.idle":"2024-02-20T18:40:31.864591Z","shell.execute_reply.started":"2024-02-20T18:39:39.972816Z","shell.execute_reply":"2024-02-20T18:40:31.863878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch.nn as nn\n# import itertools\n# from collections.abc import Sequence\n# from monai.networks.layers import Conv, trunc_normal_\n# class 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":{"execution":{"iopub.status.busy":"2024-02-20T17:04:29.999537Z","iopub.execute_input":"2024-02-20T17:04:30.000462Z","iopub.status.idle":"2024-02-20T17:04:30.043348Z","shell.execute_reply.started":"2024-02-20T17:04:30.000411Z","shell.execute_reply":"2024-02-20T17:04:30.041357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch.nn as nn\n# import itertools\n# from collections.abc import Sequence\n# from monai.networks.layers import Conv, trunc_normal_\n# class 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#         print(x.shape , \"for enteting into window attention\")\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#         print(\"Exiting winfdoe attention\", x.shape)\n#         return x\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-02-20T16:52:42.369455Z","iopub.execute_input":"2024-02-20T16:52:42.369944Z","iopub.status.idle":"2024-02-20T16:53:24.570042Z","shell.execute_reply.started":"2024-02-20T16:52:42.369898Z","shell.execute_reply":"2024-02-20T16:53:24.568732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch.nn as nn\n# import torch.nn.functional as F\n# from monai.networks.layers import Conv, trunc_normal_\n\n# class WindowAttention(nn.Module):\n#     def __init__(self, dim, num_heads, window_size=(7,7,7), qkv_bias=False, attn_drop=0.0, proj_drop=0.0):\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\n#         # VSR module\n#         self.avgpool = nn.AvgPool2d(window_size, stride=1)\n#         self.leakyrelu = nn.LeakyReLU(negative_slope=0.01)\n#         self.conv = ConvConv.CONV, 2\n\n#         # Rest of the code remains the same\n\n#     def forward(self, x, mask=None):\n#         B, N, C = x.shape\n#         x = x.view(B, N, self.num_heads, C // self.num_heads).permute(0, 2, 1, 3)\n\n#         # VSR module\n#         x = self.avgpool(x)\n#         x = self.leakyrelu(x)\n#         S_w, O_w = self.conv(x).chunk(2, dim=-1)\n#         S_w = S_w.sigmoid()\n#         O_w = O_w.tanh()\n\n#         # Apply S_w and O_w to x here according to your methodology\n\n#         # Rest of the code\n#         relative_position_bias = self.relative_position_bias_table[self.relative_position_index.view(-1)].view(\n#             b, n, n, -1\n#         )\n#         relative_position_bias = relative_position_bias.permute(3, 0, 1, 2).contiguous()\n\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","metadata":{"execution":{"iopub.status.busy":"2024-02-20T13:44:47.602161Z","iopub.execute_input":"2024-02-20T13:44:47.603309Z","iopub.status.idle":"2024-02-20T13:44:47.626018Z","shell.execute_reply.started":"2024-02-20T13:44:47.603259Z","shell.execute_reply":"2024-02-20T13:44:47.624496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom collections.abc import Sequence\n\nclass WindowAttention(nn.Module):\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        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\n        # Define relative position bias table\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            coords = torch.stack(torch.meshgrid(coords_d, coords_h, coords_w, indexing=\"ij\"))\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 * self.window_size[0] - 1) * (2 * self.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            coords = torch.stack(torch.meshgrid(coords_h, coords_w, indexing=\"ij\"))\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\n        # Define query, key, value linear layers\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        self.softmax = nn.Softmax(dim=-1)\n\n    def forward(self, x, mask=None):\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","metadata":{"execution":{"iopub.status.busy":"2024-02-20T18:40:34.098384Z","iopub.execute_input":"2024-02-20T18:40:34.098851Z","iopub.status.idle":"2024-02-20T18:40:35.183335Z","shell.execute_reply.started":"2024-02-20T18:40:34.098804Z","shell.execute_reply":"2024-02-20T18:40:35.182505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch\n# import torch.nn as nn\n\n# class VariedSizeWindowAttention(nn.Module):\n#     def __init__(\n#         self,\n#         dim,\n#         num_heads,\n#         window_size,\n#         qkv_bias=False,\n#         attn_drop=0.0,\n#         proj_drop=0.0,\n#     ):\n#         super().__init__()\n#         self.dim = dim\n#         self.num_heads = num_heads\n#         self.window_size = window_size\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#         self.softmax = nn.Softmax(dim=-1)\n\n#         # Window regression module\n#         self.window_regressor = nn.Linear(dim, 4)  # Predicts size and location of the target window\n\n#     def forward(self, x):\n#         b, n, c = x.shape\n        \n#         # Perform linear transformation for query, key, and value\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#         q, k, v = qkv[0], qkv[1], qkv[2]\n#         q = q * (self.dim // self.num_heads) ** -0.5\n\n#         # Predict size and location of the target window\n#         window_params = self.window_regressor(x)  # Shape: (batch_size, sequence_length, 4)\n#         target_window = self.get_target_window(window_params)  # Get target window coordinates\n\n#         # Sample key and value tokens from the target window\n#         k, v = self.sample_from_window(k, v, target_window)\n\n#         # Perform attention computation\n#         attn = (q @ k.transpose(-2, -1)) * (self.window_size[0] ** -0.5)  # Scale by window size\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#         x = self.proj(x)\n#         x = self.proj_drop(x)\n#         return x\n\n#     def get_target_window(self, window_params):\n#         # Extract predicted parameters for the target window\n#         target_center_x = window_params[:, :, 0]  # Predicted x-coordinate of the window center\n#         target_center_y = window_params[:, :, 1]  # Predicted y-coordinate of the window center\n#         target_size_x = window_params[:, :, 2]    # Predicted width of the window\n#         target_size_y = window_params[:, :, 3]    # Predicted height of the window\n\n#         # Calculate coordinates of the target window\n#         target_left = target_center_x - target_size_x / 2\n#         target_right = target_center_x + target_size_x / 2\n#         target_top = target_center_y - target_size_y / 2\n#         target_bottom = target_center_y + target_size_y / 2\n\n#         return target_left, target_right, target_top, target_bottom\n\n#     def sample_from_window(self, k, v, target_window):\n#         # Extract target window coordinates\n#         target_left, target_right, target_top, target_bottom = target_window\n\n#         # Sample key and value tokens from the target window\n#         sampled_k = self.sample_tokens(k, target_left, target_right, target_top, target_bottom)\n#         sampled_v = self.sample_tokens(v, target_left, target_right, target_top, target_bottom)\n\n#         return sampled_k, sampled_v\n\n#     def sample_tokens(self, tokens, left, right, top, bottom):\n#         # Perform token sampling from the target window\n#         batch_size, num_heads, seq_len, head_dim = tokens.size()\n#         tokens = tokens.permute(0, 2, 1, 3).reshape(batch_size, seq_len, -1)  # Reshape for sampling\n#         sampled_tokens = []\n\n#         for i in range(batch_size):\n#             for j in range(seq_len):\n#                 token_x, token_y = j % self.window_size[0], j // self.window_size[0]  # Token coordinates\n#                 if left[i, j] <= token_x <= right[i, j] and top[i, j] <= token_y <= bottom[i, j]:\n#                     sampled_tokens.append(tokens[i, j])  # Token is inside the target window\n\n#         sampled_tokens = torch.stack(sampled_tokens, dim=0) if len(sampled_tokens) > 0 else torch.zeros(0)  # Stack sampled tokens\n#         return sampled_tokens.view(batch_size, num_heads, -1, head_dim).permute(0, 2, 1, 3)  # Reshape back\n\n# import torch.nn as nn\n# # from models.TransBTS.IntmdSequential import IntermediateSequential \n# class 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\n# class 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\n# class 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\n# class 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\n# class 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\n# class 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# class 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#         window_size=(7, 7, 7)\n#     ):\n#         super().__init__()\n#         self.dim = dim\n#         self.depth = depth\n#         self.heads = heads\n#         self.mlp_dim = mlp_dim\n#         self.dropout_rate = dropout_rate\n#         self.attn_dropout_rate = attn_dropout_rate\n#         self.window_size = window_size\n\n#         self.layers = nn.ModuleList([\n#             nn.ModuleList([\n#                 Residual(PreNormDrop(dim, dropout_rate,\n#                             VariedSizeWindowAttention(dim, num_heads=heads,\n#                                             window_size=window_size,\n#                                             attn_drop=attn_dropout_rate,\n#                                             proj_drop=dropout_rate)),\n#                 ),\n#                 Residual(PreNorm(dim, FeedForward(dim, mlp_dim, dropout_rate))),\n#             ])\n#             for _ in range(depth)\n#         ])\n\n#     def forward(self, x):\n#         for attn, ff in self.layers:\n#             x = attn(x)\n#             x = ff(x)\n#         return x\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T13:31:34.802072Z","iopub.execute_input":"2024-02-20T13:31:34.802561Z","iopub.status.idle":"2024-02-20T13:31:34.819476Z","shell.execute_reply.started":"2024-02-20T13:31:34.802523Z","shell.execute_reply":"2024-02-20T13:31:34.818405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch.nn as nn\n# from collections.abc import Sequence\n# from monai.networks.layers import Conv, trunc_normal_\n\n# class WindowAttention(nn.Module):\n\n#     def __init__(\n#         self,\n#         dim: int,\n#         num_heads: int,\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.num_heads = num_heads\n#         head_dim = dim // num_heads\n#         self.scale = head_dim**-0.5\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\n#     def forward(self, x, mask=None):\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\n#         attn = self.softmax(attn)\n#         attn = self.attn_drop(attn).to(v.dtype)\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","metadata":{"execution":{"iopub.status.busy":"2024-02-20T13:31:35.142532Z","iopub.execute_input":"2024-02-20T13:31:35.143004Z","iopub.status.idle":"2024-02-20T13:31:35.152137Z","shell.execute_reply.started":"2024-02-20T13:31:35.142965Z","shell.execute_reply":"2024-02-20T13:31:35.150589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch\n# import torch.nn as nn\n# from monai.networks.layers import Conv, trunc_normal_\n# import torch.nn as nn\n# import torch.nn.functional as F\n# from collections.abc import Sequence\n# from monai.networks.layers import Conv, trunc_normal_\n\n# class VariedSizeWindowAttention(nn.Module):\n\n#     def __init__(\n#         self,\n#         dim: int,\n#         num_heads: int,\n#         window_size: Sequence[int] = (7, 7),  # Adjusted for 2D input\n#         qkv_bias: bool = False,\n#         attn_drop: float = 0.0,\n#         proj_drop: float = 0.0,\n#     ) -> None:\n#         super().__init__()\n#         self.dim = dim\n#         self.num_heads = num_heads\n#         self.window_size = window_size\n#         head_dim = dim // num_heads\n#         self.scale = head_dim**-0.5\n\n#         # VSA components\n#         self.vsr_module = nn.Sequential(\n#             nn.AdaptiveAvgPool2d(output_size=window_size),  # Adjusted for adaptive pooling\n#             nn.LeakyReLU(),\n#             nn.Conv2d(dim, 2 * num_heads, kernel_size=1, stride=1)  # Predict scales and offsets\n#         )\n#         self.cpe = Conv(dim, dim, kernel_size=window_size, stride=1, padding=window_size[0] // 2)\n\n#         # Window attention components\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\n#     def forward(self, x, mask=None):\n#         b, n, c = x.shape\n\n#         # Varied-size window regression\n#         sw, ow = self.vsr_module(x).chunk(2, dim=1)\n#         sw = sw.sigmoid()  # Scale to [0, 1]\n#         ow = ow.tanh()  # Offset to [-1, 1]\n\n#         # Get key/value tokens and apply conditional position embedding\n#         k, v = self.qkv(x).reshape(b, n, 3, self.num_heads, c // self.num_heads).permute(2, 0, 3, 1, 4)\n#         k, v = self.cpe(k), self.cpe(v)\n\n#         # Sample from varied-size windows\n#         k_sampled = F.grid_sample(k, grid=(sw + ow).permute(0, 2, 3, 1), mode='bilinear', padding_mode='zeros')  # Adaptive sampling\n#         v_sampled = F.grid_sample(v, grid=(sw + ow).permute(0, 2, 3, 1), mode='bilinear', padding_mode='zeros')\n\n#         # Attention calculation with sampled tokens\n#         q = q * self.scale\n#         attn = q @ k_sampled.transpose(-2, -1)\n#         # ... (rest of the attention calculation and output)\n\n#         attn = attn.softmax(dim=-1)\n#         attn = self.attn_drop(attn)\n\n#         x = (attn @ kv).transpose(1, 2).reshape(b, h * w, c)\n#         x = x.permute(0, 2, 1).reshape(b, c, h, w)  # Reshape back\n\n#         # Projection\n#         x = self.proj(x)\n#         x = self.proj_drop(x)\n\n#         return x\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T14:24:22.206651Z","iopub.execute_input":"2024-02-20T14:24:22.207186Z","iopub.status.idle":"2024-02-20T14:24:22.231444Z","shell.execute_reply.started":"2024-02-20T14:24:22.20715Z","shell.execute_reply":"2024-02-20T14:24:22.230408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch.nn as nn\n# import torch\n\n# class VariedSizeWindowRegression(nn.Module):\n#     def __init__(self, in_channels, out_channels):\n#         super().__init__()\n         \n#         self.avg_pool = nn.AvgPool2d(kernel_size=window_size, stride=window_size),\n#         self.conv1x1 = nn.Conv2d(512, 4 * 2, kernel_size=1, stride=1)\n#         self.leaky_relu = nn.LeakyReLU()\n\n#     def forward(self, x):\n#         x = self.avg_pool(x)\n#         x = self.leaky_relu(x)\n#         x = self.conv1x1(x)\n#         return x\n\n# class WindowAttention(nn.Module):\n#     def __init__(\n#         self,\n#         dim: int,\n#         num_heads: int,\n#         window_size: torch.Size,\n#         qkv_bias: bool = False,\n#         attn_drop: float = 0.0,\n#         proj_drop: float = 0.0,\n#     ) -> None:\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#         self.vs_regression = VariedSizeWindowRegression(dim, dim)\n#         self.mlp = nn.Sequential(\n#             nn.Linear(dim, dim),\n#             nn.LeakyReLU(),\n#             nn.Linear(dim, dim)\n#         )\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#         print(x.shape, \"for entering into window attention\")\n#         b, n, c = x.shape\n#         print(x.shape, b, n, c, \"...........................................\")\n        \n#         # Preprocessing with Varied-size Window Regression\n#         x = self.vs_regression(x)\n        \n#         qkv = self.mlp(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        \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#         print(\"Exiting window attention\", x.shape)\n#         return x\n\n# # Example usage:\n# input_shape = (32, 8, 512)\n# num_heads = 4\n# window_size = (7, 7, 7)\n# attention = WindowAttention(dim=512, num_heads=num_heads, window_size=window_size)\n# input_tensor = torch.randn(input_shape)\n# output = attention(input_tensor)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:54:37.344417Z","iopub.execute_input":"2024-02-20T16:54:37.344805Z","iopub.status.idle":"2024-02-20T16:54:37.842961Z","shell.execute_reply.started":"2024-02-20T16:54:37.344774Z","shell.execute_reply":"2024-02-20T16:54:37.841311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch\n# import torch.nn as nn\n# import torch.nn.functional as F\n# from collections.abc import Sequence\n# from monai.networks.layers import Conv, trunc_normal_\n\n# class WindowAttention(nn.Module):\n#     def __init__(\n#         self,\n#         dim: int,\n#         num_heads: int,\n#         window_size: Sequence[int] = (7, 7),  # Change default window_size to 2D\n#         qkv_bias: bool = False,\n#         attn_drop: float = 0.0,\n#         proj_drop: float = 0.0,\n#     ) -> None:\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#         # Define sampling coefficients layers\n#         self.sampling_coefficients = nn.Sequential(\n#             nn.AvgPool2d(kernel_size=(window_size[0], window_size[1])),  # Omit stride for no striding\n#             nn.LeakyReLU(),\n#             nn.Conv2d(36, self.num_heads * 2, kernel_size=1, stride=1)\n#         )\n        \n#         if len(self.window_size) == 3:\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#             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#         self.relative_position_bias_table = nn.Parameter(\n#             torch.zeros_like(relative_position_index, dtype=torch.float32)\n#         )\n#         trunc_normal_(self.relative_position_bias_table, std=0.02)\n#         self.softmax = nn.Softmax(dim=-1)\n\n#         # Initialize parameters\n#         self.apply(self._init_weights)\n\n#     def forward(self, x, mask=None):\n#         b, n, c = x.shape\n\n#         # Compute sampling coefficients\n#         sampling_coeffs = self.sampling_coefficients(x)\n\n#         # Compute sampling offsets and scales if applicable\n#         if hasattr(self, 'sampling_offsets'):\n#             sampling_offsets = self.sampling_offsets(x)\n#             sampling_scales = self.sampling_scales(x)\n#             # Adjust coordinates using sampling offsets and scales\n#             # ...\n\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#     def _compute_relative_position_index(self):\n#         if len(self.window_size) == 3:\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#             coords = torch.stack(torch.meshgrid(coords_d, coords_h, coords_w, indexing=\"ij\"))\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#             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, indexing=\"ij\"))\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#         relative_position_index = relative_coords.sum(-1)\n#         return relative_position_index\n\n#     def _init_weights(self, m):\n#         if isinstance(m, nn.Linear):\n#             trunc_normal_(m.weight, std=0.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.Conv2d):\n#             trunc_normal_(m.weight, std=0.02)\n#             if isinstance(m, nn.Conv2d) and m.bias is not None:\n#                 nn.init.constant_(m.bias, 0)\n# # Example usage:\n# input_shape = (32, 8, 512)\n# num_heads = 4\n# window_size = (7, 7, 7)\n# attention = WindowAttention(dim=512, num_heads=num_heads, window_size=window_size)\n# input_tensor = torch.randn(input_shape)\n# output = attention(input_tensor)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T15:40:55.896429Z","iopub.execute_input":"2024-02-20T15:40:55.897091Z","iopub.status.idle":"2024-02-20T15:40:56.347935Z","shell.execute_reply.started":"2024-02-20T15:40:55.897039Z","shell.execute_reply":"2024-02-20T15:40:56.345824Z"},"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\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 i in range(depth):\n            # Alternate between SelfAttention and WindowAttention\n            if i % 2 == 0:\n                attention_layer = SelfAttention(dim, heads=heads, dropout_rate=attn_dropout_rate)\n            else:\n                attention_layer = WindowAttention(\n                    dim=dim,\n                    num_heads=heads,\n                    window_size=(7, 7, 7),\n                    qkv_bias=True,\n                    attn_drop=attn_dropout_rate,\n                    proj_drop=0.1\n                )\n\n            layers.extend(\n                [\n                    Residual(\n                        PreNormDrop(\n                            dim,\n                            dropout_rate,\n                            attention_layer\n                        )\n                    ),\n                    Residual(\n                        PreNorm(dim, FeedForward(dim, mlp_dim, dropout_rate))\n                    ),\n                ]\n            )\n        self.net = IntermediateSequential(*layers)\n\n    def forward(self, x):\n        return self.net(x)\n\nmodel = TransformerModel(512, 4, 8, 4096, 0.1, 0.1)\nprint(model)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T18:40:44.413749Z","iopub.execute_input":"2024-02-20T18:40:44.414614Z","iopub.status.idle":"2024-02-20T18:40:44.692396Z","shell.execute_reply.started":"2024-02-20T18:40:44.414582Z","shell.execute_reply":"2024-02-20T18:40:44.691582Z"},"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\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        self.conv_layers = nn.Sequential(\n            nn.Conv3d(dim, dim, kernel_size=(1, 3, 3), stride=(1, 1, 1), padding=(0, 1, 1)),\n            nn.BatchNorm3d(dim),\n            nn.ReLU(),\n            nn.Conv3d(dim, dim, kernel_size=(1, 3, 3), stride=(1, 1, 1), padding=(0, 1, 1)),\n            nn.BatchNorm3d(dim),\n            nn.ReLU(),\n            nn.Conv3d(dim, dim, kernel_size=(1, 3, 3), stride=(1, 1, 1), padding=(0, 1, 1)),\n            nn.BatchNorm3d(dim),\n            nn.ReLU()\n        )\n#         self.transformer = nn.TransformerEncoder(\n#             nn.TransformerEncoderLayer(d_model=dim, nhead=num_heads, dim_feedforward=dim_feedforward, dropout=dropout),\n#             num_layers=num_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        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_si\n# ze[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        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)\nprint(model)\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(\"window attention \",classification_outputs)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T18:40:50.68134Z","iopub.execute_input":"2024-02-20T18:40:50.68163Z","iopub.status.idle":"2024-02-20T18:41:09.593973Z","shell.execute_reply.started":"2024-02-20T18:40:50.681604Z","shell.execute_reply":"2024-02-20T18:41:09.593307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch.nn as nn\n# import torch\n# class VariedSizeWindowRegression(nn.Module):\n#     def __init__(self, in_channels, out_channels, window_size=(7, 7, 7)):\n#         super().__init__()\n#         self.avg_pool = nn.AdaptiveAvgPool3d(1)\n#         self.conv1x1 = nn.Conv3d(in_channels, out_channels, kernel_size=1, stride=1)\n#         self.leaky_relu = nn.LeakyReLU()\n#         self.window_size = window_size\n\n#     def forward(self, x):\n#         b, c, d, h, w = x.shape\n#         # Reshape the input tensor to make it compatible with the Conv3d layer\n#         x = x.reshape(b * c, 1, d, h, w)\n#         x = self.conv1x1(x)\n#         x = self.leaky_relu(x)\n#         # Reshape back to the original shape\n#         x = x.reshape(b, c, -1)\n#         return x\n\n\n\n# class WindowAttention(nn.Module):\n#     def __init__(\n#         self,\n#         dim: int,\n#         num_heads: int,\n#         window_size: torch.Size,\n#         qkv_bias: bool = False,\n#         attn_drop: float = 0.0,\n#         proj_drop: float = 0.0,\n#     ) -> None:\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#         self.vs_regression = VariedSizeWindowRegression(dim, dim)\n#         self.mlp = nn.Sequential(\n#             nn.Linear(dim, dim),\n#             nn.LeakyReLU(),\n#             nn.Linear(dim, dim)\n#         )\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#         print(x.shape, \"for entering into window attention\")\n#         b, n, c = x.shape\n#         print(x.shape, b, n, c, \"...........................................\")\n        \n#         # Preprocessing with Varied-size Window Regression\n#         x = self.vs_regression(x)\n        \n#         qkv = self.mlp(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        \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#         print(\"Exiting window attention\", x.shape)\n#         return x\n\n# # Example usage:\n# input_shape = (32, 8, 512)\n# num_heads = 4\n# window_size = (7, 7, 7)\n# attention = WindowAttention(dim=input_shape[-1], num_heads=num_heads, window_size=window_size)\n# input_tensor = torch.randn(input_shape)\n# output = attention(input_tensor)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T14:21:05.722234Z","iopub.execute_input":"2024-02-20T14:21:05.722678Z","iopub.status.idle":"2024-02-20T14:21:06.072419Z","shell.execute_reply.started":"2024-02-20T14:21:05.722644Z","shell.execute_reply":"2024-02-20T14:21:06.069971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import os\n# import pandas as pd\n# import numpy as np\n# import nibabel as nib\n# import torch\n# import torch.nn as nn\n# import torch.optim as optim\n# from torch.utils.data import Dataset, DataLoader\n# from torchvision import transforms\n# import numpy as np\n# from sklearn import metrics\n# from sklearn.metrics import precision_recall_fscore_support\n# from sklearn.model_selection import train_test_split\n# from scipy.ndimage import zoom\n# from sklearn.metrics import confusion_matrix\n# from sklearn.metrics import accuracy_score, precision_score\n# # Create a function to move data to the device\n# def move_data_to_device(data, device):\n#     return data.to(torch.float32).to(device)\n\n# class 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\n# import torch.nn.functional as F\n\n# # Function to resize NIfTI data\n# def 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\n# csv_file = '/kaggle/input/unhealthy-csv-file/combined_data (6).csv'  # Update with the correct path\n# batch_size = 32\n# num_workers = 4  # Number of CPU cores to use for data loading\n# num_classes = 14  # Number of classes\n# desired_shape = (128, 128, 128)\n# device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n# # Define transformations if needed\n# transform = transforms.Compose([\n#     transforms.ToTensor(),  # Convert to tensor\n#     # Add more transformations if necessary\n# ])\n\n# # Load the CSV file\n# data = pd.read_csv(csv_file).head(1280)\n\n# # Remove the extra space from the column name\n# data.columns = data.columns.str.strip()\n\n# # Assuming 'data' is your DataFrame\n# data_length = len(data)\n# print(\"Length of DataFrame:\", data_length)\n\n# # Split the data into training, validation, and test sets\n# train_data, temp_data = train_test_split(data, test_size=0.2, random_state=42)\n# val_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\n# pd.set_option('display.max_rows', None)\n\n# index_values = train_data.index.values\n\n# # Reset the display option to its default value (if needed)\n# pd.reset_option('display.max_rows')\n\n# # Extract file paths and labels from the data\n# train_paths = train_data['file_path'].values\n# train_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\n# val_paths = val_data['file_path'].values\n# val_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\n# test_paths = test_data['file_path'].values\n# test_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\n# train_dataset = CustomDataset(train_paths, train_labels, transform=transform)\n# print('len of train_dataset', len(train_dataset))\n# val_dataset = CustomDataset(val_paths, val_labels, transform=transform)\n# test_dataset = CustomDataset(test_paths, test_labels, transform=transform)\n\n# # Instantiate the data loaders\n# train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers, drop_last=True)\n# print('train_loader', len(train_loader))\n# val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers)\n# test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers)\n# print(train_loader)\n\n# # Instantiate the model with the appropriate number of classes for classification\n# in_channels = 1  # Input channels (e.g., for grayscale images or volumes)\n# num_classes_classification = 14  # Number of classes for classification\n# model_class = ConvolutionalVisionTransformer(in_channels, num_classes_classification)\n\n# # Count the number of parameters\n# total_params_class = sum(p.numel() for p in model_class.parameters())\n# print(f\"Total Trainable Parameters for Classification: {total_params_class}\")\n\n# # Define loss function and optimizer\n# class_criterion = nn.CrossEntropyLoss()  # Binary Cross-Entropy loss for classification\n# class_optimizer = optim.Adam(model_class.parameters(), lr=0.001)\n\n# # Training loop\n# # Training loop\n# class_labels = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen', 'any_injury']\n\n\n# # Training loop\n# num_epochs = 500\n# for 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    \n# with 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\n# with 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-02-20T13:30:18.105666Z","iopub.execute_input":"2024-02-20T13:30:18.106141Z","iopub.status.idle":"2024-02-20T13:30:18.808639Z","shell.execute_reply.started":"2024-02-20T13:30:18.106103Z","shell.execute_reply":"2024-02-20T13:30:18.806848Z"},"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\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\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\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\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.BCEWithLogitsLoss()  # 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\ntrue_pos=[]\ntrue_neg=[]\nfalse_pos=[]\nfalse_neg=[]\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    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\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        \n        # Apply sigmoid activation to the classification outputs\n        classification_outputs = torch.sigmoid(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        # Calculate accuracy and precision\n        predicted_labels = (classification_outputs > 0.5).float()\n        true_positives = (predicted_labels * batch_labels).sum(dim=0)\n        true_pos.append(true_positives)\n        false_positives = ((1 - batch_labels) * predicted_labels).sum(dim=0)\n        false_pos.append(false_positives)\n        false_negatives = (batch_labels * (1 - predicted_labels)).sum(dim=0)\n        false_neg.append(false_negatives)\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        true_neg.append(true_negatives)\n        precision = true_positives / (true_positives + false_positives)\n        \n        print(\"Total Classification Loss:\", class_loss.item())\n        print(\"precision\",precision)\n        print(\"accuracy\",accuracy)\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/transformerWA_epoch_{epoch}.pth')\n    accuracy = (sum(true_pos) + sum(true_neg)) / (sum(true_pos) + sum(true_neg) + sum(false_pos) + sum(false_neg))\n    precision = sum(true_pos)/ (sum(true_pos) + sum(false_pos))\n    print(\"precision\",precision)\n    print(\"accuracy\",accuracy)\n\n    # Calculate evaluation metrics after all epochs\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\n        all_predicted_labels = np.concatenate(all_predicted_labels, axis=0)\n        all_batch_labels = np.concatenate(all_batch_labels, axis=0)\n        \n        # Calculate accuracy, precision, recall, and F1 score\n        true_positives = np.sum(all_predicted_labels * all_batch_labels, axis=0)\n        false_positives = np.sum(all_predicted_labels * (1 - all_batch_labels), axis=0)\n        false_negatives = np.sum((1 - all_predicted_labels) * all_batch_labels, axis=0)\n\n        micro_precision = np.sum(true_positives) / (np.sum(true_positives) + np.sum(false_positives))\n        micro_accuracy = np.mean((all_predicted_labels == all_batch_labels).all(axis=1))  # Correctly predict all labels\n        micro_recall = np.sum(true_positives) / (np.sum(true_positives) + np.sum(false_negatives))\n        micro_f1_score = 2 * micro_precision * micro_recall / (micro_precision + micro_recall)\n        num_classes = all_batch_labels.shape[1]\n        macro_precision = np.mean(true_positives / (true_positives + false_positives), axis=0)\n        macro_accuracy = np.mean((all_predicted_labels == all_batch_labels).all(axis=1))  # Same as micro-accuracy\n        macro_recall = np.mean(true_positives / (true_positives + false_negatives), axis=0)\n        macro_f1_score = 2 * macro_precision * macro_recall / (macro_precision + macro_recall)\n\n        print(\"Micro Precision:\", micro_precision)\n        print(\"Micro Accuracy:\", micro_accuracy)\n        print(\"Micro F1-score:\", micro_f1_score)\n        print(\"Macro Precision:\", macro_precision)\n        print(\"Macro Accuracy:\", macro_accuracy)\n        print(\"Macro Recall:\", macro_recall)\n        print(\"Macro F1-score:\", macro_f1_score)\n        j=0\n        for i in range(0,4,2):\n                    total_true_positives = 0\n                    total_false_positives = 0\n                    total_false_negatives = 0\n                    total_true_negatives = 0\n                    \n                    bowel_pred=predicted_labels[:,i:i+2]\n                    bowel_truth_label=batch_labels[:,i:i+2]\n                    print(bowel_pred.shape)\n                    print(bowel_truth_label.shape)\n                    true_positives = (bowel_pred * bowel_truth_label).sum(dim=0)\n                    false_positives = ((1 - bowel_truth_label) * bowel_pred).sum(dim=0)\n                    false_negatives = (bowel_truth_label * (1 - bowel_pred)).sum(dim=0)\n                    true_negatives = ((1 - bowel_truth_label) * (1 - bowel_pred)).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                    print(class_labels[j])\n                    j=j+1\n                    total_true_positives += true_positives.sum()\n                    total_false_positives += false_positives.sum()\n                    total_false_negatives += false_negatives.sum()\n                    total_true_negatives += true_negatives.sum()\n\n# Calculate micro-averaged accuracy and precision\n                    micro_accuracy = (total_true_positives + total_true_negatives) / (total_true_positives + total_true_negatives + total_false_positives + total_false_negatives)\n                    micro_precision = total_true_positives / (total_true_positives + total_false_positives)\n                    print(\"Micro Precision:\", micro_precision)\n                    print(\"Micro Accuracy:\", micro_accuracy)\n                    \n        for i in range(4,14,3):\n                    total_true_positives = 0\n                    total_false_positives = 0\n                    total_false_negatives = 0\n                    total_true_negatives = 0\n                    bowel_pred=predicted_labels[:,i:i+3]\n                    bowel_truth_label=batch_labels[:,i:i+3]\n                    print(bowel_pred.shape)\n                    print(bowel_truth_label.shape)\n                    true_positives = (bowel_pred * bowel_truth_label).sum(dim=0)\n                    false_positives = ((1 - bowel_truth_label) * bowel_pred).sum(dim=0)\n                    false_negatives = (bowel_truth_label * (1 - bowel_pred)).sum(dim=0)\n                    true_negatives = ((1 - bowel_truth_label) * (1 - bowel_pred)).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                    print(class_labels[j])\n                    j=j+1\n                    total_true_positives += true_positives.sum()\n                    total_false_positives += false_positives.sum()\n                    total_false_negatives += false_negatives.sum()\n                    total_true_negatives += true_negatives.sum()\n\n# Calculate micro-averaged accuracy and precision\n                    micro_accuracy = (total_true_positives + total_true_negatives) / (total_true_positives + total_true_negatives + total_false_positives + total_false_negatives)\n                    micro_precision = total_true_positives / (total_true_positives + total_false_positives)\n                    print(\"Micro Precision:\", micro_precision)\n                    print(\"Micro Accuracy:\", micro_accuracy)\n        ","metadata":{"execution":{"iopub.status.busy":"2024-02-16T05:36:39.038218Z","iopub.status.idle":"2024-02-16T05:36:39.038871Z","shell.execute_reply.started":"2024-02-16T05:36:39.038556Z","shell.execute_reply":"2024-02-16T05:36:39.038588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}