{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":36363,"databundleVersionId":4050810,"sourceType":"competition"},{"sourceId":4407134,"sourceType":"datasetVersion","datasetId":2556595},{"sourceId":4454029,"sourceType":"datasetVersion","datasetId":2607864},{"sourceId":7910796,"sourceType":"datasetVersion","datasetId":4647323},{"sourceId":7911066,"sourceType":"datasetVersion","datasetId":4647528},{"sourceId":7921148,"sourceType":"datasetVersion","datasetId":4654903},{"sourceId":7921218,"sourceType":"datasetVersion","datasetId":4654958},{"sourceId":7921255,"sourceType":"datasetVersion","datasetId":4654987},{"sourceId":7971413,"sourceType":"datasetVersion","datasetId":4690515},{"sourceId":7988724,"sourceType":"datasetVersion","datasetId":4702748},{"sourceId":8392679,"sourceType":"datasetVersion","datasetId":4882013},{"sourceId":7144712,"sourceType":"kernelVersion"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install einops","metadata":{"execution":{"iopub.status.busy":"2024-05-12T18:37:23.79892Z","iopub.execute_input":"2024-05-12T18:37:23.79922Z","iopub.status.idle":"2024-05-12T18:37:38.064786Z","shell.execute_reply.started":"2024-05-12T18:37:23.799193Z","shell.execute_reply":"2024-05-12T18:37:38.063569Z"},"trusted":true},"execution_count":1,"outputs":[{"name":"stdout","text":"Collecting einops\n  Downloading einops-0.8.0-py3-none-any.whl.metadata (12 kB)\nDownloading einops-0.8.0-py3-none-any.whl (43 kB)\n\u001b[2K   \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m43.2/43.2 kB\u001b[0m \u001b[31m1.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hInstalling collected packages: einops\nSuccessfully installed einops-0.8.0\n","output_type":"stream"}]},{"cell_type":"code","source":"import torch\nprint(torch.version.cuda)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T09:54:37.98741Z","iopub.execute_input":"2024-04-25T09:54:37.987738Z","iopub.status.idle":"2024-04-25T09:54:39.691992Z","shell.execute_reply.started":"2024-04-25T09:54:37.987703Z","shell.execute_reply":"2024-04-25T09:54:39.69108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install python\n# !pip install pytorch\n!pip install torchaudio\n!pip install torchvision\n# !pip install cudatoolkit\n!pip install pytorch-lightning","metadata":{"execution":{"iopub.status.busy":"2024-04-26T05:19:51.56343Z","iopub.execute_input":"2024-04-26T05:19:51.563879Z","iopub.status.idle":"2024-04-26T05:20:30.717475Z","shell.execute_reply.started":"2024-04-26T05:19:51.563848Z","shell.execute_reply":"2024-04-26T05:20:30.716369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install einops","metadata":{"execution":{"iopub.status.busy":"2024-04-25T09:55:17.119272Z","iopub.execute_input":"2024-04-25T09:55:17.119675Z","iopub.status.idle":"2024-04-25T09:55:29.381645Z","shell.execute_reply.started":"2024-04-25T09:55:17.119639Z","shell.execute_reply":"2024-04-25T09:55:29.380527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install imageio\n!pip install importlib-metadata\n!pip install itk\n!pip install matplotlib\n!pip install MedPy\n!pip install monai\n!pip install timm\n!pip install torchio","metadata":{"execution":{"iopub.status.busy":"2024-04-25T09:55:40.669848Z","iopub.execute_input":"2024-04-25T09:55:40.670228Z","iopub.status.idle":"2024-04-25T09:57:19.825811Z","shell.execute_reply.started":"2024-04-25T09:55:40.670196Z","shell.execute_reply":"2024-04-25T09:57:19.824502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install python_gdcm\n!pip install pylibjpeg\n!pip install pydicom\n!pip install einops\n!pip install monai\nimport time","metadata":{"execution":{"iopub.status.busy":"2024-04-25T09:57:19.828244Z","iopub.execute_input":"2024-04-25T09:57:19.828638Z","iopub.status.idle":"2024-04-25T09:58:22.102418Z","shell.execute_reply.started":"2024-04-25T09:57:19.828605Z","shell.execute_reply":"2024-04-25T09:58:22.101284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport random\nimport re\nfrom dataclasses import dataclass\nfrom typing import Dict\nfrom typing import List\n\nimport albumentations\nimport cv2\nimport numpy as np\nimport pydicom\nimport tifffile\nimport torch\nimport torch.hub\nfrom albumentations import ReplayCompose\nfrom skimage import measure\nfrom torch.functional import Tensor\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2024-04-25T09:58:22.104055Z","iopub.execute_input":"2024-04-25T09:58:22.104472Z","iopub.status.idle":"2024-04-25T09:58:23.139519Z","shell.execute_reply.started":"2024-04-25T09:58:22.104437Z","shell.execute_reply":"2024-04-25T09:58:23.138504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cd /kaggle/input/efficientvit-caps/3D-EffiViTCaps-main/3D-EffiViTCaps-main","metadata":{"execution":{"iopub.status.busy":"2024-04-25T09:58:23.141121Z","iopub.execute_input":"2024-04-25T09:58:23.141528Z","iopub.status.idle":"2024-04-25T09:58:23.149698Z","shell.execute_reply.started":"2024-04-25T09:58:23.141501Z","shell.execute_reply":"2024-04-25T09:58:23.148636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# --------------------------------------------------------\n# EfficientViT3D Model Block\n# Copyright (c) 2024 UESTC\n# Build the EfficientViT3D Model\n# Written by: Dongwei Gan\n# --------------------------------------------------------\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom einops import rearrange\nimport itertools\n\nfrom timm.models.vision_transformer import trunc_normal_\n# from timm.models.layers import SqueezeExcite, SEModule3D\n\n\nclass Conv3d_BN(torch.nn.Sequential):\n    def __init__(self, a, b, ks=1, stride=1, pad=0, dilation=1,\n                 groups=1, bn_weight_init=1, resolution=-10000):\n        super().__init__()\n        self.add_module('c', torch.nn.Conv3d(\n            a, b, ks, stride, pad, dilation, groups, bias=False))\n        self.add_module('bn', torch.nn.BatchNorm3d(b))\n        torch.nn.init.constant_(self.bn.weight, bn_weight_init)\n        torch.nn.init.constant_(self.bn.bias, 0)\n\n    @torch.no_grad()\n    def fuse(self):\n        c, bn = self._modules.values()\n        w = bn.weight / (bn.running_var + bn.eps) ** 0.5\n        w = c.weight * w[:, None, None, None, None]\n        b = bn.bias - bn.running_mean * bn.weight / \\\n            (bn.running_var + bn.eps) ** 0.5\n        m = torch.nn.Conv3d(w.size(1) * self.c.groups, w.size(\n            0), w.shape[2:], stride=self.c.stride, padding=self.c.padding, dilation=self.c.dilation,\n                            groups=self.c.groups)\n        m.weight.data.copy_(w)\n        m.bias.data.copy_(b)\n        return m\n\n\nclass BN_Linear(torch.nn.Sequential):\n    def __init__(self, a, b, bias=True, std=0.02):\n        super().__init__()\n        self.add_module('bn', torch.nn.BatchNorm1d(a))\n        self.add_module('l', torch.nn.Linear(a, b, bias=bias))\n        trunc_normal_(self.l.weight, std=std)\n        if bias:\n            torch.nn.init.constant_(self.l.bias, 0)\n\n    @torch.no_grad()\n    def fuse(self):\n        bn, l = self._modules.values()\n        w = bn.weight / (bn.running_var + bn.eps) ** 0.5\n        b = bn.bias - self.bn.running_mean * \\\n            self.bn.weight / (bn.running_var + bn.eps) ** 0.5\n        w = l.weight * w[None, :]\n        if l.bias is None:\n            b = b @ self.l.weight.T\n        else:\n            b = (l.weight @ b[:, None]).view(-1) + self.l.bias\n        m = torch.nn.Linear(w.size(1), w.size(0))\n        m.weight.data.copy_(w)\n        m.bias.data.copy_(b)\n        return m\n\n\nclass Residual(torch.nn.Module):\n    def __init__(self, m, drop=0.):\n        super().__init__()\n        self.m = m\n        self.drop = drop\n\n    def forward(self, x):\n        if self.training and self.drop > 0:\n            return x + self.m(x) * torch.rand(x.size(0), 1, 1, 1, 1,\n                                              device=x.device).ge_(self.drop).div(1 - self.drop).detach()\n        else:\n            return x + self.m(x)\n\n\nclass FFN(torch.nn.Module):\n    def __init__(self, ed, h, resolution):\n        super().__init__()\n        self.pw1 = Conv3d_BN(ed, h, resolution=resolution)\n        self.act = torch.nn.ReLU()\n        self.pw2 = Conv3d_BN(h, ed, bn_weight_init=0, resolution=resolution)\n\n    def forward(self, x):\n        x = self.pw2(self.act(self.pw1(x)))\n        return x\n\n\nclass CascadedGroupAttention3D(torch.nn.Module):\n    r\"\"\" Cascaded Group Attention.\n\n    Args:\n        dim (int): Number of input channels.\n        key_dim (int): The dimension for query and key.\n        num_heads (int): Number of attention heads.\n        attn_ratio (int): Multiplier for the query dim for value dimension.\n        resolution (int): Input resolution, correspond to the window size.\n        kernels (List[int]): The kernel size of the dw conv on query.\n    \"\"\"\n\n    def __init__(self, dim, key_dim, num_heads=4,\n                 attn_ratio=4,\n                 resolution=14,\n                 kernels=[5, 5, 5, 5], ):\n        super().__init__()\n        self.num_heads = num_heads\n        self.scale = key_dim ** -0.5\n        self.key_dim = key_dim\n        self.d = int(attn_ratio * key_dim)\n        self.attn_ratio = attn_ratio\n\n        qkvs = []\n        dws = []\n        for i in range(num_heads):\n            qkvs.append(Conv3d_BN(dim // (num_heads), self.key_dim * 2 + self.d, resolution=resolution))\n            dws.append(Conv3d_BN(self.key_dim, self.key_dim, kernels[i], 1, kernels[i] // 2, groups=self.key_dim,\n                                 resolution=resolution))\n        self.qkvs = torch.nn.ModuleList(qkvs)\n        self.dws = torch.nn.ModuleList(dws)\n        self.proj = torch.nn.Sequential(torch.nn.ReLU(),\n            Conv3d_BN(self.d * num_heads, dim, bn_weight_init=0, resolution=resolution),)\n        print(resolution, \"resolution\")\n        points = list(itertools.product(range(resolution), range(resolution), range(resolution)))\n        N = len(points)\n        attention_offsets = {}\n        idxs = []\n        for p1 in points:\n            for p2 in points:\n                offset = (abs(p1[0] - p2[0]), abs(p1[1] - p2[1]))\n                if offset not in attention_offsets:\n                    attention_offsets[offset] = len(attention_offsets)\n                idxs.append(attention_offsets[offset])\n        self.attention_biases = torch.nn.Parameter(\n            torch.zeros(num_heads, len(attention_offsets)))\n        self.register_buffer('attention_bias_idxs',\n                             torch.LongTensor(idxs).view(N, N))\n\n    @torch.no_grad()\n    def train(self, mode=True):\n        super().train(mode)\n        if mode and hasattr(self, 'ab'):\n            del self.ab\n        else:\n            self.ab = self.attention_biases[:, self.attention_bias_idxs]\n\n    def forward(self, x):  # x (B,C,H,W,D)\n        B, C, H, W, D = x.shape\n        trainingab = self.attention_biases[:, self.attention_bias_idxs]\n        feats_in = x.chunk(len(self.qkvs), dim=1)\n        feats_out = []\n        feat = feats_in[0]\n        for i, qkv in enumerate(self.qkvs):\n            if i > 0:  # add the previous output to the input\n                feat = feat + feats_in[i]\n            feat = qkv(feat)\n            q, k, v = feat.view(B, -1, H, W, D).split([self.key_dim, self.key_dim, self.d], dim=1)  # B, C/h, H, W, D\n            q = self.dws[i](q)\n            q, k, v = q.flatten(2), k.flatten(2), v.flatten(2)  # B, C/h, N\n            print(q.shape, k.shape, v.shape, \"q.shape, k.shape, v.shape\")\n            attn = (\n                    (q.transpose(-2, -1) @ k) * self.scale\n#                     +\n#                     (trainingab[i] if self.training else self.ab[i])\n            )\n            attn = attn.softmax(dim=-1)  # B N N\n            feat = (v @ attn.transpose(-2, -1)).view(B, self.d, H, W, D)  # B C H W D\n            feats_out.append(feat)\n        x = self.proj(torch.cat(feats_out, 1))\n        return x\n\n\nclass LocalWindowAttention3D(torch.nn.Module):\n    r\"\"\" Local Window Attention.\n\n    Args:\n        dim (int): Number of input channels.\n        key_dim (int): The dimension for query and key.\n        num_heads (int): Number of attention heads.\n        attn_ratio (int): Multiplier for the query dim for value dimension.\n        resolution (int): Input resolution.\n        window_resolution (int): Local window resolution.\n        kernels (List[int]): The kernel size of the dw conv on query.\n    \"\"\"\n\n    def __init__(self, dim, key_dim, num_heads=4,\n                 attn_ratio=4,\n                 resolution=14,\n                 window_resolution=7,\n                 kernels=[5, 5, 5, 5], ):\n        super().__init__()\n        self.dim = dim\n        self.num_heads = num_heads\n        self.resolution = resolution\n        assert window_resolution > 0, 'window_size must be greater than 0'\n        self.window_resolution = window_resolution\n\n        window_resolution = min(window_resolution, resolution)\n        self.attn = CascadedGroupAttention3D(dim, key_dim, num_heads,\n                                           attn_ratio=attn_ratio,\n                                           resolution=window_resolution,\n                                           kernels=kernels, )\n\n    def forward(self, x):\n        \n        B, C, H_, W_, D_ = x.shape\n        H = H_\n        W = W_\n        D = D_ \n        # Only check this for classifcation models\n#         assert H == H_ and W == W_ and D == D_, 'input feature has wrong size, expect {}, got {}'.format((H, W, D), (H_, W_, D_))\n\n        if H <= self.window_resolution and W <= self.window_resolution and D <= self.window_resolution:\n            x = self.attn(x)\n        else:\n            x = x.permute(0, 2, 3, 4, 1)\n            pad_h = (self.window_resolution - H %\n                     self.window_resolution) % self.window_resolution\n            pad_w = (self.window_resolution - W %\n                     self.window_resolution) % self.window_resolution\n            pad_d = (self.window_resolution - D %\n                     self.window_resolution) % self.window_resolution\n            padding = pad_h > 0 or pad_w > 0 or pad_d > 0\n\n            if padding:\n                x = torch.nn.functional.pad(x, (0, 0, 0, pad_d, 0, pad_w, 0, pad_h))\n\n            pH, pW , pD = H + pad_h, W + pad_w, D + pad_d\n            nH = pH // self.window_resolution\n            nW = pW // self.window_resolution\n            nD = pD // self.window_resolution\n            # window partition, B H W D C -> B nH h nW w nD d C -> B nH h nW w nD d C -> B*nH*nW*nD h w d C -> B*nH*nW*nD C h w d\n            print(x.shape) \n            x = x.view(B, nH, self.window_resolution, nW, self.window_resolution, nD, self.window_resolution, C).\\\n                permute(0, 1, 3, 5, 2, 4, 6, 7).reshape(\n                B * nH * nW * nD, self.window_resolution, self.window_resolution, self.window_resolution, C\n            ).permute(0, 4, 1, 2, 3)\n            print(x.shape, \"After permute view..................................\")\n            x = self.attn(x)\n            # window reverse, B*nH*nW*nD C h w d -> B*nH*nW*nD h w d C -> B nH nW nD h w d C -> B nH h nW w nD d C -> B H W D C\n            x = x.permute(0, 2, 3, 4, 1).view(B, nH, nW, nD, self.window_resolution, self.window_resolution,\n            self.window_resolution, C).permute(0, 1, 4, 2, 5, 3, 6, 7).reshape(B, pH, pW, pD, C)\n            if padding:\n                x = x[:, :H, :W, :D].contiguous()\n            x = x.permute(0, 4, 1, 2, 3)\n        return x\n\n\nclass EfficientViTBlock3D(torch.nn.Module):\n    \"\"\" A basic 3D EfficientViT building block.\n\n    Args:\n        type (str): Type for token mixer. Default: 's' for self-attention.\n        ed (int): Number of input channels.\n        kd (int): Dimension for query and key in the token mixer.\n        nh (int): Number of attention heads.\n        ar (int): Multiplier for the query dim for value dimension.\n        resolution (int): Input resolution.\n        window_resolution (int): Local window resolution.\n        kernels (List[int]): The kernel size of the dw conv on query.\n    \"\"\"\n\n    def __init__(self, type,\n                 ed, kd, nh=4,\n                 ar=4,\n                 resolution=14,\n                 window_resolution=7,\n                 kernels=[5, 5, 5, 5], ):\n        super().__init__()\n\n        self.dw0 = Residual(Conv3d_BN(ed, ed, 3, 1, 1, groups=ed, bn_weight_init=0., resolution=resolution))\n        self.ffn0 = Residual(FFN(ed, int(ed * 2), resolution))\n\n        if type == 's':\n            self.mixer = Residual(LocalWindowAttention3D(ed, kd, nh, attn_ratio=ar, \\\n                                                       resolution=resolution, window_resolution=window_resolution,\n                                                       kernels=kernels))\n\n        self.dw1 = Residual(Conv3d_BN(ed, ed, 3, 1, 1, groups=ed, bn_weight_init=0., resolution=resolution))\n        self.ffn1 = Residual(FFN(ed, int(ed * 2), resolution))\n\n    def forward(self, x):\n        return self.ffn1(self.dw1(self.mixer(self.ffn0(self.dw0(x)))))\n\n\nclass PatchMerging3D(nn.Module):\n    \"\"\" 3D Patch Merging Layer\n\n    Args:\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_dim, output_dim, norm_layer=nn.LayerNorm):\n        super().__init__()\n        self.reduction = nn.Linear(8 * input_dim, output_dim, bias=False)\n        self.norm = norm_layer(8 * input_dim)\n\n    def forward(self, x):\n        \"\"\" Forward function.\n\n        Args:\n            x: Input feature, tensor size (B, C, H, W, D).\n        \"\"\"\n        x = x.transpose(1, 4)\n        B, D, H, W, C = x.shape\n\n        # padding\n        pad_input = (H % 2 == 1) or (W % 2 == 1) or (D % 2 == 1)\n        if pad_input:\n            x = F.pad(x, (0, 0, 0, D % 2, 0, W % 2, 0, H % 2))\n\n        x0 = x[:, 0::2, 0::2, 0::2, :]  # B D/2 H/2 W/2 C\n        x1 = x[:, 0::2, 0::2, 1::2, :]  # B D/2 H/2 W/2 C\n        x2 = x[:, 0::2, 1::2, 0::2, :]  # B D/2 H/2 W/2 C\n        x3 = x[:, 0::2, 1::2, 1::2, :]  # B D/2 H/2 W/2 C\n        x4 = x[:, 1::2, 0::2, 0::2, :]  # B D/2 H/2 W/2 C\n        x5 = x[:, 1::2, 0::2, 1::2, :]  # B D/2 H/2 W/2 C\n        x6 = x[:, 1::2, 1::2, 0::2, :]  # B D/2 H/2 W/2 C\n        x7 = x[:, 1::2, 1::2, 1::2, :]  # B D/2 H/2 W/2 C\n        x = torch.cat([x0, x1, x2, x3, x4, x5, x6, x7], -1)  # B D/2 H/2 W/2 8*C\n\n        x = self.norm(x)\n        x = self.reduction(x)\n        x = x.transpose(1, 4)\n        return x\n\n\nclass PatchExpand3D(nn.Module):\n    \"\"\" 3D Patch Merging Layer\n\n    Args:\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_dim, output_dim, norm_layer=nn.LayerNorm):\n        super().__init__()\n        self.expand = nn.Linear(input_dim, output_dim * 8, bias=False)\n        self.norm = norm_layer(output_dim)\n\n    def forward(self, x):\n        \"\"\" Forward function.\n\n        Args:\n            x: Input feature, tensor size (B, C, H, W, D).\n        \"\"\"\n        x = x.transpose(1, 4)\n        B, D, H, W, C = x.shape\n\n        x = self.expand(x)\n        # assert L == D * H * W, \"input feature has wrong size\"\n\n        x = x.view(B, D * 2, H * 2, W * 2, -1)\n        x = self.norm(x)\n        x = x.transpose(1, 4)\n\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-04-25T09:58:31.103514Z","iopub.execute_input":"2024-04-25T09:58:31.103863Z","iopub.status.idle":"2024-04-25T09:58:34.91784Z","shell.execute_reply.started":"2024-04-25T09:58:31.103835Z","shell.execute_reply":"2024-04-25T09:58:34.916807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from __future__ import absolute_import, division, print_function\n\nfrom collections import OrderedDict\n\nimport pytorch_lightning as pl\nimport torch\nimport torch.nn.functional as F\nfrom main_block.capsule_layers import ConvSlimCapsule3D, MarginLoss\n# from main_block.efficientViT3D import EfficientViTBlock3D, PatchMerging3D, PatchExpand3D\nfrom monai.data import decollate_batch\nfrom monai.inferers import sliding_window_inference\nfrom monai.losses import DiceCELoss\nfrom monai.metrics import ConfusionMatrixMetric, DiceMetric, SurfaceDistanceMetric\nfrom monai.networks import one_hot\nfrom monai.networks.blocks import Convolution, UpSample\nfrom monai.networks.layers.factories import Conv\nfrom monai.transforms import AsDiscrete, Compose, EnsureType\nfrom monai.visualize.img2tensorboard import plot_2d_or_3d_image\nfrom torch import nn\n\nimport warnings\n\nwarnings.filterwarnings(\"ignore\")\n\ntorch.set_printoptions(profile=\"full\")\n\n\n# Pytorch Lightning module\nclass EffiViTCaps3D(pl.LightningModule):\n    def __init__(\n            self,\n            in_channels=2,\n            out_channels=4,\n            lr_rate=2e-4,\n            rec_loss_weight=0.1,\n            margin_loss_weight=1.0,\n            class_weight=None,\n            share_weight=False,\n            sw_batch_size=128,\n            cls_loss=\"CE\",\n            val_patch_size=(40, 256, 256),\n            overlap=0.75,\n            connection=\"skip\",\n            val_frequency=100,\n            weight_decay=2e-6,\n            **kwargs,\n    ):\n        super().__init__()\n        self.save_hyperparameters()\n        self.in_channels = self.hparams.in_channels\n        self.out_channels = self.hparams.out_channels\n        self.share_weight = self.hparams.share_weight\n        self.connection = self.hparams.connection\n\n        self.lr_rate = self.hparams.lr_rate\n        self.weight_decay = self.hparams.weight_decay\n\n        self.cls_loss = self.hparams.cls_loss\n        self.margin_loss_weight = self.hparams.margin_loss_weight\n        self.rec_loss_weight = self.hparams.rec_loss_weight\n        self.class_weight = self.hparams.class_weight\n\n        # Defining losses\n        self.classification_loss1 = MarginLoss(class_weight=self.class_weight, margin=0.2)\n\n        if self.cls_loss == \"DiceCE\":\n            self.classification_loss2 = DiceCELoss(softmax=True, to_onehot_y=True, ce_weight=self.class_weight)\n        elif self.cls_loss == \"CE\":\n            self.classification_loss2 = DiceCELoss(\n                softmax=True, to_onehot_y=True, ce_weight=self.class_weight, lambda_dice=0.0\n            )\n        elif self.cls_loss == \"Dice\":\n            self.classification_loss2 = DiceCELoss(softmax=True, to_onehot_y=True, lambda_ce=0.0)\n        self.reconstruction_loss = nn.MSELoss(reduction=\"none\")\n\n        self.val_frequency = self.hparams.val_frequency\n        self.val_patch_size = self.hparams.val_patch_size\n        self.sw_batch_size = self.hparams.sw_batch_size\n        self.overlap = self.hparams.overlap\n        \n        self.final = nn.Linear(256, out_features=8)\n        self.avg_pool = nn.AdaptiveAvgPool3d((1, 1, 1)) \n        self.dropout = nn.Dropout(0.5)\n\n        # Building model\n        self.feature_extractor = nn.Sequential(\n            OrderedDict(\n                [\n                    (\n                        \"conv1\",\n                        Convolution(\n                            spatial_dims=3,\n                            in_channels=self.in_channels,\n                            out_channels=16,\n                            kernel_size=5,\n                            strides=1,\n                            padding=2,\n                            bias=False,\n                        ),\n                    ),\n                    (\n                        \"conv2\",\n                        Convolution(\n                            spatial_dims=3,\n                            in_channels=16,\n                            out_channels=32,\n                            kernel_size=5,\n                            strides=1,\n                            dilation=2,\n                            padding=4,\n                            bias=False,\n                        ),\n                    ),\n                    (\n                        \"conv3\",\n                        Convolution(\n                            spatial_dims=3,\n                            in_channels=32,\n                            out_channels=64,\n                            kernel_size=5,\n                            strides=1,\n                            padding=4,\n                            dilation=2,\n                            bias=False,\n                            act=\"tanh\",\n                        ),\n                    ),\n                ]\n            )\n        )\n\n        self._build_encoder()\n        self._build_decoder()\n        self._build_reconstruct_branch()\n        self._build_3D_EfficientViTBlock()\n\n        # For validation\n#         self.post_pred = Compose([EnsureType(), AsDiscrete(argmax=True, to_onehot=True, n_classes=self.out_channels)])\n#         self.post_label = Compose([EnsureType(), AsDiscrete(to_onehot=num_classes, n_classes=self.out_channels)])\n\n#         self.dice_metric = DiceMetric(include_background=False, reduction=\"mean_batch\", get_not_nans=False)\n#         self.precision_metric = ConfusionMatrixMetric(\n#             include_background=False, metric_name=\"precision\", compute_sample=True, reduction=\"mean_batch\", get_not_nans=False\n#         )\n#         self.recall_metric = ConfusionMatrixMetric(\n#             include_background=False, metric_name=\"recall\", compute_sample=True, reduction=\"mean_batch\", get_not_nans=False\n#         )\n\n#         self.example_input_array = torch.rand(1, self.in_channels, 32, 32, 32)\n\n\n    @staticmethod\n    def add_model_specific_args(parent_parser):\n        parser = parent_parser.add_argument_group(\"EffiViTCaps3D\")\n        # Architecture params\n        parser.add_argument(\"--in_channels\", type=int, default=2)\n        parser.add_argument(\"--out_channels\", type=int, default=4)\n        parser.add_argument(\"--share_weight\", type=int, default=1)\n        parser.add_argument(\"--connection\", type=str, default=\"skip\")\n\n        # Validation params\n        parser.add_argument(\"--val_patch_size\", nargs=\"+\", type=int, default=[32, 32, 32])\n        parser.add_argument(\"--val_frequency\", type=int, default=1000)\n        parser.add_argument(\"--sw_batch_size\", type=int, default=16)\n        parser.add_argument(\"--overlap\", type=float, default=0.75)\n\n        # Loss params\n        parser.add_argument(\"--margin_loss_weight\", type=float, default=0.1)\n        parser.add_argument(\"--rec_loss_weight\", type=float, default=1e-1)\n        parser.add_argument(\"--cls_loss\", type=str, default=\"CE\")\n\n        # Optimizer params\n        parser.add_argument(\"--lr_rate\", type=float, default=1e-4)\n        parser.add_argument(\"--weight_decay\", type=float, default=2e-6)\n        return parent_parser, parser\n\n    def forward(self, x):\n        # Contracting\n        x = self.feature_extractor(x)\n        # fe_0 = self.feature_extractor.conv1(x)\n        # fe_1 = self.feature_extractor.conv2(fe_0)\n        # x = self.feature_extractor.conv3(fe_1)\n\n        conv_1_1 = x\n\n        # conv_2_1 = self.encoder_convs[0](conv_1_1)\n        conv_2_1 = self.patchMergingblock_1(conv_1_1)\n        conv_2_1 = self.relu(conv_2_1)\n        conv_2_1 = self.efficientViT3Dblock_encoder_2(conv_2_1)\n\n        conv_3_1 = self.patchMergingblock_2(conv_2_1)\n        # conv_3_1 = self.encoder_convs[1](conv_2_1)\n        conv_3_1 = self.relu(conv_3_1)\n        conv_3_1 = self.efficientViT3Dblock_encoder_3(conv_3_1)\n        print(conv_3_1.shape, \"After encoder 3 in EffiViTCaps3D\")\n        conv_3_1_reshaped = conv_3_1.view(-1, 8, 16, conv_3_1.shape[2], conv_3_1.shape[3], conv_3_1.shape[4])\n\n        x = self.encoder_conv_caps[0](conv_3_1_reshaped.contiguous())\n        # conv_cap_4_1 = x\n        conv_cap_4_1 = self.encoder_conv_caps[1](x)\n\n        shape = conv_cap_4_1.size()\n        conv_cap_4_1 = conv_cap_4_1.view(shape[0], -1, shape[-3], shape[-2], shape[-1])\n        conv_cap_4_1 = self.efficientViT3Dblock_bottleneck(conv_cap_4_1)\n        print(conv_cap_4_1.shape)\n        # Expanding\n         \n        x = self.avg_pool(conv_cap_4_1)\n        x = self.dropout(x)\n        x = x.flatten(1)\n        print(x.shape)\n        x = self.final(x)\n        \n\n        return x\n\n    def training_step(self, batch, batch_idx):\n        images, labels = batch[\"image\"], batch[\"label\"]\n        # Contracting\n        x = self.feature_extractor(images)\n        # fe_0 = self.feature_extractor.conv1(x)\n        # fe_1 = self.feature_extractor.conv2(fe_0)\n        # x = self.feature_extractor.conv3(fe_1)\n\n        conv_1_1 = x\n\n        # conv_2_1 = self.encoder_convs[0](conv_1_1)\n        conv_2_1 = self.patchMergingblock_1(conv_1_1)\n        conv_2_1 = self.relu(conv_2_1)\n        conv_2_1 = self.efficientViT3Dblock_encoder_2(conv_2_1)\n\n        conv_3_1 = self.patchMergingblock_2(conv_2_1)\n        # conv_3_1 = self.encoder_convs[1](conv_2_1)\n        conv_3_1 = self.relu(conv_3_1)\n        conv_3_1 = self.efficientViT3Dblock_encoder_3(conv_3_1)\n        \n        print(\"After  encoder 3\", conv_3_1.shape)\n        conv_3_1_reshaped = conv_3_1.view(-1, 8, 16, conv_3_1.shape[-1], conv_3_1.shape[-1], conv_3_1.shape[-1])\n\n        x = self.encoder_conv_caps[0](conv_3_1_reshaped.contiguous())\n        # conv_cap_4_1 = x\n        conv_cap_4_1 = self.encoder_conv_caps[1](x)\n\n        shape = conv_cap_4_1.size()\n        conv_cap_4_1 = conv_cap_4_1.view(shape[0], -1, shape[-3], shape[-2], shape[-1])\n        conv_cap_4_1 = self.efficientViT3Dblock_bottleneck(conv_cap_4_1)\n\n        # Downsampled predictions\n        norm = torch.linalg.norm(conv_cap_4_1, dim=2)\n\n        # Expanding\n        if self.connection == \"skip\":\n            # ###########################################################################################\n            # x = self.patchExpandingblock_3(conv_cap_4_1)\n            # x = torch.cat((x, conv_3_1), dim=1)\n            # x = self.efficientViT3Dblock_decoder_3(x)\n            # x = self.relu(x)\n            # x = self.patchExpandingblock_2(x)\n            # x = torch.cat((x, conv_2_1), dim=1)\n            # x = self.efficientViT3Dblock_decoder_2(x)\n            # x = self.relu(x)\n            # x = self.patchExpandingblock_1(x)\n            # x = torch.cat((x, conv_1_1), dim=1)\n            # ###########################################################################################\n            x = self.decoder_conv[0](conv_cap_4_1)\n            x = torch.cat((x, conv_3_1), dim=1)\n            x = self.decoder_conv[1](x)\n            x = self.efficientViT3Dblock_decoder_3(x)\n            x = self.decoder_conv[2](x)\n            x = torch.cat((x, conv_2_1), dim=1)\n            x = self.decoder_conv[3](x)\n            x = self.efficientViT3Dblock_decoder_2(x)\n            x = self.decoder_conv[4](x)\n            x = torch.cat((x, conv_1_1), dim=1)\n\n            # extend decover and skip connection\n            # x = self.add_deconvs[0](x)\n            # x = torch.cat((x, fe_1), dim=1)\n            # x = self.add_deconvs[1](x)\n            # x = torch.cat((x, fe_0), dim=1)\n\n        logits = self.decoder_conv[5](x)\n\n        # Reconstructing\n        reconstructions = self.reconstruct_branch(x)\n\n        # Calculating losses\n        loss, cls_loss, rec_loss = self.losses(images, labels, norm, logits, reconstructions)\n\n        self.log(\"margin_loss\", cls_loss[0], on_step=False, on_epoch=True, sync_dist=True)\n        self.log(f\"{self.cls_loss}_loss\", cls_loss[1], on_step=False, on_epoch=True, sync_dist=True)\n        self.log(\"reconstruction_loss\", rec_loss, on_step=False, on_epoch=True, sync_dist=True)\n\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        images, labels = batch[\"image\"], batch[\"label\"]\n\n        val_outputs = sliding_window_inference(\n            images,\n            roi_size=self.val_patch_size,\n            sw_batch_size=self.sw_batch_size,\n            predictor=self.forward,\n            overlap=self.overlap,\n        )\n\n        val_outputs = [self.post_pred(val_output) for val_output in decollate_batch(val_outputs)]\n        labels = [self.post_label(label) for label in decollate_batch(labels)]\n        self.dice_metric(y_pred=val_outputs, y=labels)\n        self.precision_metric(y_pred=val_outputs, y=labels)\n        self.recall_metric(y_pred=val_outputs, y=labels)\n\n    def validation_epoch_end(self, outputs):\n        dice_scores = self.dice_metric.aggregate()\n        mean_val_dice = torch.mean(dice_scores)\n        print(f\"mean_val_dice:{mean_val_dice}\")\n        self.log(\"val_dice\", mean_val_dice, sync_dist=True)\n        for i, dice_score in enumerate(dice_scores):\n            self.log(f\"val_dice_class {i + 1}\", dice_score, sync_dist=True)\n            print(f\"val_dice_class {i + 1}:{dice_score}\")\n        self.dice_metric.reset()\n\n        precisions = self.precision_metric.aggregate()[0]\n        mean_val_precision = torch.mean(precisions)\n        print(f\"mean_val_precision:{mean_val_precision}\")\n        for i, precision in enumerate(precisions):\n            print(f\"val_precision_class {i + 1}:{precision}\")\n        self.precision_metric.reset()\n\n        recalls = self.recall_metric.aggregate()[0]\n        mean_val_recall = torch.mean(recalls)\n        print(f\"mean_val_recall:{mean_val_recall}\")\n        for i, recall in enumerate(recalls):\n            print(f\"val_recall_class {i + 1}:{recall}\")\n        self.recall_metric.reset()\n\n    def predict_step(self, batch, batch_idx, dataloader_idx=None):\n        images = batch[\"image\"]\n        outputs = sliding_window_inference(\n            images,\n            roi_size=self.val_patch_size,\n            sw_batch_size=self.sw_batch_size,\n            predictor=self.forward,\n            overlap=self.overlap,\n        )\n        return outputs\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=self.lr_rate, weight_decay=self.weight_decay)\n        scheduler = {\n            \"scheduler\": torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, \"max\", factor=0.1, patience=5),\n            \"monitor\": \"val_dice\",\n            \"frequency\": self.val_frequency,\n        }\n        '''\n        optimizer = torch.optim.SGD(self.parameters(), lr=self.lr_rate, weight_decay=self.weight_decay)\n        scheduler = {\n            \"scheduler\": torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=100, T_mult=2),\n            \"monitor\": \"val_dice\",\n            \"frequency\": self.val_frequency\n        }\n        '''\n        return [optimizer], [scheduler]\n\n    def losses(self, volumes, labels, norm, pred, reconstructions):\n        mask = torch.gt(labels, 0)\n        rec_loss = torch.sum(self.reconstruction_loss(volumes * mask, reconstructions * mask), dim=(1, 2, 3, 4)) / (\n                torch.sum(mask, dim=(1, 2, 3, 4)) + 1e-8\n        )\n        rec_loss = torch.mean(rec_loss)\n\n        downsample_labels = F.interpolate(\n            one_hot(labels, self.out_channels), scale_factor=0.125, mode=\"trilinear\", align_corners=False\n        )\n        cls_loss1 = self.classification_loss1(norm, downsample_labels)\n        cls_loss2 = self.classification_loss2(pred, labels)\n\n        return (\n            self.margin_loss_weight * cls_loss1 + cls_loss2 + self.rec_loss_weight * rec_loss,\n            [cls_loss1, cls_loss2],\n            rec_loss,\n        )\n\n    def _build_encoder(self):\n        # self.encoder_convs = nn.ModuleList()\n        # self.encoder_convs.append(\n        #     nn.Conv3d(64, 128, 3, stride=2, padding=1)\n        # )\n        # # self.encoder_convs.append(\n        # #     nn.Conv3d(128, 128, 5, stride=1, padding=2)\n        # # )\n        # self.encoder_convs.append(\n        #     nn.Conv3d(128, 128, 3, stride=2, padding=1)\n        # )\n        # # self.encoder_convs.append(\n        # #     nn.Conv3d(128, 128, 5, stride=1, padding=2)\n        # # )\n        #\n        # for i in range(len(self.encoder_convs)):\n        #     torch.nn.init.normal_(self.encoder_convs[i].weight, std=0.1)\n\n        self.relu = nn.ReLU(inplace=True)\n\n        self.encoder_conv_caps = nn.ModuleList()\n        self.encoder_kernel_size = 3\n\n        self.encoder_conv_caps.append(\n            ConvSlimCapsule3D(\n                kernel_size=self.encoder_kernel_size,\n                input_dim=8,\n                output_dim=8,\n                input_atoms=16,\n                output_atoms=32,\n                stride=2,\n                padding=1,\n                dilation=1,\n                num_routing=3,\n                share_weight=self.share_weight,\n            )\n        )\n\n        self.encoder_conv_caps.append(\n            ConvSlimCapsule3D(\n                kernel_size=self.encoder_kernel_size,\n                input_dim=8,\n                output_dim=self.out_channels,\n                input_atoms=32,\n                output_atoms=64,\n                stride=1,\n                padding=1,\n                dilation=1,\n                num_routing=3,\n                share_weight=self.share_weight,\n            )\n        )\n\n\n    def _build_decoder(self):\n        # self.add_deconvs = nn.ModuleList()\n        # self.add_deconvs.append(\n        #     nn.ConvTranspose3d(128, 32, 1, 1)\n        # )\n        # self.add_deconvs.append(\n        #     nn.ConvTranspose3d(64, 16, 1, 1)\n        # )\n        self.decoder_conv = nn.ModuleList()\n        if self.connection == \"skip\":\n            self.decoder_in_channels = [self.out_channels * 64, 384, 128, 256, 64, 128]\n            self.decoder_out_channels = [256, 128, 128, 64, 64, self.out_channels]\n\n        for i in range(6):\n            if i == 5:\n                self.decoder_conv.append(\n                    Conv[\"conv\", 3](self.decoder_in_channels[i], self.decoder_out_channels[i], kernel_size=1)\n                )\n\n            elif i % 2 == 0:\n                self.decoder_conv.append(\n                    UpSample(\n                        spatial_dims=3,\n                        in_channels=self.decoder_in_channels[i],\n                        out_channels=self.decoder_out_channels[i],\n                        scale_factor=2,\n                    )\n                )\n            else:\n                self.decoder_conv.append(\n                    Convolution(\n                        spatial_dims=3,\n                        kernel_size=3,\n                        in_channels=self.decoder_in_channels[i],\n                        out_channels=self.decoder_out_channels[i],\n                        strides=1,\n                        padding=1,\n                        bias=False,\n                    )\n                )\n\n    def _build_reconstruct_branch(self):\n        self.reconstruct_branch = nn.Sequential(\n            nn.Conv3d(self.decoder_in_channels[-1], 64, 1),\n            # nn.Conv3d(32, 64, 1),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(64, 128, 1),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(128, self.in_channels, 1),\n            nn.Sigmoid(),\n        )\n\n    def _build_3D_EfficientViTBlock(self):\n        # self.efficientViT3Dblock_encoder_1 = EfficientViTBlock3D(type='s', ed=64, kd=16, ar=1, resolution=self.val_patch_size[0])\n        self.efficientViT3Dblock_encoder_2 = EfficientViTBlock3D(type='s', ed=128, kd=16, ar=2, resolution=self.val_patch_size[0]//2)\n        self.efficientViT3Dblock_encoder_3 = EfficientViTBlock3D(type='s', ed=128, kd=16, ar=2, resolution=self.val_patch_size[0]//4)\n\n        # self.efficientViT3Dblock_decoder_1 = EfficientViTBlock3D(type='s', ed=128, kd=16, ar=2, resolution=self.val_patch_size[0])\n#         self.efficientViT3Dblock_decoder_2 = EfficientViTBlock3D(type='s', ed=64, kd=16, ar=1, resolution=self.val_patch_size[0]//2)\n#         self.efficientViT3Dblock_decoder_3 = EfficientViTBlock3D(type='s', ed=128, kd=16, ar=2, resolution=self.val_patch_size[0]//4)\n\n        self.efficientViT3Dblock_bottleneck = EfficientViTBlock3D(type='s', ed=self.out_channels * 64, kd=16,\n                                                                  ar=self.out_channels, resolution=self.val_patch_size[0]//8)\n\n        self.patchMergingblock_1 = PatchMerging3D(input_dim=64, output_dim=128)\n        self.patchMergingblock_2 = PatchMerging3D(input_dim=128, output_dim=128)\n\n        # self.patchExpandingblock_1 = PatchExpand3D(input_dim=256, output_dim=64)\n        # self.patchExpandingblock_2 = PatchExpand3D(input_dim=256, output_dim=128)\n        # self.patchExpandingblock_3 = PatchExpand3D(input_dim=self.out_channels * 64, output_dim=128)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T09:58:34.919953Z","iopub.execute_input":"2024-04-25T09:58:34.920332Z","iopub.status.idle":"2024-04-25T09:59:17.874164Z","shell.execute_reply.started":"2024-04-25T09:58:34.9203Z","shell.execute_reply":"2024-04-25T09:59:17.873173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = EffiViTCaps3D(in_channels=3)\ninput = torch.rand(1,3,40,256,256)\noutput = model(input)\n# print(output.shape)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T10:00:10.155342Z","iopub.execute_input":"2024-04-25T10:00:10.155709Z","iopub.status.idle":"2024-04-25T10:00:41.051593Z","shell.execute_reply.started":"2024-04-25T10:00:10.15568Z","shell.execute_reply":"2024-04-25T10:00:41.050555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\n# sys.path.append('/kaggle/input/rsnazoopublic')","metadata":{"execution":{"iopub.status.busy":"2024-04-25T10:00:41.053806Z","iopub.execute_input":"2024-04-25T10:00:41.054168Z","iopub.status.idle":"2024-04-25T10:00:41.058451Z","shell.execute_reply.started":"2024-04-25T10:00:41.054136Z","shell.execute_reply":"2024-04-25T10:00:41.057521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n@dataclass\nclass BatchSlice:\n    i_from: int\n    i_to: int\n    i_start: int\n\n\ndef get_slices(batch: Tensor, dim=1, window: int = 16, overlap: int = 8) -> List[BatchSlice]:\n    num_imgs = batch.size(dim)\n    if num_imgs <= window:\n        return [BatchSlice(0, num_imgs, 0)]\n    stride = window - overlap\n    result = []\n    current_idx = 0\n    while True:\n        next_idx = current_idx + window\n\n        if next_idx >= num_imgs:\n            current_idx = num_imgs - window\n            offset = overlap // 2 if current_idx > 0 else 0\n            next_idx = num_imgs\n            result.append(BatchSlice(current_idx, next_idx, offset))\n            break\n        else:\n            offset = overlap // 2 if current_idx > 0 else 0\n            result.append(BatchSlice(current_idx, next_idx, offset))\n        current_idx += stride\n    return result","metadata":{"execution":{"iopub.status.busy":"2024-04-25T10:00:41.059679Z","iopub.execute_input":"2024-04-25T10:00:41.060037Z","iopub.status.idle":"2024-04-25T10:00:41.072903Z","shell.execute_reply.started":"2024-04-25T10:00:41.060005Z","shell.execute_reply":"2024-04-25T10:00:41.072033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\ndef _read_labels():\n    labels_df = pd.read_csv(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\n    labels_dict = {}\n    for index, row in labels_df.iterrows():\n        cube_id = row['StudyInstanceUID']\n        overall_patient = row['patient_overall']\n        c1, c2, c3, c4, c5, c6, c7 = row['C1'], row['C2'], row['C3'], row['C4'], row['C5'], row['C6'], row['C7']\n        labels_dict[cube_id] = [overall_patient, c1, c2, c3, c4, c5, c6, c7]\n    return labels_dict\n\nlabels_dict = _read_labels()\n# print(labels_dict)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T10:00:41.07495Z","iopub.execute_input":"2024-04-25T10:00:41.075288Z","iopub.status.idle":"2024-04-25T10:00:41.279645Z","shell.execute_reply.started":"2024-04-25T10:00:41.075239Z","shell.execute_reply":"2024-04-25T10:00:41.278704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def combine_scan(scan_dir: str, size=512, fix_monochrome: bool = True) -> np.ndarray:\n    num_files = len(os.listdir(scan_dir))\n    images = []\n    offset = 0\n    first = None\n    last = None\n    files = []\n    for i in range(num_files):\n        dpath = os.path.join(scan_dir, f\"{i + offset}.dcm\")\n        if i == 0:\n            while not os.path.exists(dpath):\n                offset += 1\n                dpath = os.path.join(scan_dir, f\"{i + offset}.dcm\")\n        files.append(dpath)\n\n    for dpath in files[::2]:\n        ds = pydicom.dcmread(dpath)\n        if not first:\n            first = ds\n        last = ds\n\n        data = ds.pixel_array\n        data = cv2.resize(data, (size, size))\n        if fix_monochrome and ds.PhotometricInterpretation == \"MONOCHROME1\":\n            data = np.amax(data) - data\n        images.append(data)\n\n    if first and last:\n        if last.ImagePositionPatient[2] > first.ImagePositionPatient[2]:\n            images = images[::-1]\n    return np.array(images)\n\nclass DatasetSeg(Dataset):\n    def __init__(\n            self,\n            dataset_dir: str,\n            cases: List[str],\n    ):\n        self.dataset_dir = dataset_dir\n        self.cases = cases\n\n    def __getitem__(self, i):\n        cube_id = self.cases[i]\n        image_cube = combine_scan(os.path.join(self.dataset_dir, cube_id), size=256)\n        image_mean = image_cube.mean()\n        image_std = image_cube.std()\n        h = image_cube.shape[0]\n\n        images = image_cube\n        if h % 32 > 0:\n            tmp = np.zeros(((h // 32 + 1) * 32, 256, 256))\n            tmp[:h] = images\n            images = tmp\n        images = (images - image_mean) / image_std\n        images = np.expand_dims(images, 0)\n        sample = {}\n        sample['image'] = torch.from_numpy(images).float()\n        sample['cube_id'] = cube_id\n        sample['h'] = h\n        return sample\n\n    def __len__(self):\n        return len(self.cases)\n\n\ncrop_augs =  albumentations.ReplayCompose([\n            albumentations.LongestMaxSize(256),\n            albumentations.PadIfNeeded(256, 256, border_mode=cv2.BORDER_CONSTANT),\n        ])\n\nclass DatasetCrops(Dataset):\n    def __init__(\n            self,\n            dataset_dir: str,\n            cases: List[str],\n            transforms=crop_augs,\n            slice_size=40,\n    ):\n        self.dataset_dir = dataset_dir\n        self.transforms = transforms\n        self.slice_size = slice_size\n        self.cases = cases\n\n    def __getitem__(self, i):\n        cube_id = self.cases[i]\n#         mask_cube = tifffile.imread(os.path.join(\"seg_preds\", f\"{cube_id}.tif\"))\n        mask_cube_path = os.path.join(\"/kaggle/input/masks-part1\", f\"{cube_id}.tif\")\n        if os.path.exists(mask_cube_path):\n            mask_cube = tifffile.imread(mask_cube_path)\n        else:\n            # If the file is not found in the first directory, try the second directory\n            mask_cube_path = os.path.join(\"/kaggle/input/masks-part2\", f\"{cube_id}.tif\")\n            if os.path.exists(mask_cube_path):\n                mask_cube = tifffile.imread(mask_cube_path)\n            else:\n                mask_cube_path = os.path.join(\"/kaggle/input/masks-part3\", f\"{cube_id}.tif\")\n                if os.path.exists(mask_cube_path):\n                    mask_cube = tifffile.imread(mask_cube_path)\n                else:\n                    mask_cube_path = os.path.join(\"/kaggle/input/masks-part4\", f\"{cube_id}.tif\")\n                    if os.path.exists(mask_cube_path):\n                        mask_cube = tifffile.imread(mask_cube_path)\n                    else:\n                        mask_cube_path = os.path.join(\"/kaggle/input/masks-part5\", f\"{cube_id}.tif\")\n                        if os.path.exists(mask_cube_path):\n                            mask_cube = tifffile.imread(mask_cube_path)\n                        else:\n                            mask_cube = torch.rand(256, 256, 256)\n                            mask_cube = mask_cube.cpu().numpy().astype(np.int)\n                            print(\"Error: mask_cube.tif not found in both directories.\")\n\n        image_cube = combine_scan(os.path.join(self.dataset_dir, cube_id) ,size=512)\n        boxes = {}\n        for rprop in measure.regionprops(mask_cube):\n            boxes[rprop.label] = rprop.bbox, rprop.area\n\n        image_mean = image_cube.mean()\n        image_std = image_cube.std()\n        slice_size = self.slice_size\n        all_images = []\n#         labels = np.zeros((8,))\n        labels = labels_dict[cube_id]\n        for li in range(1, 8):\n            if li not in boxes:\n                all_images.append(np.zeros((3, self.slice_size, 256, 256)))\n            else:\n                bbox, area = boxes[li]\n                z1, z2 = bbox[0], bbox[3]\n                y1, y2 = max(bbox[1] - 16, 0), min(bbox[4] + 16, 256)\n                x1, x2 = max(bbox[2] - 16, 0), min(bbox[5] + 16, 256)\n                # if z2 - z1 < slice_size:\n                #     z1 = random.randint(max(z2 - slice_size, 0), z1)\n                #     z2 = z1 + slice_size\n                # todo: verify\n                if z2 - z1 < slice_size:\n                    diff = (slice_size - z2 + z1) // 2\n                    z1 = max(0, z1 - diff)\n                    z2 = z1 + slice_size\n                images = image_cube[z1:z2, y1 * 2:y2 * 2, x1 * 2:x2 * 2].copy()\n                masks = mask_cube[z1:z2, y1:y2, x1:x2].copy()\n                slice_size = self.slice_size\n\n                replay = None\n                image_crops = []\n                mask_crops = []\n                for i in range(images.shape[0]):\n                    image = images[i]\n                    mask = masks[i]\n                    h, w, = mask.shape\n                    mask = cv2.resize(mask, (w * 2, h * 2), interpolation=cv2.INTER_NEAREST)\n                    if replay is None:\n                        sample = self.transforms(image=image, mask=mask)\n                        replay = sample[\"replay\"]\n                    else:\n                        sample = ReplayCompose.replay(replay, image=image, mask=mask)\n                    image_ = sample[\"image\"]\n                    image_crops.append(image_)\n                    mask_crops.append(sample[\"mask\"])\n                images = np.array(image_crops).astype(np.float32)\n                masks = np.array(mask_crops).astype(np.float32)\n                images = np.expand_dims(images, -1)\n                masks = np.expand_dims(masks, -1)\n                images = (images - image_mean) / image_std\n\n                images = np.concatenate([images, images, masks], axis=-1)\n                h = images.shape[0]\n                if h > slice_size:\n                    images = images[: slice_size]\n                    all_images.append(np.moveaxis(images, -1, 0))\n                    images = images[-slice_size:]\n                    all_images.append(np.moveaxis(images, -1, 0))\n                else:\n                    if h != slice_size:\n                        tmp = np.zeros((slice_size, *images.shape[1:]))\n                        tmp[:h] = images\n                        images = tmp\n                    all_images.append(np.moveaxis(images, -1, 0))\n\n        sample = {}\n        sample['image'] = torch.from_numpy(np.array(all_images)).float()\n        sample['label'] = torch.from_numpy(np.array(labels)).float()\n        sample['cube_id'] = cube_id\n        return sample\n\n    def __len__(self):\n        return len(self.cases)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T10:00:41.281153Z","iopub.execute_input":"2024-04-25T10:00:41.281521Z","iopub.status.idle":"2024-04-25T10:00:41.317487Z","shell.execute_reply.started":"2024-04-25T10:00:41.28149Z","shell.execute_reply":"2024-04-25T10:00:41.316627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import random_split\nimport numpy as np\nfrom torch.utils.data import Subset \n\nvalidation_split = 0.1\ntest_split = 0.1\n\ndataset_dir = \"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train_images/\"\n\ncases = os.listdir(dataset_dir)\n# Create dataset\ndataset = DatasetCrops(dataset_dir=dataset_dir, cases=cases)\n\n# Compute sizes\ntrain_size = int(len(dataset) * (1 - validation_split - test_split))\nval_test_size = len(dataset) - train_size\nval_size = int(val_test_size / 2)\ntest_size = val_test_size - val_size\n\n# Random split into train, validation, and test datasets\ntrain_dataset, val_test_dataset = random_split(dataset, [train_size, val_test_size])\nval_dataset, test_dataset = random_split(val_test_dataset, [val_size, test_size])\n\n# train_dataset = Subset(train_dataset, range(5))\n# val_dataset = Subset(val_dataset, range(5))\n# test_dataset = Subset(test_dataset, range(5))\n\n# Create data loaders\ntrain_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=1, shuffle=False)\ntest_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\nprint(len(train_loader), len(val_loader), len(test_loader)) \n\n# train_loader = train_loader[:100]","metadata":{"execution":{"iopub.status.busy":"2024-04-25T10:00:57.645609Z","iopub.execute_input":"2024-04-25T10:00:57.645983Z","iopub.status.idle":"2024-04-25T10:00:57.894512Z","shell.execute_reply.started":"2024-04-25T10:00:57.645953Z","shell.execute_reply":"2024-04-25T10:00:57.893564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import csv\n# CSV file path\ncsv_file_path = '/kaggle/working/EfficientNet_data.csv'\n# Column names\nfieldnames = ['epoch','Train_Loss', 'Train_Acc', 'Val_Loss', 'Val_Acc']\n# Writing lists to a CSV file with specific column names\n# csvfile = open(csv_file_path, 'w', newline='')\nwith open(csv_file_path, 'w', newline='') as csvfile:  # Change 'w' to 'a' for append mode\n    # Create a CSV writer object with DictWriter\n    csv_writer = csv.DictWriter(csvfile, fieldnames=fieldnames)\n    csv_writer.writeheader()","metadata":{"execution":{"iopub.status.busy":"2024-04-25T10:00:58.206134Z","iopub.execute_input":"2024-04-25T10:00:58.206796Z","iopub.status.idle":"2024-04-25T10:00:58.212603Z","shell.execute_reply.started":"2024-04-25T10:00:58.206762Z","shell.execute_reply":"2024-04-25T10:00:58.211308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import nn\nimport torch.optim as optim\nfrom sklearn.metrics import accuracy_score\nfrom torch.utils.data import random_split\n\nlearning_rate = 0.001  # Define the learning rate\n\ndef train_model(model, imgs):\n        preds = []\n        with torch.no_grad():\n            for i in range(len(imgs)):\n                with torch.cuda.amp.autocast():\n#                     output = model(imgs[i:i + 1])[\"cls\"][0]\n                    output = model(imgs[i:i + 1])\n                pred_slice = (output.float()).cpu().numpy().astype(np.float32)\n#                 with torch.cuda.amp.autocast():\n#                     output = model(torch.flip(imgs[i:i + 1], dims=(-1,)))\n                pred_slice += (output.float()).cpu().numpy().astype(np.float32)\n                pred_slice /= 2\n                preds.append(pred_slice)\n                break\n        preds = np.max(np.array(preds), axis=0)\n        preds[np.isnan(preds)] = 0.01\n        return preds\n\n\n\ndef train_classification(model, cases: List, num_epochs: int = 1):\n#     dataset = DatasetCrops(dataset_dir=dataset_dir, cases=cases)\n#     train_size = int(len(dataset) * (1 - validation_split))\n#     val_size = len(dataset) - train_size\n#     train_dataset, val_dataset = random_split(dataset, [train_size, val_size])\n#     train_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)\n#     val_loader = DataLoader(val_dataset, batch_size=1, shuffle=False)\n    \n    \n    # Define loss criterion for multi-label classification\n    optimizer = optim.Adam(model.parameters(), lr=learning_rate)  # Define optimizer\n    loss_criterion = nn.BCEWithLogitsLoss()  # You can adjust the loss function as needed\n    best_metric = -1\n    best_metric_epoch = -1\n    best_metrics_epochs_and_time = [[], [], []]\n    total_start=time.time()\n    for epoch in range(start_epoch, num_epochs):\n        print(f\"Epoch [{epoch + 1}/{num_epochs}]\")\n        \n        # Training phase\n        train_loss = 0.0\n        train_preds = []\n        train_labels = []\n        \n        for sample in tqdm(train_loader, desc=\"Training\"):\n            imgs = sample[\"image\"].cuda().float()[0]\n            cube_id = sample[\"cube_id\"][0]\n            with torch.no_grad():\n                preds = []\n#                 for model in models:\n                model.train()  # Set the model to training mode\n                x = train_model(model, imgs)\n                preds.append(x.squeeze())\n                preds = np.average(np.array(preds), axis=0)\n                preds = np.clip(preds, 0.01, 0.99)\n                \n                print(sample['label'].shape)\n                print(torch.tensor(preds).shape)\n                loss = loss_criterion(torch.tensor(preds).squeeze(), sample['label'].squeeze())  # Calculate loss\n                loss.requires_grad = True\n                print(loss)\n                # Backpropagation and optimization (assuming models are trainable)\n#                 for model in models:\n#                     model.train()  # Set the model to training mode\n                optimizer.zero_grad()  # Zero the gradients\n                loss.backward()  # Backpropagate the gradients\n                optimizer.step()  # Update model parameters\n                \n                # Compute metrics\n                train_loss += loss.item() * imgs.size(0)\n                train_preds.extend(preds)\n                train_labels.extend(sample['label'].squeeze().cpu().numpy())\n        \n#         print(train_labels.shape, train_preds.shape)\n        train_loss /= len(train_loader.dataset)\n        train_accuracy = accuracy_score(np.array(train_labels), np.array(train_preds) >= 0.5)\n        print(f\"Train Loss: {train_loss:.4f} | Train Accuracy: {train_accuracy:.4f}\")\n        \n        # Validation phase\n        val_loss = 0.0\n        val_preds = []\n        val_labels = []\n        model.eval()  # Set the model to evaluation mode\n        with torch.no_grad():\n            for val_sample in tqdm(val_loader, desc=\"Validation\"):\n                val_imgs = val_sample[\"image\"].cuda().float()[0]\n                val_preds_batch = []\n#                 for val_model in models:\n                val_x = train_model(model, val_imgs)\n                val_preds_batch.append(val_x.squeeze())\n                val_preds_batch = np.average(np.array(val_preds_batch), axis=0)\n                val_preds_batch = np.clip(val_preds_batch, 0.01, 0.99)\n                val_loss += loss_criterion(torch.tensor(val_preds_batch), val_sample['label'].squeeze()).item() * val_imgs.size(0)\n                val_preds.extend(val_preds_batch)\n                val_labels.extend(val_sample['label'].squeeze().cpu().numpy())\n        \n        val_loss /= len(val_loader.dataset)\n        val_accuracy = accuracy_score(np.array(val_labels), np.array(val_preds) >= 0.5)\n        print(f\"Validation Loss: {val_loss:.4f} | Validation Accuracy: {val_accuracy:.4f}\")\n        with open(csv_file_path, 'a', newline='') as csvfile:\n            csv_writer = csv.DictWriter(csvfile, fieldnames=fieldnames)\n            csv_writer.writerow({\n                'epoch': epoch+1,\n                'Train_Loss': f\"{train_loss:.4f}\",\n                'Train_Acc': f\"{train_accuracy:.4f}\",\n                'Val_Loss': f\"{val_loss:.4f}\",\n                'Val_Acc': f\"{val_accuracy:.4f}\"           \n            })\n        \n        if val_accuracy > best_metric:\n            best_metric = val_accuracy\n            best_metric_epoch = epoch + 1\n            best_metrics_epochs_and_time[0].append(best_metric)\n            best_metrics_epochs_and_time[1].append(best_metric_epoch)\n            best_metrics_epochs_and_time[2].append(time.time() - total_start)\n            torch.save(\n                model.state_dict(),\n                os.path.join(\"/kaggle/working/\", f\"efficientnet_best_metric_model_{epoch+1}.pth\"),\n            )\n            print(\"saved new best metric model\")\n        else:\n            chk_file_name = \"efficientnet_epoch_\" + str(epoch+1) + \"_model.pth\"    \n            torch.save(\n                model.state_dict(),\n                os.path.join(\"/kaggle/working/\", chk_file_name),\n            )\n","metadata":{"execution":{"iopub.status.busy":"2024-04-25T10:00:58.944018Z","iopub.execute_input":"2024-04-25T10:00:58.944454Z","iopub.status.idle":"2024-04-25T10:00:58.968938Z","shell.execute_reply.started":"2024-04-25T10:00:58.944423Z","shell.execute_reply":"2024-04-25T10:00:58.967962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"output = []\n\nstart_epoch = 0\nlatest_checkpoint_path = \"/kaggle/input/efficientvit-dataset/efficientnet_best_metric_model_1.pth\"\nif os.path.exists(latest_checkpoint_path):\n    checkpoint = torch.load(latest_checkpoint_path)\n    model.load_state_dict(checkpoint)\n    start_epoch = 5\n    print(\"checkpoint loaded\")\nelse:\n    print(\"No checkpoint. Starting from scratch\")\n# print(model)\ndevice=torch.device(\"cuda:0\")\nmodel = model.to(device)\n# train_classification(model, cases, 100)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T11:58:19.196536Z","iopub.execute_input":"2024-04-25T11:58:19.196843Z","iopub.status.idle":"2024-04-25T11:58:19.538546Z","shell.execute_reply.started":"2024-04-25T11:58:19.196818Z","shell.execute_reply":"2024-04-25T11:58:19.537237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import List\nimport torch\nimport numpy as np\nfrom sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, roc_auc_score\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\n\ndef predict_classification(models: List[nn.Module]):\n#     test_dataset = DatasetCrops(dataset_dir=test_dataset_dir, cases=cases)\n#     dataloader = DataLoader(\n#         test_dataset, batch_size=1, sampler=None, shuffle=False, num_workers=1, pin_memory=False\n#     )\n    \n    def predict_model(model, imgs):\n        preds = []\n        with torch.no_grad():\n            for i in range(len(imgs)):\n                with torch.cuda.amp.autocast():\n                    output = model(imgs[i:i + 1])\n                pred_slice = torch.sigmoid(output.float()).cpu().numpy().astype(np.float32)\n                with torch.cuda.amp.autocast():\n                    output = model(torch.flip(imgs[i:i + 1], dims=(-1,)))\n                pred_slice += torch.sigmoid(output.float()).cpu().numpy().astype(np.float32)\n                pred_slice /= 2\n                preds.append(pred_slice)\n        preds = np.max(np.array(preds), axis=0)\n        preds[np.isnan(preds)] = 0.01\n        return preds\n        \n    output_list = []  # Initialize the output list\n    all_labels = []\n    all_predictions = []\n    print(\".......................................................................\", len(test_loader))\n    for sample in tqdm(test_loader):\n        imgs = sample[\"image\"].cuda().float()[0]\n        cube_id = sample[\"cube_id\"][0]\n        labels = sample[\"label\"].cpu().numpy()\n        # Option 1: Remove outer list\n        labels = labels[0]\n        print(\"len(labels), labels\", len(labels), labels)\n        all_labels.extend(labels)\n        \n        with torch.no_grad():\n            preds = []\n            \n            preds.append(predict_model(model, imgs))  # Pass imgs to predict_model\n            preds = np.average(np.array(preds), axis=0)\n            preds = np.clip(preds, 0.01, 0.99)\n            print(preds)\n            all_predictions.extend(preds.squeeze())\n#             output_list.append([cube_id, preds])\n    \n    # Calculate metrics\n    \n    accuracy = accuracy_score(all_labels, (np.array(all_predictions) >= 0.5).astype(int))\n    precision = precision_score(all_labels, (np.array(all_predictions) >= 0.5).astype(int))\n    recall = recall_score(all_labels, (np.array(all_predictions) >= 0.5).astype(int))\n    f1 = f1_score(all_labels, (np.array(all_predictions) >= 0.5).astype(int))\n    print(\"accuracy, precision, recall, f1\", accuracy, precision, recall, f1)\n   # auc = roc_auc_score(all_labels, (np.array(all_predictions)))\n    \n    return output_list, accuracy, precision, recall, f1\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"output = []\n# cases = os.listdir(train_dataset_dir)\npredict_classification(model)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import shutil\n# shutil.rmtree('/kaggle/working/seg_preds')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}