{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":10446384,"sourceType":"datasetVersion","datasetId":6466271},{"sourceId":10662541,"sourceType":"datasetVersion","datasetId":6555362},{"sourceId":210375552,"sourceType":"kernelVersion"},{"sourceId":210540318,"sourceType":"kernelVersion"}],"dockerImageVersionId":30823,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This is a compact version of WMCSFB, based in Mednext with Fourier Bassel found here\n\nModified trained version of https://www.kaggle.com/code/perdigao1/wmcsfb-train\nwith added GRN, similar to mednext v2:\n\nhttps://github.com/ZhaoWenzhao/WMCSFB\n\nreference\nhttps://arxiv.org/abs/2402.16825\n\nNew:\n- back to 'S'\n- Adjusted weights\n- Retrain from best","metadata":{}},{"cell_type":"code","source":"!pip install zarr --find-links=/kaggle/input/zarr-installation/zarr-package -q --no-index","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:01:00.735025Z","iopub.execute_input":"2025-02-01T18:01:00.735367Z","iopub.status.idle":"2025-02-01T18:01:07.852453Z","shell.execute_reply.started":"2025-02-01T18:01:00.735330Z","shell.execute_reply":"2025-02-01T18:01:07.851126Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models_pytorch -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:01:26.239758Z","iopub.execute_input":"2025-02-01T18:01:26.240065Z","iopub.status.idle":"2025-02-01T18:01:32.966164Z","shell.execute_reply.started":"2025-02-01T18:01:26.240044Z","shell.execute_reply":"2025-02-01T18:01:32.965340Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install torch-lr-finder -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:01:32.967201Z","iopub.execute_input":"2025-02-01T18:01:32.967413Z","iopub.status.idle":"2025-02-01T18:01:36.305442Z","shell.execute_reply.started":"2025-02-01T18:01:32.967394Z","shell.execute_reply":"2025-02-01T18:01:36.304502Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q /kaggle/input/connected-components-3d-installation/cc3d-package/connected_components_3d-3.21.0-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:01:36.707852Z","iopub.execute_input":"2025-02-01T18:01:36.708179Z","iopub.status.idle":"2025-02-01T18:01:40.143335Z","shell.execute_reply.started":"2025-02-01T18:01:36.708156Z","shell.execute_reply":"2025-02-01T18:01:40.142501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!set","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# WMCSFB_grn compact code\n\nmake sure file `spherical_bessel.npy` is available","metadata":{}},{"cell_type":"code","source":"SPH_BESSEL_FILE_LOCATION = \"/kaggle/input/wmcsfb-spherical-bessel/spherical_bessel.npy\"\n#SPH_BESSEL_FILE_LOCATION = \"./spherical_bessel.npy\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:01:56.696905Z","iopub.execute_input":"2025-02-01T18:01:56.697195Z","iopub.status.idle":"2025-02-01T18:01:56.700896Z","shell.execute_reply.started":"2025-02-01T18:01:56.697176Z","shell.execute_reply":"2025-02-01T18:01:56.700085Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.utils.checkpoint as checkpoint\n\nimport torch.nn.functional as F\n\nfrom typing import Union, List, Tuple\n\nimport scipy.special as special\nimport numpy as np\n\nimport os\n\n\n\"\"\"\nWMCSFB  blocks_fb.py\n\n\"\"\"\n\nclass WMSCSFB_grn_MedNeXtBlock(nn.Module):\n\n    def __init__(self, \n                in_channels:int, \n                out_channels:int, \n                exp_r:int=4, \n                kernel_size:int=7, \n                do_res:int=True,\n                norm_type:str = 'group',\n                n_groups:int or None = None,\n                grn=False\n                ):\n\n        super().__init__()\n\n        self.do_res = do_res\n\n        # First convolution layer with DepthWise Convolutions\n        self.conv1 = Conv3d_sfb(\n            in_channels = in_channels,\n            out_channels = in_channels,\n            kernel_size = kernel_size,\n            stride = 1,\n            padding = kernel_size//2,\n            groups = in_channels if n_groups is None else n_groups,\n            #numbasis = 100,# 64,##24\n        )\n\n        # Normalization Layer. GroupNorm is used by default.\n        if norm_type=='group':\n            self.norm = nn.GroupNorm(\n                num_groups=in_channels, \n                num_channels=in_channels\n                )\n        elif norm_type=='layer':\n            self.norm = LayerNorm(\n                normalized_shape=in_channels, \n                data_format='channels_first'\n                )\n\n        # Second convolution (Expansion) layer with Conv3D 1x1x1\n        self.conv2 = nn.Conv3d(\n            in_channels = in_channels,\n            out_channels = exp_r*in_channels,\n            kernel_size = 1,\n            stride = 1,\n            padding = 0\n        )\n        \n        # GeLU activations\n        self.act = nn.GELU()\n        \n        # Third convolution (Compression) layer with Conv3D 1x1x1\n        self.conv3 = nn.Conv3d(\n            in_channels = exp_r*in_channels,\n            out_channels = out_channels,\n            kernel_size = 1,\n            stride = 1,\n            padding = 0\n        )\n\n        self.grn = grn\n        if grn:\n            # gamma, beta: learnable affine transform parameters\n            self.grn_beta = nn.Parameter(torch.zeros(1,exp_r*in_channels,1,1,1), requires_grad=True)\n            self.grn_gamma = nn.Parameter(torch.zeros(1,exp_r*in_channels,1,1,1), requires_grad=True)\n\n \n    def forward(self, x, dummy_tensor=None):\n        \n        x1 = x\n        x1 = self.conv1(x1)\n        x1 = self.act(self.conv2(self.norm(x1)))\n        \n        if self.grn:\n            # gamma, beta: learnable affine transform parameters\n            # X: input of shape (N,C,H,W,D)\n            gx = torch.norm(x1, p=2, dim=(-3, -2, -1), keepdim=True)\n            nx = gx / (gx.mean(dim=1, keepdim=True)+1e-6)\n            x1 = self.grn_gamma * (x1 * nx) + self.grn_beta + x1\n        \n        x1 = self.conv3(x1)\n        if self.do_res:\n            x1 = x + x1  \n        return x1\n\n\nclass WMSCSFB_MedNeXtDownBlock(WMSCSFB_grn_MedNeXtBlock):\n\n    def __init__(self, in_channels, out_channels, exp_r=4, kernel_size=7, \n                do_res=False, norm_type = 'group'):\n\n        super().__init__(in_channels, out_channels, exp_r, kernel_size, \n                        do_res = False, norm_type = norm_type\n                        )\n\n        self.resample_do_res = do_res\n        if do_res:\n            self.res_conv = nn.Conv3d(\n                in_channels = in_channels,\n                out_channels = out_channels,\n                kernel_size = 1,\n                stride = 2\n            )\n            \n        self.conv1 = nn.Conv3d(\n            in_channels = in_channels,\n            out_channels = in_channels,\n            kernel_size = kernel_size,\n            stride = 2,\n            padding = kernel_size//2,\n            groups = in_channels,\n        )\n\n        \n    def forward(self, x, dummy_tensor=None):\n        \n        x1 = super().forward(x)\n        \n        if self.resample_do_res:\n            res = self.res_conv(x)\n            x1 = x1 + res\n\n        return x1\n\n\nclass WMSCSFB_MedNeXtUpBlock(WMSCSFB_grn_MedNeXtBlock):\n\n    def __init__(self, in_channels, out_channels, exp_r=4, kernel_size=7, \n                do_res=False, norm_type = 'group'):\n        super().__init__(in_channels, out_channels, exp_r, kernel_size,\n                         do_res=False, norm_type = norm_type)\n\n        self.resample_do_res = do_res\n        if do_res:            \n            self.res_conv = nn.ConvTranspose3d(\n                in_channels = in_channels,\n                out_channels = out_channels,\n                kernel_size = 1,\n                stride = 2\n                )\n\n        self.conv1 = nn.ConvTranspose3d(\n            in_channels = in_channels,\n            out_channels = in_channels,\n            kernel_size = kernel_size,\n            stride = 2,\n            padding = kernel_size//2,\n            groups = in_channels,\n        )\n\n\n    def forward(self, x, dummy_tensor=None):\n        \n        x1 = super().forward(x)\n        # Asymmetry but necessary to match shape\n        x1 = torch.nn.functional.pad(x1, (1,0,1,0,1,0))       \n        \n        if self.resample_do_res:\n            res = self.res_conv(x)\n            res = torch.nn.functional.pad(res, (1,0,1,0,1,0))\n            x1 = x1 + res\n\n        return x1\n\n\nclass OutBlock(nn.Module):\n\n    def __init__(self, in_channels, n_classes):\n        super().__init__()\n        self.conv_out = nn.Conv3d(in_channels, n_classes, kernel_size=1)\n    \n    def forward(self, x, dummy_tensor=None): \n        return self.conv_out(x)\n\n\nclass LayerNorm(nn.Module):\n    \"\"\" LayerNorm that supports two data formats: channels_last (default) or channels_first. \n    The ordering of the dimensions in the inputs. channels_last corresponds to inputs with \n    shape (batch_size, height, width, channels) while channels_first corresponds to inputs \n    with shape (batch_size, channels, height, width).\n    \"\"\"\n    def __init__(self, normalized_shape, eps=1e-5, data_format=\"channels_last\"):\n        super().__init__()\n        self.weight = nn.Parameter(torch.ones(normalized_shape))        # beta\n        self.bias = nn.Parameter(torch.zeros(normalized_shape))         # gamma\n        self.eps = eps\n        self.data_format = data_format\n        if self.data_format not in [\"channels_last\", \"channels_first\"]:\n            raise NotImplementedError \n        self.normalized_shape = (normalized_shape, )\n    \n    def forward(self, x, dummy_tensor=False):\n        if self.data_format == \"channels_last\":\n            return F.layer_norm(x, self.normalized_shape, self.weight, self.bias, self.eps)\n        elif self.data_format == \"channels_first\":\n            u = x.mean(1, keepdim=True)\n            s = (x - u).pow(2).mean(1, keepdim=True)\n            x = (x - u) / torch.sqrt(s + self.eps)\n            x = self.weight[:, None, None, None] * x + self.bias[:, None, None, None]\n            return x\n\n\"\"\"\nWMCSFB conv3d_sfb.py\n\"\"\"\n\ndef cartesian_to_polar_coordinates3D(x, y, z):\n    rho = np.sqrt(x**2 + y**2 + z**2)\n    phi_y = np.arctan2(y, x)\n    theta_z = np.arccos(z/rho)\n    theta_z[np.isnan(theta_z)]=1e-12#\n    return phi_y, theta_z, rho\n\ndef Jn(n, r):\n  return special.spherical_jn(n,r)#\n\ndef calculate_FB_3Dbases_shear(L1, alpha, rot_theta, rot_z, shear_xz, shear_yz, shear_xy, shear_zy, shear_yx, shear_zx, sph_bessel):\n    '''\n    s = 2^alpha is the scale\n    alpha <= 0\n    maxK is the maximum num of bases you need\n    '''\n    maxK = (2 * L1 + 1)**3-1#\n\n    L = L1 + 1\n    R = L1 + 0.5\n\n    truncate_freq_factor = 2.5\n\n    if L1 < 2:\n        truncate_freq_factor = 2.5\n\n    xx, yy, zz = np.meshgrid(range(-L, L+1), range(-L, L+1), range(-L, L+1))\n    #xx = xx+shear_xy*yy#################################\n    #zz = zz+shear_zy*yy#################################\n    xx = xx+shear_xz*zz#################################\n    yy = yy+shear_yz*zz#################################\n    \n    xx = xx+shear_xy*yy#################################\n    zz = zz+shear_zy*yy#################################\n    \n    yy = yy+shear_yx*xx#################################\n    zz = zz+shear_zx*xx################################# \n    \n    xx = alpha*xx/(R)\n    yy = alpha*yy/(R)\n    zz = alpha*zz/(R)\n\n    ugrid = np.concatenate([xx.reshape(-1,1), yy.reshape(-1,1), zz.reshape(-1,1)], 1)\n    # angleGrid, lengthGrid\n    tgrid, zgrid, rgrid = cartesian_to_polar_coordinates3D(ugrid[:,0], ugrid[:,1], ugrid[:,2])\n    \n    ########################################\n    tgrid = tgrid+rot_theta ################\n    zgrid = zgrid+rot_z #############\n\n    num_grid_points = ugrid.shape[0]\n\n    maxAngFreq = 15\n\n    B = sph_bessel[(sph_bessel[:,0] <= maxAngFreq) & (sph_bessel[:,3]<= np.pi*R*truncate_freq_factor)]\n    # print(\"B.shape\")\n    # print(B.shape)\n\n    idxB = np.argsort(B[:,2])\n\n    mu_ns = B[idxB, 2]**2\n\n    ang_freqs = B[idxB, 0]\n    rad_freqs = B[idxB, 1]\n    R_ns = B[idxB, 2]\n    \n    z=np.cos(zgrid)#\n\n    num_kq_all = len(ang_freqs)\n    max_ang_freqs = max(ang_freqs)\n\n    Phi_ns = np.zeros((num_grid_points, num_kq_all), np.float32)\n\n    Psi = []\n    kq_Psi = []\n    num_bases = 0\n    Pmn_list = []\n\n    max_ki = max(ang_freqs[:])\n    for i in range(B.shape[0]):\n        ki = np.int64(ang_freqs[i])\n        qi = rad_freqs[i]\n        rkqi = R_ns[i]\n        \n        r0grid = rgrid*R_ns[i]\n        #print(r0grid.shape)\n        F = special.spherical_jn(ki, np.array(r0grid.copy()))#special.jv(ki, r0grid)\n        \n        #Legendre\n        for li in range(np.int64(max_ki+1-ki)):\n            n = li+ki#\n\n            Phi = 1./np.abs(special.spherical_jn(np.int32(ki)+1, R_ns[i]))*F#\n\n            Phi[rgrid >=1] = 0\n\n            Phi_ns[:, i] = Phi\n\n            if ki == 0:#\n                Pmn_z = special.lpmv(ki, n, z)\n                \n                Psi.append(Phi*Pmn_z)\n                kq_Psi.append([ki,qi,rkqi,n])\n                num_bases = num_bases+1\n                \n                Pmn_list.append(Pmn_z)##########\n\n            else:\n                Pmn_z = special.lpmv(ki, n, z)\n                Psi.append(Phi*Pmn_z*np.cos(ki*tgrid)*np.sqrt(2))\n                Psi.append(Phi*Pmn_z*np.sin(ki*tgrid)*np.sqrt(2))\n                kq_Psi.append([ki,qi,rkqi,n])\n                kq_Psi.append([ki,qi,rkqi,n])\n                num_bases = num_bases+2\n                \n                Pmn_list.append(Pmn_z)##########\n                Pmn_list.append(Pmn_z)##########\n                ######################################################\n                        \n    Psi = np.array(Psi)\n    kq_Psi = np.array(kq_Psi)\n    #print(Psi.shape)\n    num_bases = Psi.shape[1]\n\n    Pmn_list = np.array(Pmn_list)##########\n\n    if num_bases > maxK:\n        Psi = Psi[:maxK]\n        kq_Psi = kq_Psi[:maxK]\n        \n        Pmn_list = Pmn_list[:maxK]\n        \n    num_bases = Psi.shape[0]\n    p = Psi.reshape(num_bases, 2*L+1, 2*L+1, 2*L+1).transpose(1,2,3,0)\n    psi = p[1:-1, 1:-1, 1:-1, :]\n    # print(psi.shape)\n    psi = psi.reshape((2*L1+1)**3, num_bases)    \n        \n    # normalize\n    # using the sum of psi_0 to normalize.\n    c = np.sum(psi[:,0])\n    \n    psi = psi/c\n\n    # Add zero frequency basis ################\n    psi = np.concatenate((psi,1.0/(2*L1+1)/(2*L1+1)/(2*L1+1)*np.ones(((2*L1+1)**3, 1))),axis=1).reshape(((2*L1+1)**3, num_bases+1))\n\n    return psi, c, kq_Psi,Pmn_list\n#################################################################################\ndef tensor_fourier_bessel_affine3D(c_o,c_in,size, rot_theta_limit, rot_z_limit, \n    shear_xz_limit, shear_yz_limit, shear_xy_limit, shear_zy_limit, shear_yx_limit, shear_zx_limit, scale_up, scale_low, bessel_folder, num_funcs=None):\n\n    print('sfb 3d applied')\n    # change path based on environment\n    sph_file = SPH_BESSEL_FILE_LOCATION\n    sph_bessel = np.load(sph_file)\n      \n    base_rot_theta = (np.random.rand(c_o,c_in)*2-1)*rot_theta_limit*np.pi\n    base_rot_z = (np.random.rand(c_o,c_in)*2-1)*rot_z_limit*np.pi\n    \n    base_scale = np.random.rand(c_o,c_in)*(scale_up-scale_low)+scale_low\n    base_scale = 2**(base_scale)########################################\n\n    base_shear_xy = (np.random.rand(c_o,c_in)*2-1)*shear_xy_limit\n    base_shear_xy = np.tan(np.pi*base_shear_xy)\n    \n    base_shear_zy = (np.random.rand(c_o,c_in)*2-1)*shear_zy_limit\n    base_shear_zy = np.tan(np.pi*base_shear_zy)\n    \n    base_shear_xz = (np.random.rand(c_o,c_in)*2-1)*shear_xz_limit\n    base_shear_xz = np.tan(np.pi*base_shear_xz)\n    \n    base_shear_yz = (np.random.rand(c_o,c_in)*2-1)*shear_yz_limit\n    base_shear_yz = np.tan(np.pi*base_shear_yz)\n    \n    base_shear_yx = (np.random.rand(c_o,c_in)*2-1)*shear_yx_limit\n    base_shear_yx = np.tan(np.pi*base_shear_yx)\n    \n    base_shear_zx = (np.random.rand(c_o,c_in)*2-1)*shear_zx_limit\n    base_shear_zx = np.tan(np.pi*base_shear_zx)\n    #print(base_rotation)    \n    \n    t_x = np.random.uniform(low=0, high=size, size=(c_o,c_in)).astype(int)\n    t_y = np.random.uniform(low=0, high=size, size=(c_o,c_in)).astype(int)\n    t_z = np.random.uniform(low=0, high=size, size=(c_o,c_in)).astype(int)\n    \n    max_order = size-1\n    \n    num_funcs = num_funcs or size ** 3\n\n    basis_xy = []\n\n    bxy = []\n    \n    for i in range(c_o):\n        for j in range(c_in):\n\n            psi, c, kq_Psi, _ = calculate_FB_3Dbases_shear(size//2, base_scale[i,j], base_rot_theta[i,j], base_rot_z[i,j],\n                            base_shear_xz[i,j], base_shear_yz[i,j],base_shear_xy[i,j], base_shear_zy[i,j],base_shear_yx[i,j], base_shear_zx[i,j],sph_bessel)\n            \n            #print('psi.shape')\n            #print(psi.shape)\n            psi = psi.transpose((1,0))\n            base_n_m = (psi).reshape((-1,size,size,size))\n            \n            ##########################################################################\n            base_n_m = np.roll(base_n_m,(t_x[i,j],t_y[i,j],t_z[i,j]),axis=(1,2,3))\n            ##########################################################################\n            \n            #print('base_n_m',base_n_m.shape)\n            # print(base_n_m.shape)                \n            bxy.append(base_n_m)\n    #print('np.array(bxy).shape',np.array(bxy).shape)\n    basis_xy.extend(bxy)\n\n    basis = torch.Tensor(np.stack(basis_xy))#[:num_funcs]\n    #print(basis.shape)\n    basis = basis.reshape(c_o, c_in, num_funcs, size, size, size).permute((2,0,1,3,4,5)).contiguous()\n\n    return basis,base_scale,base_rot_theta,base_rot_z,base_shear_xz,base_shear_yz,base_shear_xy,base_shear_zy,base_shear_yx,base_shear_zx,t_x,t_y,t_z\n\n\nclass Conv3d_sfb(nn.Module):\n    \"\"\"\n    Convolution with weighted spherical Fourier Bessel filters, 2024, by Wenzhao Zhao\n    \"\"\"\n    '''\n    __constants__ = ['stride', 'padding', 'dilation', 'groups',\n                     'padding_mode', 'output_padding', 'in_channels',\n                     'out_channels', 'kernel_size']\n    '''\n    def __init__(self, \n        in_channels: int,\n        out_channels: int,\n        kernel_size: Union[int, List[int], Tuple[int, ...]],#: int,\n        stride: int = 1,\n        padding: int = 0,\n        dilation: int = 1,\n        groups: int = 1,\n        bias: bool = False,\n        padding_mode: str = 'zeros',\n        numbasis: int = -1,\n        filter_path: str = \"./\",\n        ):#49):#  \n        # TODO: refine this type             \n        #         height, width, patch_x, patch_y, channel,ks=1, stride=1, c_feature=24, c_in=3):#):\n        #height, width, patch_x,patch_y, channel, ks=1, stride=2, c_feature=64, c_in=3):\n        super(Conv3d_sfb, self).__init__()\n\n        if not np.isscalar(kernel_size):\n            self.kernel_size = kernel_size[0]#\n        else:\n            self.kernel_size = kernel_size#\n        self.patch_x = self.kernel_size#\n        self.patch_y = self.kernel_size#\n        self.patch_z = self.kernel_size\n        self.stride = stride\n        self.in_channels = in_channels\n        self.out_channels = out_channels\n        self.padding = padding#[x+1 for x in padding]#\n        self.dilation = dilation\n\n        if numbasis == -1:\n            #if not np.isreal(kernel_size):\n            if not np.isscalar(kernel_size):\n                self.numbasis = (kernel_size[0]-2)**3#\n            else:\n                self.numbasis = (kernel_size-2)**3#\n        else:\n            self.numbasis = numbasis\n        self.groups = groups\n        '''\n        if in_channels<4:\n            self.groups = 1 #in_channels#1#groups\n        else:\n            self.groups = 4 #24 #groups//4\n        '''\n        if bias is False:\n            self.bias = None #False#\n        else:\n            #self.bias = bias\n            self.bias = torch.nn.Parameter(torch.empty(out_channels))\n            torch.nn.init.zeros_(self.bias)\n            \n        self.padding_mode = padding_mode\n\n        save_path = filter_path\n        save_file = save_path+\"basisFB_ic\"+str(in_channels//self.groups)+\"_oc\"+str(self.out_channels)+\"_k\"+str(self.kernel_size)+\"_nb\"+str(self.numbasis)+\".pt\"\n        if not os.path.exists(save_file):\n            wt,base_scale,base_rot_theta,base_rot_z,base_shear_xz,base_shear_yz,base_shear_xy,base_shear_zy,base_shear_yx,base_shear_zx,t_x,t_y,t_z=self.get_str_fb_filter_tensor(self.numbasis,\n                self.in_channels//self.groups,self.out_channels,self.patch_x,self.patch_y,self.patch_z, filter_path)\n            \n            torch.save(wt, save_file)\n        \n        wt = torch.load(save_file, weights_only=False)            \n            \n        self.register_buffer('wt_filter', wt)\n        \n        self.weight = torch.nn.Parameter(torch.Tensor(self.numbasis,self.out_channels,self.in_channels//self.groups))#(c_feature*c_in,3))\n        torch.nn.init.xavier_uniform_(\n        self.weight,\n        gain=torch.nn.init.calculate_gain(\"linear\"))        \n        \n    def forward(self, x):\n        assert len(x.shape) == 5, 'x must been 5 dimensions, but got ' + str(len(x.shape))\n        b,c,h,w,d = x.shape\n        ##b, t, h, w = x.shape\n        filter = torch.einsum('nct,nctijk->nctijk', self.weight, self.wt_filter).contiguous().mean(dim=0).contiguous()\n\n        if isinstance(self.stride, tuple) or isinstance(self.stride, list):\n            stride1 = self.stride[0]\n        elif isinstance(self.stride,int):\n            stride1 = self.stride\n\n        padding1 = self.kernel_size//2\n        \n        result = torch.nn.functional.conv3d(x, filter, \n            stride = stride1, padding = padding1, dilation = self.dilation, groups = self.groups, bias = self.bias)\n\n        return result #\n\n    def extra_repr(self):\n        s = ('{in_channels}, {out_channels}, kernel_size={kernel_size}'\n             ', stride={stride}')\n        if self.padding != 0:#(0,) * len(self.padding):\n            s += ', padding={padding}'\n        if self.dilation != 1:#(1,) * len(self.dilation):\n            s += ', dilation={dilation}'\n        #if self.output_padding != 0:#(0,) * len(self.output_padding):\n        #    s += ', output_padding={output_padding}'\n        if self.groups != 1:\n            s += ', groups={groups}'\n        if self.bias is None:\n            s += ', bias=False'\n        if self.padding_mode != 'zeros':\n            s += ', padding_mode={padding_mode}'\n        if self.numbasis != -1:\n            s += ', numbasis={numbasis}'\n        #if self.affine_grid is True:\n        #    s += ', affine_grid=True'\n        return s.format(**self.__dict__)\n\n    def get_str_fb_filter_tensor(self, numbasis, c_in, channel, patch_x, patch_y, patch_z, bessel_folder, seeds=20):    \n\n        c_o = channel#64\n        c_in = c_in#3\n        size = patch_x#9\n        \n        rot_theta_limit = 1.0\n        rot_z_limit = 0.5\n\n        shear_xz_limit = 0.25#0.4#0.125#0\n        shear_yz_limit = 0.25#0.4#0.125#0\n        shear_xy_limit = 0.25#0.4#0.125#0\n        shear_zy_limit = 0.25#0.4#0.125#0\n        shear_yx_limit = 0.25#0.4#0.125#0\n        shear_zx_limit = 0.25#0.4#0.125#0\n        scale_up = 1.0     # 2.0\n        scale_low = 0.0    # 1.0\n\n        np.random.seed(seeds)\n\n        wt_filter,base_scale,base_rot_theta,base_rot_z,base_shear_xz,base_shear_yz,base_shear_xy,base_shear_zy,base_shear_yx,base_shear_zx,t_x,t_y,t_z = tensor_fourier_bessel_affine3D(c_o,c_in,size, \n            rot_theta_limit, rot_z_limit, shear_xz_limit, shear_yz_limit, shear_xy_limit, shear_zy_limit, shear_yx_limit, shear_zx_limit, scale_up, scale_low, bessel_folder)#, num_funcs=num_funcs)\n\n        wt_filter = wt_filter[0:numbasis,:,:,:,:,:]######################\n        \n        wt_filter = torch.tensor(wt_filter,dtype=torch.float)\n        base_scale = torch.tensor(base_scale)\n        base_rot_theta = torch.tensor(base_rot_theta)\n        base_rot_z = torch.tensor(base_rot_z)\n        base_shear_xz = torch.tensor(base_shear_xz)\n        base_shear_yz = torch.tensor(base_shear_yz)\n        base_shear_xy = torch.tensor(base_shear_xy)\n        base_shear_zy = torch.tensor(base_shear_zy)\n        base_shear_yx = torch.tensor(base_shear_yx)\n        base_shear_zx = torch.tensor(base_shear_zx)\n        t_x = torch.tensor(t_x)\n        t_y = torch.tensor(t_y)\n        t_z = torch.tensor(t_z)     \n\n        return wt_filter,base_scale,base_rot_theta,base_rot_z,base_shear_xz,base_shear_yz,base_shear_xy,base_shear_zy,base_shear_yx,base_shear_zx,t_x,t_y,t_z\n      \n\n\n\n\"\"\"\nWMCSFB MedNextV1.py\n\"\"\"\nclass EmptyDecoder:\n    def __init__(self):\n        self.deep_supervision = False #True\n\n    #def add(self, x):\n    #    self.data.append(x)\n\nclass WMSCSFB_grn_MedNeXt(nn.Module):\n\n    def __init__(self, \n        in_channels: int, \n        n_channels: int,\n        n_classes: int, \n        exp_r: int = 4,                            # Expansion ratio as in Swin Transformers\n        kernel_size: int = 7,                      # Ofcourse can test kernel_size\n        enc_kernel_size: int = None,\n        dec_kernel_size: int = None,\n        deep_supervision: bool = False,             # Can be used to test deep supervision\n        do_res: bool = False,                       # Can be used to individually test residual connection\n        do_res_up_down: bool = False,             # Additional 'res' connection on up and down convs\n        checkpoint_style: bool = None,            # Either inside block or outside block\n        block_counts: list = [2,2,2,2,2,2,2,2,2], # Can be used to test staging ratio: \n                                            # [3,3,9,3] in Swin as opposed to [2,2,2,2,2] in nnUNet\n        norm_type = 'group'\n    ):\n\n        super().__init__()\n\n        self.do_ds = deep_supervision\n        self.decoder = EmptyDecoder() ###############################\n        #self.decoder.deep_supervision = True ###########################################\n        self.do_ds = self.decoder.deep_supervision########################################\n        deep_supervision = self.decoder.deep_supervision#####################################\n        \n        assert checkpoint_style in [None, 'outside_block']\n        self.inside_block_checkpointing = False\n        self.outside_block_checkpointing = False\n        if checkpoint_style == 'outside_block':\n            self.outside_block_checkpointing = True\n\n        if kernel_size is not None:\n            enc_kernel_size = kernel_size\n            dec_kernel_size = kernel_size\n\n        self.stem = nn.Conv3d(in_channels, n_channels, kernel_size=1)\n        if type(exp_r) == int:\n            exp_r = [exp_r for i in range(len(block_counts))]\n        \n        self.enc_block_0 = nn.Sequential(*[\n            WMSCSFB_grn_MedNeXtBlock(\n                in_channels=n_channels,\n                out_channels=n_channels,\n                exp_r=exp_r[0],\n                kernel_size=enc_kernel_size,\n                do_res=do_res,\n                norm_type=norm_type,\n                grn=True\n            ) \n            for i in range(block_counts[0])]\n        ) \n\n        self.down_0 = WMSCSFB_MedNeXtDownBlock(\n            in_channels=n_channels,\n            out_channels=2*n_channels,\n            exp_r=exp_r[1],\n            kernel_size=enc_kernel_size,\n            do_res=do_res_up_down,\n            norm_type=norm_type,\n        )\n    \n        self.enc_block_1 = nn.Sequential(*[\n            WMSCSFB_grn_MedNeXtBlock(\n                in_channels=n_channels*2,\n                out_channels=n_channels*2,\n                exp_r=exp_r[1],\n                kernel_size=enc_kernel_size,\n                do_res=do_res,\n                norm_type=norm_type,\n                grn=True\n            )\n            for i in range(block_counts[1])]\n        )\n\n        self.down_1 = WMSCSFB_MedNeXtDownBlock(\n            in_channels=2*n_channels,\n            out_channels=4*n_channels,\n            exp_r=exp_r[2],\n            kernel_size=enc_kernel_size,\n            do_res=do_res_up_down,\n            norm_type=norm_type,\n        )\n\n        self.enc_block_2 = nn.Sequential(*[\n            WMSCSFB_grn_MedNeXtBlock(\n                in_channels=n_channels*4,\n                out_channels=n_channels*4,\n                exp_r=exp_r[2],\n                kernel_size=enc_kernel_size,\n                do_res=do_res,\n                norm_type=norm_type,\n                grn=True\n            )\n            for i in range(block_counts[2])]\n        )\n\n        self.down_2 = WMSCSFB_MedNeXtDownBlock(\n            in_channels=4*n_channels,\n            out_channels=8*n_channels,\n            exp_r=exp_r[3],\n            kernel_size=enc_kernel_size,\n            do_res=do_res_up_down,\n            norm_type=norm_type,\n        )\n        \n        self.enc_block_3 = nn.Sequential(*[\n            WMSCSFB_grn_MedNeXtBlock(\n                in_channels=n_channels*8,\n                out_channels=n_channels*8,\n                exp_r=exp_r[3],\n                kernel_size=enc_kernel_size,\n                do_res=do_res,\n                norm_type=norm_type,\n                grn=True\n            )            \n            for i in range(block_counts[3])]\n        )\n        \n        self.down_3 = WMSCSFB_MedNeXtDownBlock(\n            in_channels=8*n_channels,\n            out_channels=16*n_channels,\n            exp_r=exp_r[4],\n            kernel_size=enc_kernel_size,\n            do_res=do_res_up_down,\n            norm_type=norm_type,\n        )\n\n        self.bottleneck = nn.Sequential(*[\n            WMSCSFB_grn_MedNeXtBlock(\n                in_channels=n_channels*16,\n                out_channels=n_channels*16,\n                exp_r=exp_r[4],\n                kernel_size=dec_kernel_size,\n                do_res=do_res,\n                norm_type=norm_type,\n                grn=True\n            )\n            for i in range(block_counts[4])]\n        )\n\n        self.up_3 = WMSCSFB_MedNeXtUpBlock(\n            in_channels=16*n_channels,\n            out_channels=8*n_channels,\n            exp_r=exp_r[5],\n            kernel_size=dec_kernel_size,\n            do_res=do_res_up_down,\n            norm_type=norm_type,\n        )\n\n        self.dec_block_3 = nn.Sequential(*[\n            WMSCSFB_grn_MedNeXtBlock(\n                in_channels=n_channels*8,\n                out_channels=n_channels*8,\n                exp_r=exp_r[5],\n                kernel_size=dec_kernel_size,\n                do_res=do_res,\n                norm_type=norm_type,\n                grn=True\n            )\n            for i in range(block_counts[5])]\n        )\n\n        self.up_2 = WMSCSFB_MedNeXtUpBlock(\n            in_channels=8*n_channels,\n            out_channels=4*n_channels,\n            exp_r=exp_r[6],\n            kernel_size=dec_kernel_size,\n            do_res=do_res_up_down,\n            norm_type=norm_type,\n        )\n\n        self.dec_block_2 = nn.Sequential(*[\n            WMSCSFB_grn_MedNeXtBlock(\n                in_channels=n_channels*4,\n                out_channels=n_channels*4,\n                exp_r=exp_r[6],\n                kernel_size=dec_kernel_size,\n                do_res=do_res,\n                norm_type=norm_type,\n                grn=True\n            )\n            for i in range(block_counts[6])]\n        )\n\n        self.up_1 = WMSCSFB_MedNeXtUpBlock(\n            in_channels=4*n_channels,\n            out_channels=2*n_channels,\n            exp_r=exp_r[7],\n            kernel_size=dec_kernel_size,\n            do_res=do_res_up_down,\n            norm_type=norm_type,\n        )\n\n        self.dec_block_1 = nn.Sequential(*[\n            WMSCSFB_grn_MedNeXtBlock(\n                in_channels=n_channels*2,\n                out_channels=n_channels*2,\n                exp_r=exp_r[7],\n                kernel_size=dec_kernel_size,\n                do_res=do_res,\n                norm_type=norm_type,\n                grn=True\n            )\n            for i in range(block_counts[7])]\n        )\n\n        self.up_0 = WMSCSFB_MedNeXtUpBlock(\n            in_channels=2*n_channels,\n            out_channels=n_channels,\n            exp_r=exp_r[8],\n            kernel_size=dec_kernel_size,\n            do_res=do_res_up_down,\n            norm_type=norm_type\n        )\n\n        self.dec_block_0 = nn.Sequential(*[\n            WMSCSFB_grn_MedNeXtBlock(\n                in_channels=n_channels,\n                out_channels=n_channels,\n                exp_r=exp_r[8],\n                kernel_size=dec_kernel_size,\n                do_res=do_res,\n                norm_type=norm_type,\n                grn=True\n            )\n            for i in range(block_counts[8])]\n        )\n\n        self.out_0 = OutBlock(in_channels=n_channels, n_classes=n_classes)\n\n        # Used to fix PyTorch checkpointing bug\n        self.dummy_tensor = nn.Parameter(torch.tensor([1.]), requires_grad=True)  \n\n        if True:# deep_supervision:\n            self.out_1 = OutBlock(in_channels=n_channels*2, n_classes=n_classes)\n            self.out_2 = OutBlock(in_channels=n_channels*4, n_classes=n_classes)\n            self.out_3 = OutBlock(in_channels=n_channels*8, n_classes=n_classes)\n            self.out_4 = OutBlock(in_channels=n_channels*16, n_classes=n_classes)\n\n        self.block_counts = block_counts\n\n\n    def iterative_checkpoint(self, sequential_block, x):\n        \"\"\"\n        This simply forwards x through each block of the sequential_block while\n        using gradient_checkpointing. This implementation is designed to bypass\n        the following issue in PyTorch's gradient checkpointing:\n        https://discuss.pytorch.org/t/checkpoint-with-no-grad-requiring-inputs-problem/19117/9\n        \"\"\"\n        for l in sequential_block:\n            x = checkpoint.checkpoint(l, x, self.dummy_tensor)\n        return x\n\n\n    def forward(self, x):\n        #print(x.shape)\n        x = self.stem(x)\n        if self.outside_block_checkpointing:\n            x_res_0 = self.iterative_checkpoint(self.enc_block_0, x)\n            x = checkpoint.checkpoint(self.down_0, x_res_0, self.dummy_tensor)\n            x_res_1 = self.iterative_checkpoint(self.enc_block_1, x)\n            x = checkpoint.checkpoint(self.down_1, x_res_1, self.dummy_tensor)\n            x_res_2 = self.iterative_checkpoint(self.enc_block_2, x)\n            x = checkpoint.checkpoint(self.down_2, x_res_2, self.dummy_tensor)\n            x_res_3 = self.iterative_checkpoint(self.enc_block_3, x)\n            x = checkpoint.checkpoint(self.down_3, x_res_3, self.dummy_tensor)\n\n            x = self.iterative_checkpoint(self.bottleneck, x)\n            if True:#self.do_ds:\n                x_ds_4 = checkpoint.checkpoint(self.out_4, x, self.dummy_tensor)\n\n            x_up_3 = checkpoint.checkpoint(self.up_3, x, self.dummy_tensor)\n            dec_x = x_res_3 + x_up_3 \n            x = self.iterative_checkpoint(self.dec_block_3, dec_x)\n            if True:#self.do_ds:\n                x_ds_3 = checkpoint.checkpoint(self.out_3, x, self.dummy_tensor)\n            del x_res_3, x_up_3\n\n            x_up_2 = checkpoint.checkpoint(self.up_2, x, self.dummy_tensor)\n            dec_x = x_res_2 + x_up_2 \n            x = self.iterative_checkpoint(self.dec_block_2, dec_x)\n            if True:#self.do_ds:\n                x_ds_2 = checkpoint.checkpoint(self.out_2, x, self.dummy_tensor)\n            del x_res_2, x_up_2\n\n            x_up_1 = checkpoint.checkpoint(self.up_1, x, self.dummy_tensor)\n            dec_x = x_res_1 + x_up_1 \n            x = self.iterative_checkpoint(self.dec_block_1, dec_x)\n            if True:#self.do_ds:\n                x_ds_1 = checkpoint.checkpoint(self.out_1, x, self.dummy_tensor)\n            del x_res_1, x_up_1\n\n            x_up_0 = checkpoint.checkpoint(self.up_0, x, self.dummy_tensor)\n            dec_x = x_res_0 + x_up_0 \n            x = self.iterative_checkpoint(self.dec_block_0, dec_x)\n            del x_res_0, x_up_0, dec_x\n\n            x = checkpoint.checkpoint(self.out_0, x, self.dummy_tensor)\n\n        else:\n            x_res_0 = self.enc_block_0(x)\n            x = self.down_0(x_res_0)\n            x_res_1 = self.enc_block_1(x)\n            x = self.down_1(x_res_1)\n            x_res_2 = self.enc_block_2(x)\n            x = self.down_2(x_res_2)\n            x_res_3 = self.enc_block_3(x)\n            x = self.down_3(x_res_3)\n\n            x = self.bottleneck(x)\n            if True:#self.do_ds:\n                x_ds_4 = self.out_4(x)\n\n            x_up_3 = self.up_3(x)\n            dec_x = x_res_3 + x_up_3 \n            x = self.dec_block_3(dec_x)\n\n            if True:# self.do_ds:\n                x_ds_3 = self.out_3(x)\n            del x_res_3, x_up_3\n\n            x_up_2 = self.up_2(x)\n            dec_x = x_res_2 + x_up_2 \n            x = self.dec_block_2(dec_x)\n            if True:# self.do_ds:\n                x_ds_2 = self.out_2(x)\n            del x_res_2, x_up_2\n\n            x_up_1 = self.up_1(x)\n            dec_x = x_res_1 + x_up_1 \n            x = self.dec_block_1(dec_x)\n            if True:# self.do_ds:\n                x_ds_1 = self.out_1(x)\n            del x_res_1, x_up_1\n\n            x_up_0 = self.up_0(x)\n            dec_x = x_res_0 + x_up_0 \n            x = self.dec_block_0(dec_x)\n            del x_res_0, x_up_0, dec_x\n\n            x = self.out_0(x)\n\n\n        if self.decoder.deep_supervision:#self.do_ds:\n            return [x, x_ds_1, x_ds_2, x_ds_3, x_ds_4]\n        else: \n            return x\n\n\n\"\"\"\nWMCSFB create_mednext_v1.py\n\n\"\"\"\n\ndef create_WMSCSFB_grn_mednextv1_small(num_input_channels, num_classes, kernel_size=3, ds=False):\n\n    return WMSCSFB_grn_MedNeXt(\n        in_channels = num_input_channels, \n        n_channels = 32,\n        n_classes = num_classes, \n        exp_r=2,                         \n        kernel_size=kernel_size,         \n        deep_supervision=ds,             \n        do_res=True,                     \n        do_res_up_down = True,\n        block_counts = [2,2,2,2,2,2,2,2,2]\n    )\n\n\n\ndef create_WMSCSFB_grn_mednextv1_base(num_input_channels, num_classes, kernel_size=3, ds=False):\n\n    return WMSCSFB_grn_MedNeXt(\n        in_channels = num_input_channels, \n        n_channels = 32,\n        n_classes = num_classes, \n        exp_r=[2,3,4,4,4,4,4,3,2],       \n        kernel_size=kernel_size,         \n        deep_supervision=ds,             \n        do_res=True,                     \n        do_res_up_down = True,\n        block_counts = [2,2,2,2,2,2,2,2,2]\n    )\n\n\ndef create_WMSCSFB_grn_mednextv1_medium(num_input_channels, num_classes, kernel_size=3, ds=False):\n\n    return WMSCSFB_grn_MedNeXt(\n        in_channels = num_input_channels, \n        n_channels = 32,\n        n_classes = num_classes, \n        exp_r=[2,3,4,4,4,4,4,3,2],       \n        kernel_size=kernel_size,         \n        deep_supervision=ds,             \n        do_res=True,                     \n        do_res_up_down = True,\n        block_counts = [3,4,4,4,4,4,4,4,3],\n        checkpoint_style = 'outside_block'\n    )\n\n\ndef create_WMSCSFB_grn_mednextv1_large(num_input_channels, num_classes, kernel_size=3, ds=False):\n\n    return WMSCSFB_grn_MedNeXt(\n        in_channels = num_input_channels, \n        n_channels = 32,\n        n_classes = num_classes, \n        exp_r=[3,4,8,8,8,8,8,4,3],                          \n        kernel_size=kernel_size,                     \n        deep_supervision=ds,             \n        do_res=True,                     \n        do_res_up_down = True,\n        block_counts = [3,4,8,8,8,8,8,4,3],\n        #checkpoint_style = 'outside_block'\n    )\n\n\ndef create_WMSCSFB_grn_mednext_v1(num_input_channels, num_classes, model_id, kernel_size=3,\n                      deep_supervision=False):\n\n    model_dict = {\n        'S': create_WMSCSFB_grn_mednextv1_small,\n        'B': create_WMSCSFB_grn_mednextv1_base,\n        'M': create_WMSCSFB_grn_mednextv1_medium,\n        'L': create_WMSCSFB_grn_mednextv1_large,\n        }\n    \n    return model_dict[model_id](\n        num_input_channels, num_classes, kernel_size, deep_supervision\n        )","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:01:57.215553Z","iopub.execute_input":"2025-02-01T18:01:57.215895Z","iopub.status.idle":"2025-02-01T18:02:00.610308Z","shell.execute_reply.started":"2025-02-01T18:01:57.215867Z","shell.execute_reply":"2025-02-01T18:02:00.609559Z"},"trusted":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import numpy as np \nimport json\nimport matplotlib.pyplot as plt\nimport zarr\nfrom tqdm import tqdm\nimport os\nimport sys\nimport datetime\n\nfrom pathlib import Path\n\nimport torch\nimport torch.nn\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom sklearn.model_selection import StratifiedKFold\nimport segmentation_models_pytorch as smp\nimport random\nimport time\nimport math\n\nimport torch\n\ntorch_device = 'cuda' if torch.cuda.is_available() else 'cpu'\n\nfrom torch_lr_finder import LRFinder\n\nimport cc3d\nimport pandas\n#from dataclasses import dataclass\n\nfrom types import SimpleNamespace # for CFG","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:00.611107Z","iopub.execute_input":"2025-02-01T18:02:00.611455Z","iopub.status.idle":"2025-02-01T18:02:06.136088Z","shell.execute_reply.started":"2025-02-01T18:02:00.611431Z","shell.execute_reply":"2025-02-01T18:02:06.135158Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEBUG=False\nif os.environ.get('KAGGLE_KERNEL_RUN_TYPE','') == 'Interactive':\n    print(\"Kaggle in interactive mode. Setting to DEBUG\")\n    DEBUG = True","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:06.137239Z","iopub.execute_input":"2025-02-01T18:02:06.137823Z","iopub.status.idle":"2025-02-01T18:02:06.142182Z","shell.execute_reply.started":"2025-02-01T18:02:06.137787Z","shell.execute_reply":"2025-02-01T18:02:06.141388Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import logging\n# logging.basicConfig(level=logging.INFO,\n#                     format=\"%(asctime)s — %(name)s — %(levelname)s — %(funcName)s:%(lineno)d — %(message)s\",\n#         )","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:09.045535Z","iopub.execute_input":"2025-02-01T18:02:09.045853Z","iopub.status.idle":"2025-02-01T18:02:09.049320Z","shell.execute_reply.started":"2025-02-01T18:02:09.045830Z","shell.execute_reply":"2025-02-01T18:02:09.048427Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"config","metadata":{}},{"cell_type":"code","source":"#class CFG:\nCFG = SimpleNamespace(\n    # all ids: ['TS_86_3', 'TS_6_6', 'TS_6_4', 'TS_5_4', 'TS_73_6', 'TS_99_9', 'TS_69_2']\n\n    #IDS_TO_TRAIN = ['TS_86_3', 'TS_6_6', 'TS_73_6', 'TS_99_9', 'TS_6_4'],\n    IDS_TO_TRAIN = ['TS_86_3', 'TS_6_6', 'TS_6_4', 'TS_5_4', 'TS_73_6', 'TS_99_9', 'TS_69_2'], #all volumes\n    IDS_TO_TEST = ['TS_5_4', 'TS_69_2'],\n\n    #DATA_NORM_FUNC = \"stdev3_with_clip\" # stdev_no_clip, stdev3_no_clip, stdev3_with_clip\n    DATA_NORM_FUNC = \"stdev_no_clip\",\n\n    PARTICLES = ['ribosome', 'virus-like-particle', 'apo-ferritin', 'beta-galactosidase', 'thyroglobulin'],\n\n    PARTICLE_RADIUS = {\n        'ribosome': 150.0,\n        'virus-like-particle': 135.0,\n        'apo-ferritin': 60.0,\n        'beta-galactosidase': 90.0,\n        'thyroglobulin': 130.0,\n    },\n\n    #CLASS_WEIGHTS = [ 0.5, 1.0, 1.0, 1.0, 2.0, 2.0 ], # Use same order as PARTICLES\n    CLASS_WEIGHTS = [ 0.2, 1.0, 1.0, 1.0, 3.0, 4.0 ], # balanced\n    \n    GND_MASK_RADIUS_SCALE_FACTOR = 0.3, # First training 3/4, later reduce to 1/3\n    \n    TRAIN_AUGM_ROTS = True,\n    TRAIN_AUGM_FLIPS = True,\n    TRAIN_AUGM_NEIGHB_VOXEL_SWAP_PROB= 0.0, #0.1\n    TRAIN_AUGM_GAUSS_NOISE_STD = 0.10,\n    TRAIN_RANDOM_VOXEL_SHIFT_PROB = 0.10,\n\n    WMCSFB_V1_INP_CHN = 1,\n    WMCSFB_V1_NUM_CLASSES = 6, # including background\n    WMCSFBL_V1_ID = 'S', # S, B, M, L\n    WMCSFB_V1_KERNEL_SIZE = 3,\n    WMCSFB_GRN = True,\n    \n    #################\n    WMCSFB_USE_PRETRAINED_MODEL = \"/kaggle/input/notebook-output-wmcsfb-grn-train/2025-02-04-0141_wmcsfb_grn_step1.ptd\",\n    #################\n\n    \n    TRAIN_CUBE_SIZE = 64, # note that the Mednext only accepts cube sizes 64, 80, 96, 112, 128. for training 80 would require too much VRAM\n    TRAIN_BATCH_SIZE= 5,\n    \n    TRAIN_EPOCHS = 4,\n    TRAIN_K_FOLDS_PER_EPOCH = 4,\n    TRAIN_LR = 5.0e-6,\n    TRAIN_MAX_LR = 5.0e-5,\n\n    # second part of the training, using stepLR\n    TRAIN2_EPOCHS = 0, #don't train using stepLR, not working\n    TRAIN2_INIT_LR = 1e-5,\n    TRAIN2_GAMMA = 0.31622, # 1/sqrt(10) every 2 steps to reduce by 10\n\n    # Train with several steps=epoch of cosine annealing\n    TRAIN3_COSWREST_EPOCHS = 4, # Note that each epoch runs a cosine cycle. Each epoch uses mini epochs, using sfk\n    TRAIN3_COSWREST_INIT_LR = 3.0e-5,\n    \n    # This is training by excluding the worst training discriminate% in metrics\n    TRAIN4_EPOCHS = 0,\n    TRAIN4_DISCRIMINATE_PERC = 3,\n    TRAIN4_K_FOLDS_PER_EPOCH = 4,\n    TRAIN4_LR = 5.0e-6,\n    TRAIN4_MAX_LR = 5.0e-5,\n\n    #CC3D_CENTROID_SHIFT_PIX = -0.5,\n    CC3D_CENTROID_SHIFT_PIX = 0.0,\n\n    # DATALOADER_NUM_WORKERS = 3 # cannot use CUDA with multiprocess\n)","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:09.557515Z","iopub.execute_input":"2025-02-01T18:02:09.557856Z","iopub.status.idle":"2025-02-01T18:02:09.564118Z","shell.execute_reply.started":"2025-02-01T18:02:09.557830Z","shell.execute_reply":"2025-02-01T18:02:09.563117Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if DEBUG:\n    print(\"Debug mode\")\n    CFG.IDS_TO_TRAIN = ['TS_86_3'] # Debug\n    CFG.IDS_TO_TEST = ['TS_69_2']\n    \n    CFG.TRAIN_EPOCHS = 1\n    CFG.TRAIN2_EPOCHS = 0\n    CFG.TRAIN3_COSWREST_EPOCHS=1\n    CFG.TRAIN4_EPOCHS = 1","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:12.203895Z","iopub.execute_input":"2025-02-01T18:02:12.204181Z","iopub.status.idle":"2025-02-01T18:02:12.209037Z","shell.execute_reply.started":"2025-02-01T18:02:12.204160Z","shell.execute_reply":"2025-02-01T18:02:12.208077Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAIN_BASE = '/kaggle/input/czii-cryo-et-object-identification/train/static/ExperimentRuns'\nTEST_BASE =  '/kaggle/input/czii-cryo-et-object-identification/test/static/ExperimentRuns'\nOVERLAY_BASE = '/kaggle/input/czii-cryo-et-object-identification/train/overlay/ExperimentRuns'\n\nif os.environ.get('COMPUTERNAME','') != '':\n    print(\"notebook runnning in local PC\")\n    TRAIN_BASE = '../../data/train/static/ExperimentRuns'\n    TEST_BASE =  '../../data/test/static/ExperimentRuns'\n    OVERLAY_BASE = '../../data/train/overlay/ExperimentRuns'","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:12.396939Z","iopub.execute_input":"2025-02-01T18:02:12.397151Z","iopub.status.idle":"2025-02-01T18:02:12.400889Z","shell.execute_reply.started":"2025-02-01T18:02:12.397134Z","shell.execute_reply":"2025-02-01T18:02:12.400134Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"exp_folder_path = Path(TRAIN_BASE)\nexp_runs = [ d.parts[-1] for d in exp_folder_path.iterdir()]\nexp_runs","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:14.398880Z","iopub.execute_input":"2025-02-01T18:02:14.399177Z","iopub.status.idle":"2025-02-01T18:02:14.409762Z","shell.execute_reply.started":"2025-02-01T18:02:14.399155Z","shell.execute_reply":"2025-02-01T18:02:14.409082Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_data_by_id(id):\n    \n    exp_webpzarr = Path(TRAIN_BASE)/id/ \"VoxelSpacing10.000/denoised.zarr\"\n    print(f\"{str(exp_webpzarr)} exists:{exp_webpzarr.exists()}\")\n    # zstore_ts = zarr.open(str(exp_webpzarr), mode='r')\n    zstore_ts = zarr.open(str(exp_webpzarr), mode='r')\n    data0 = zstore_ts[0][:]\n    print(f\"id:{id} , data.shape:{data0.shape}\")\n    \n    return data0","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:14.612929Z","iopub.execute_input":"2025-02-01T18:02:14.613150Z","iopub.status.idle":"2025-02-01T18:02:14.617251Z","shell.execute_reply.started":"2025-02-01T18:02:14.613131Z","shell.execute_reply":"2025-02-01T18:02:14.616353Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Particle definitions","metadata":{}},{"cell_type":"code","source":"def get_dict_labels_by_id(exp_id, scaling_factor=10):\n    labels = {}\n    for particle_type in CFG.PARTICLES:\n    \n        ### Load file with labels\n        #raw_data = open_labels_file(experiment, particle_type)\n        #path = get_label_path(experiment, particle_type)\n        path = Path(OVERLAY_BASE) / exp_id / \"Picks\"/ f\"{particle_type}.json\"\n        with open(path, 'r') as f:\n            raw_data = json.load(f)\n\n        labels[particle_type] = []\n        for item in raw_data['points']:\n            labels[particle_type].append(\n                (\n                    item['location']['z'] / scaling_factor,\n                    item['location']['y'] / scaling_factor,\n                    item['location']['x'] / scaling_factor\n                )\n            )\n    return labels","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:15.465433Z","iopub.execute_input":"2025-02-01T18:02:15.465796Z","iopub.status.idle":"2025-02-01T18:02:15.470784Z","shell.execute_reply.started":"2025-02-01T18:02:15.465765Z","shell.execute_reply":"2025-02-01T18:02:15.469920Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## mask of a single particle for a given exp_id","metadata":{}},{"cell_type":"code","source":"PARTICLE_TO_KERNEL={}\n\nfor p0, r0 in CFG.PARTICLE_RADIUS.items():\n    \n    # radius_scaled = r0/10.0\n    radius_scaled = r0/10.0 * CFG.GND_MASK_RADIUS_SCALE_FACTOR\n    \n    k_size_2 = math.ceil(radius_scaled*1.1)\n    k_size = int(k_size_2*2)+1\n    \n    # gconv_kernel = np.zeros( (k_size,k_size,k_size), np.float16)\n    z = np.arange(k_size)\n    y = np.arange(k_size)\n    x = np.arange(k_size)\n    Z, Y, X = np.meshgrid(z, y, x)\n    \n    #distance = np.sqrt((X-k_size/2)**2 + (Y-k_size/2)**2 + (Z-k_size/2)**2)\n    #distance = np.sqrt((X-k_size/2+1)**2 + (Y-k_size/2+1)**2 + (Z-k_size/2+1)**2) #correction by one pixel\n    distance = np.sqrt((X-k_size_2)**2 + (Y-k_size_2)**2 + (Z-k_size_2)**2)\n\n    # Fill the circle with class number\n    #mask_bool= (distance <= radius_scaled) #mask as True and False\n    mask_bool= (distance < radius_scaled) #mask as True and False\n    \n    PARTICLE_TO_KERNEL[p0] = mask_bool","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:17.798292Z","iopub.execute_input":"2025-02-01T18:02:17.798584Z","iopub.status.idle":"2025-02-01T18:02:17.805727Z","shell.execute_reply.started":"2025-02-01T18:02:17.798560Z","shell.execute_reply":"2025-02-01T18:02:17.804942Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(7,20))\nnplots = len(PARTICLE_TO_KERNEL)\n#print(nplots)\nfor i,(p0,k0) in enumerate(PARTICLE_TO_KERNEL.items()):\n    print(f\"{k0.shape=}\")\n    plt.subplot(1,nplots,i+1)\n    z0 = k0.shape[0]//2\n    plt.imshow(k0[z0,:,:])\n    plt.axis('off')\n    plt.title(f\"{p0} z={z0}\", size=\"x-small\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:19.640796Z","iopub.execute_input":"2025-02-01T18:02:19.641080Z","iopub.status.idle":"2025-02-01T18:02:19.959835Z","shell.execute_reply.started":"2025-02-01T18:02:19.641059Z","shell.execute_reply":"2025-02-01T18:02:19.958977Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from scipy.signal import convolve","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:19.961131Z","iopub.execute_input":"2025-02-01T18:02:19.961445Z","iopub.status.idle":"2025-02-01T18:02:20.011039Z","shell.execute_reply.started":"2025-02-01T18:02:19.961411Z","shell.execute_reply":"2025-02-01T18:02:20.010394Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PARTICLE_TO_CLASS={ part0:i+1 for i,part0 in enumerate(CFG.PARTICLES)}\nPARTICLE_TO_CLASS","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:22.529047Z","iopub.execute_input":"2025-02-01T18:02:22.529683Z","iopub.status.idle":"2025-02-01T18:02:22.535102Z","shell.execute_reply.started":"2025-02-01T18:02:22.529654Z","shell.execute_reply":"2025-02-01T18:02:22.534220Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_data_and_sphere_masks(exp_id):\n    \"\"\"\n    For the given ID returns:\n        - Data volume\n        - Mask as an int8, with each particle marked as a sphere centered in the\n        particle with corresponding int value from PARTICLE_TO_CLASS and radius\n        given in PARTICLE_RADIUS\n        - A list of particles processed ewith each item being (expid, particle_class, (z,y,x) )\n        with zyx in pixel coordinates\n\n    \"\"\"\n\n    # Generate 3D kernel mask for each particle\n\n    dataset = get_data_by_id(exp_id)\n\n    print(f\"{exp_id=}, {dataset.shape=}\")\n\n    ### Init mask with zeros\n    mask = np.zeros(dataset.shape, dtype=np.uint8)\n    all_particles = get_dict_labels_by_id(exp_id)\n\n    used_particles_expid_cls_poszyx = []\n    \n    for particle_type,pcls in tqdm(PARTICLE_TO_CLASS.items()):\n        \n        fmask_dirac = np.zeros((dataset.shape), dtype=np.uint8)\n\n        this_particle_poss_zyx = all_particles[particle_type]\n        \n        for cz, cy, cx in this_particle_poss_zyx:\n            fmask_dirac[round(cz),round(cy),round(cx)]=1\n            used_particles_expid_cls_poszyx.append( (exp_id, pcls, (cz, cy, cx) ) )\n\n        #convolve for faster mask creation\n        fmask_same = convolve(fmask_dirac, PARTICLE_TO_KERNEL[particle_type], mode='same' )\n    \n        mask[ fmask_same>0 ] = (fmask_same[ fmask_same>0 ]>0)*pcls\n\n    return dataset, mask, used_particles_expid_cls_poszyx\n","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:22.706540Z","iopub.execute_input":"2025-02-01T18:02:22.706814Z","iopub.status.idle":"2025-02-01T18:02:22.712467Z","shell.execute_reply.started":"2025-02-01T18:02:22.706791Z","shell.execute_reply":"2025-02-01T18:02:22.711672Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# test\nid0 = CFG.IDS_TO_TRAIN[0]\nprint(f\"{id0=}\")\ndata0, mask0, up0 = get_data_and_sphere_masks(id0)","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:25.361115Z","iopub.execute_input":"2025-02-01T18:02:25.361434Z","iopub.status.idle":"2025-02-01T18:02:56.365391Z","shell.execute_reply.started":"2025-02-01T18:02:25.361410Z","shell.execute_reply":"2025-02-01T18:02:56.364666Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"{data0.shape=} , {mask0.shape=}\")","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:56.366288Z","iopub.execute_input":"2025-02-01T18:02:56.366509Z","iopub.status.idle":"2025-02-01T18:02:56.370901Z","shell.execute_reply.started":"2025-02-01T18:02:56.366489Z","shell.execute_reply":"2025-02-01T18:02:56.369912Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"zslice = 50\nplt.figure(figsize=(6,6))\nplt.imshow(data0[zslice,:,:], cmap='gray')\nplt.imshow(mask0[zslice,:,:], cmap='tab10', interpolation='nearest', alpha=0.5)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:56.372911Z","iopub.execute_input":"2025-02-01T18:02:56.373157Z","iopub.status.idle":"2025-02-01T18:02:56.731974Z","shell.execute_reply.started":"2025-02-01T18:02:56.373137Z","shell.execute_reply":"2025-02-01T18:02:56.731149Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# zslice = 50\n# plt.figure(figsize=(7,7))\n# plt.imshow(data0[zslice,400:550,400:500], cmap='gray')\n# plt.imshow(mask0[zslice,400:550,400:500], cmap='tab10', interpolation='nearest', alpha=0.7)\n# plt.axis('off')\n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:56.733200Z","iopub.execute_input":"2025-02-01T18:02:56.733431Z","iopub.status.idle":"2025-02-01T18:02:56.736645Z","shell.execute_reply.started":"2025-02-01T18:02:56.733410Z","shell.execute_reply":"2025-02-01T18:02:56.735790Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#up0","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:56.737590Z","iopub.execute_input":"2025-02-01T18:02:56.737902Z","iopub.status.idle":"2025-02-01T18:02:56.750506Z","shell.execute_reply.started":"2025-02-01T18:02:56.737873Z","shell.execute_reply":"2025-02-01T18:02:56.749915Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Good","metadata":{}},{"cell_type":"code","source":"def normalise_voldata_to_stdev(datavol):\n    # Result is returned as float\n    d0_mean = np.mean(datavol)\n    d0_std = np.std(datavol)\n\n    if d0_std==0:\n        raise ValueError(\"Error. Stdev of data volume is zero.\")\n    \n    d0_corr = (datavol.astype(np.float32) - d0_mean) / d0_std\n    #d0_corr = (np.clip(d0_corr, -5.0, 3.0) + 5.0) / 8.0\n    \n    return d0_corr\n\ndef normalise_voldata_to_stdev_3(datavol):\n    # Result is returned as float\n    \n    d0_mean = np.mean(datavol)\n    d0_std = np.std(datavol)\n\n    if d0_std==0:\n        raise ValueError(\"Error. Stdev of data volume is zero.\")\n    \n    d0_corr = (datavol.astype(np.float32) - d0_mean) / d0_std\n    d0_corr = (d0_corr+3.0) / 6.0\n    \n    return d0_corr\n\ndef normalise_voldata_to_stdev_3_clip(datavol):\n    # Result is returned as float value between 0.0 and 1.0\n    \n    d0_mean = np.mean(datavol)\n    d0_std = np.std(datavol)\n\n    if d0_std==0:\n        raise ValueError(\"Error. Stdev of data volume is zero.\")\n        \n    d0_corr = (datavol.astype(np.float32) - d0_mean) / d0_std\n    d0_corr = (np.clip(d0_corr, -3.0, 3.0) +3.0) / 6.0\n    \n    return d0_corr","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:56.751143Z","iopub.execute_input":"2025-02-01T18:02:56.751343Z","iopub.status.idle":"2025-02-01T18:02:56.765277Z","shell.execute_reply.started":"2025-02-01T18:02:56.751325Z","shell.execute_reply":"2025-02-01T18:02:56.764539Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# \"stdev_no_clip\" default , \"stdev3_no_clip\"\nnorm_func = normalise_voldata_to_stdev\nif \"stdev3_no_clip\" in CFG.DATA_NORM_FUNC:\n    norm_func = normalise_voldata_to_stdev_3\nif \"stdev3_with_clip\" in CFG.DATA_NORM_FUNC:\n    norm_func = normalise_voldata_to_stdev_3_clip","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:56.766088Z","iopub.execute_input":"2025-02-01T18:02:56.766290Z","iopub.status.idle":"2025-02-01T18:02:56.780414Z","shell.execute_reply.started":"2025-02-01T18:02:56.766272Z","shell.execute_reply":"2025-02-01T18:02:56.779685Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"norm_func","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:56.782404Z","iopub.execute_input":"2025-02-01T18:02:56.782636Z","iopub.status.idle":"2025-02-01T18:02:56.794382Z","shell.execute_reply.started":"2025-02-01T18:02:56.782597Z","shell.execute_reply":"2025-02-01T18:02:56.793772Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"# class CZII_Dataset5(Dataset):\n#     \"\"\"\n#     Selects regions with particles of interested being at the center\n#     50% of the data volumes are random locations across all volumes\n#     \"\"\"\n#     def __init__(self, exp_ids,\n#                  cubesize=48,\n#                  augm_rots=True, augm_flips=True):\n        \n#         self.exp_ids = exp_ids\n\n#         self.data_expid_dict={}\n        \n#         # expid, class, particle_pos\n#         self.all_locs = []\n        \n#         for exp0 in exp_ids:\n                \n#             data0, mask0, upl0 = get_data_and_sphere_masks(exp0)\n\n#             #backup\n#             self.data_expid_dict[exp0]={ 'data': data0, 'mask': mask0}\n\n#             self.all_locs.extend(upl0)\n                \n#         self.cubesize = cubesize\n\n#         count_particles= len(self.all_locs)\n\n#         other_regions_count = count_particles #50/50\n        \n#         self.datavols_shape = data0.shape\n\n#         #create a list of random locations\n#         i_exp_rand = np.random.randint(0,len(exp_ids), size=count_particles)\n#         i_z = np.random.randint(0,self.datavols_shape[0], size=count_particles)\n#         i_y = np.random.randint(0,self.datavols_shape[1], size=count_particles)\n#         i_x = np.random.randint(0,self.datavols_shape[2], size=count_particles)\n\n#         for i in range(other_regions_count):\n#             self.all_locs.append((\n#                 exp_ids[i_exp_rand[i]],\n#                 0,\n#                 (i_z[i], i_y[i], i_x[i] ) #zyx\n#             ))\n        \n#         self.augm_rots = augm_rots\n#         self.augm_flips = augm_flips\n\n#     def __len__(self):\n#         return len(self.all_locs) # including random volumes\n\n#     def __getitem__(self, idx):\n        \n#         # for the particle location choosen, get the cube centered at that point\n#         # if cube goes out of limits, adjust cube location to limits\n\n#         exp0, cl0, loc_zyx = self.all_locs[idx]\n\n#         # Load data from the experiment\n#         data_and_labels_dict = self.data_expid_dict[exp0]\n#         data0 = data_and_labels_dict['data']\n#         lbls0 = data_and_labels_dict['mask']\n\n#         z,y,x = loc_zyx\n\n#         width2 = self.cubesize/2\n        \n#         z0 = int(z - width2)\n#         if z0<0:\n#             z0=0\n#         z1 = z0 + self.cubesize\n#         if z1>self.datavols_shape[0]:\n#             z1 = self.datavols_shape[0]\n#             z0 = z1- self.cubesize\n\n#         y0 = int(y - width2)\n#         if y0<0:\n#             y0=0\n#         y1 = y0 + self.cubesize\n#         if y1>self.datavols_shape[1]:\n#             y1 = self.datavols_shape[1]\n#             y0 = y1- self.cubesize\n\n#         x0 = int(x - width2)\n#         if x0<0:\n#             x0=0\n#         x1 = x0 + self.cubesize\n#         if x1>self.datavols_shape[2]:\n#             x1 = self.datavols_shape[2]\n#             x0 = x1- self.cubesize\n\n#         datacrop = data0[z0:z1 , y0:y1, x0:x1]\n#         labelscrop = lbls0[z0:z1 , y0:y1, x0:x1]\n\n#         data1 = datacrop\n#         labels1=labelscrop\n#         if norm_func is not None:\n#             data1 = norm_func(data1)\n\n#         # Random rotate on the XY axis\n#         if self.augm_rots:\n#             r = random.randint(0,3)\n#             if r>0:\n#                 data1 = np.rot90(data1,r,axes=(1,2))\n#                 labels1 = np.rot90(labels1,r,axes=(1,2))\n\n#         # Random flips\n#         if self.augm_flips:\n#             for ax in range(3):\n#                 #random flip y 25%\n#                 if random.randint(0,3)==3:\n#                     data1 = np.flip(data1,ax)\n#                     labels1 = np.flip(labels1,ax)\n\n#         data1_t = torch.from_numpy(data1.copy()).to(torch_device)\n#         labels1_t = torch.from_numpy(labels1.copy()).to(torch_device)\n\n#         return data1_t.unsqueeze(dim=0).float(), labels1_t.long()\n\n#     def get_classes(self):\n#         #return an array with the corresponding class for each id\n#         cls0 = [ loc0[1]  for loc0 in self.all_locs]\n#         return cls0","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:56.795533Z","iopub.execute_input":"2025-02-01T18:02:56.795831Z","iopub.status.idle":"2025-02-01T18:02:56.808172Z","shell.execute_reply.started":"2025-02-01T18:02:56.795802Z","shell.execute_reply":"2025-02-01T18:02:56.807520Z"},"jupyter":{"source_hidden":true},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# class CZII_Dataset6(Dataset):\n#     \"\"\"\n#     Selects regions with particles of interested being at the center\n#     A total 2/3 of the data volumes in this datset are random\n#     locations across all training volumes\n#     \"\"\"\n#     def __init__(self, exp_ids,\n#                  cubesize=48,\n#                  augm_rots=True, augm_flips=True):\n        \n#         self.exp_ids = exp_ids\n\n#         self.data_expid_dict={}\n        \n#         # expid, class, particle_pos\n#         self.all_locs = []\n        \n#         for exp0 in exp_ids:\n                \n#             data0, mask0, upl0 = get_data_and_sphere_masks(exp0)\n\n#             #backup\n#             self.data_expid_dict[exp0]={ 'data': data0, 'mask': mask0}\n\n#             self.all_locs.extend(upl0)\n                \n#         self.cubesize = cubesize\n\n#         count_particles= len(self.all_locs)\n\n#         other_regions_count = int(1.5*count_particles) #2/3\n        \n#         self.datavols_shape = data0.shape\n\n#         #create a list of random locations\n#         i_exp_rand = np.random.randint(0,len(exp_ids), size=other_regions_count)\n#         i_z = np.random.randint(0,self.datavols_shape[0], size=other_regions_count)\n#         i_y = np.random.randint(0,self.datavols_shape[1], size=other_regions_count)\n#         i_x = np.random.randint(0,self.datavols_shape[2], size=other_regions_count)\n\n#         for i in range(other_regions_count):\n#             self.all_locs.append((\n#                 exp_ids[i_exp_rand[i]],\n#                 0,\n#                 (i_z[i], i_y[i], i_x[i] ) #zyx\n#             ))\n        \n#         self.augm_rots = augm_rots\n#         self.augm_flips = augm_flips\n\n#     def __len__(self):\n#         return len(self.all_locs) # including random volumes\n\n#     def __getitem__(self, idx):\n        \n#         # for the particle location choosen, get the cube centered at that point\n#         # if cube goes out of limits, adjust cube location to limits\n\n#         exp0, cl0, loc_zyx = self.all_locs[idx]\n\n#         # Load data from the experiment\n#         data_and_labels_dict = self.data_expid_dict[exp0]\n#         data0 = data_and_labels_dict['data']\n#         lbls0 = data_and_labels_dict['mask']\n\n#         z,y,x = loc_zyx\n\n#         width2 = self.cubesize/2\n        \n#         z0 = int(z - width2)\n#         if z0<0:\n#             z0=0\n#         z1 = z0 + self.cubesize\n#         if z1>self.datavols_shape[0]:\n#             z1 = self.datavols_shape[0]\n#             z0 = z1- self.cubesize\n\n#         y0 = int(y - width2)\n#         if y0<0:\n#             y0=0\n#         y1 = y0 + self.cubesize\n#         if y1>self.datavols_shape[1]:\n#             y1 = self.datavols_shape[1]\n#             y0 = y1- self.cubesize\n\n#         x0 = int(x - width2)\n#         if x0<0:\n#             x0=0\n#         x1 = x0 + self.cubesize\n#         if x1>self.datavols_shape[2]:\n#             x1 = self.datavols_shape[2]\n#             x0 = x1- self.cubesize\n\n#         datacrop = data0[z0:z1 , y0:y1, x0:x1]\n#         labelscrop = lbls0[z0:z1 , y0:y1, x0:x1]\n\n#         data1 = datacrop\n#         labels1=labelscrop\n#         if norm_func is not None:\n#             data1 = norm_func(data1)\n\n#         # Random rotate on the XY axis\n#         if self.augm_rots:\n#             r = random.randint(0,3)\n#             if r>0:\n#                 data1 = np.rot90(data1,r,axes=(1,2))\n#                 labels1 = np.rot90(labels1,r,axes=(1,2))\n\n#         # Random flips\n#         if self.augm_flips:\n#             for ax in range(3):\n#                 #random flip y 25%\n#                 if random.randint(0,3)==3:\n#                     data1 = np.flip(data1,ax)\n#                     labels1 = np.flip(labels1,ax)\n\n#         data1_t = torch.from_numpy(data1.copy()).to(torch_device)\n#         labels1_t = torch.from_numpy(labels1.copy()).to(torch_device)\n\n#         return data1_t.unsqueeze(dim=0).float(), labels1_t.long()\n\n#     def get_classes(self):\n#         #return an array with the corresponding class for each id\n#         cls0 = [ loc0[1]  for loc0 in self.all_locs]\n#         return cls0","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:56.808874Z","iopub.execute_input":"2025-02-01T18:02:56.809123Z","iopub.status.idle":"2025-02-01T18:02:56.824242Z","shell.execute_reply.started":"2025-02-01T18:02:56.809104Z","shell.execute_reply":"2025-02-01T18:02:56.823516Z"},"jupyter":{"source_hidden":true},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# local voxel neighboriing swap augmentation\nimport torch.nn.functional as F\ndef swap_neighbors_LMAP(volume):\n    global CFG\n    swap_prob = CFG.TRAIN_AUGM_NEIGHB_VOXEL_SWAP_PROB\n    \n    assert volume.dim() == 3, \"Input volume must be a 3D tensor\"\n    assert 0 <= swap_prob <= 1, \"swap_prob must be in the range [0, 1]\"\n\n    D, H, W = volume.shape\n\n    # Create a swap mask based on probability\n    swap_mask = torch.rand((D, H, W), device=volume.device) < swap_prob\n    #print(\"swap_mask\\n\", swap_mask)\n    # Randomly pick one of three axes (0=Z, 1=Y, 2=X) to perform swaps along\n    axis_mask = torch.randint(0, 3, (D, H, W), device=volume.device)\n    #print(\"axis_mask\\n\", axis_mask)\n\n    # Store the original values to prevent overwriting during swapping\n    vol1 = volume.clone()\n\n    z_swap_mask = swap_mask[:-1,:,:] & (axis_mask[:-1,:,:]==0)\n    # z_swap_mask_1 = F.pad(z_swap_mask,(1,0,0,0,0,0),value=False)\n    #print(\"z_swap_mask\\n\",z_swap_mask)\n    #print(new_volume[:-1,:,:][z_swap_mask])\n    vol1[:-1,:,:][z_swap_mask] = volume[1:,:,:][z_swap_mask]\n    vol1[1:,:,:][z_swap_mask] = volume[:-1,:,:][z_swap_mask]\n    #print(\"vol1\\n\", vol1)\n\n    vol0 = vol1\n    vol1 = vol0.clone()\n    y_swap_mask = swap_mask[:,:-1,:] & (axis_mask[:,:-1,:]==1)\n    #print(\"y_swap_mask\\n\",y_swap_mask)\n    vol1[:,:-1,:][y_swap_mask] = vol0[:,1:,:][y_swap_mask]\n    vol1[:,1:,:][y_swap_mask] = vol0[:,:-1,:][y_swap_mask]\n    #print(\"vol1\\n\", vol1)\n\n    vol0 = vol1\n    vol1 = vol0.clone()\n    x_swap_mask = swap_mask[:,:,:-1] & (axis_mask[:,:,:-1]==2)\n    #print(\"x_swap_mask\\n\",x_swap_mask)\n    vol1[:,:,:-1][x_swap_mask] = vol0[:,:,1:][x_swap_mask]\n    vol1[:,:,1:][x_swap_mask] = vol0[:,:,:-1][x_swap_mask]\n    #print(\"vol1\\n\", vol1)\n\n    return vol1","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:56.824968Z","iopub.execute_input":"2025-02-01T18:02:56.825177Z","iopub.status.idle":"2025-02-01T18:02:56.840475Z","shell.execute_reply.started":"2025-02-01T18:02:56.825158Z","shell.execute_reply":"2025-02-01T18:02:56.839695Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# class CZII_Dataset7(Dataset):\n#     \"\"\"\n#     Selects regions with particles of interested being at the center\n#     A total 2/3 of the data volumes in this datset are random\n#     locations across all training volumes\n\n#     Augmentations added:\n#      - neighbouring voxel shuffle augmentation\n#      - level shift\n#      - contrast (scale)\n#      - (TODO: add gaussian noise)\n     \n#     \"\"\"\n#     def __init__(self, exp_ids,\n#                  cubesize=48,\n#                  augm_rots=True, augm_flips=True):\n        \n#         self.exp_ids = exp_ids\n\n#         self.data_expid_dict={}\n        \n#         # expid, class, particle_pos\n#         self.all_locs = []\n        \n#         for exp0 in exp_ids:\n                \n#             data0, mask0, upl0 = get_data_and_sphere_masks(exp0)\n\n#             #backup\n#             self.data_expid_dict[exp0]={ 'data': data0, 'mask': mask0}\n\n#             self.all_locs.extend(upl0)\n                \n#         self.cubesize = cubesize\n\n#         count_particles= len(self.all_locs)\n\n#         other_regions_count = int(1.5*count_particles) #2/3\n        \n#         self.datavols_shape = data0.shape\n\n#         #create a list of random locations\n#         i_exp_rand = np.random.randint(0,len(exp_ids), size=other_regions_count)\n#         i_z = np.random.randint(0,self.datavols_shape[0], size=other_regions_count)\n#         i_y = np.random.randint(0,self.datavols_shape[1], size=other_regions_count)\n#         i_x = np.random.randint(0,self.datavols_shape[2], size=other_regions_count)\n\n#         for i in range(other_regions_count):\n#             self.all_locs.append((\n#                 exp_ids[i_exp_rand[i]],\n#                 0,\n#                 (i_z[i], i_y[i], i_x[i] ) #zyx\n#             ))\n        \n#         self.augm_rots = augm_rots\n#         self.augm_flips = augm_flips\n\n#     def __len__(self):\n#         return len(self.all_locs) # including random volumes\n\n#     def _random_voxel_shift(self,p=0.2):\n#         r0 = random.random()\n#         if r0< p/2:\n#             return -1\n#         elif r0<p:\n#             return 1\n#         return 0\n        \n#     def __getitem__(self, idx):\n#         global CFG\n#         # for the particle location choosen, get the cube centered at that point\n#         # if cube goes out of limits, adjust cube location to limits\n\n#         exp0, cl0, loc_zyx = self.all_locs[idx]\n\n#         # Load data from the experiment\n#         data_and_labels_dict = self.data_expid_dict[exp0]\n#         data0 = data_and_labels_dict['data']\n#         lbls0 = data_and_labels_dict['mask']\n\n#         z,y,x = loc_zyx\n\n#         width2 = self.cubesize/2\n        \n#         z0 = int(z - width2)\n#         if z0<1:\n#             z0=1\n#         z1 = z0 + self.cubesize\n#         if z1>=self.datavols_shape[0]:\n#             z1 = self.datavols_shape[0]-1\n#             z0 = z1- self.cubesize\n\n#         y0 = int(y - width2)\n#         if y0<1:\n#             y0=1\n#         y1 = y0 + self.cubesize\n#         if y1>=self.datavols_shape[1]:\n#             y1 = self.datavols_shape[1]-1\n#             y0 = y1- self.cubesize\n\n#         x0 = int(x - width2)\n#         if x0<1:\n#             x0=1\n#         x1 = x0 + self.cubesize\n#         if x1>=self.datavols_shape[2]:\n#             x1 = self.datavols_shape[2]-1\n#             x0 = x1- self.cubesize\n        \n#         labelscrop = lbls0[z0:z1 , y0:y1, x0:x1]\n\n#         #shift augmentation\n#         z_shift = self._random_voxel_shift()\n#         y_shift = self._random_voxel_shift()\n#         x_shift = self._random_voxel_shift()\n#         datacrop = data0[z0+z_shift:z1+z_shift , y0+y_shift:y1+y_shift, x0+x_shift:x1+x_shift]\n\n#         data1 = datacrop\n#         labels1=labelscrop\n#         if norm_func is not None:\n#             data1 = norm_func(data1)\n\n#         # contrast scale augmentation\n#         if random.randint(0,9)==5: #10%\n#             data1 = data1 * random.random()/0.5\n            \n#         # signal shift augmentation\n#         if random.randint(0,9)==6: #10%\n#             data1 = data1 + (random.random()/5+0.1) # +- 0.1\n\n#         # Random rotate on the XY axis\n#         if self.augm_rots:\n#             r = random.randint(0,3)\n#             if r>0:\n#                 data1 = np.rot90(data1,r,axes=(1,2))\n#                 labels1 = np.rot90(labels1,r,axes=(1,2))\n\n#         # Random flips\n#         if self.augm_flips:\n#             for ax in range(3):\n#                 #random flip y 25%\n#                 if random.randint(0,3)==3:\n#                     data1 = np.flip(data1,ax)\n#                     labels1 = np.flip(labels1,ax)\n\n\n#         data1_t = torch.from_numpy(data1.copy()).to(torch_device)\n#         labels1_t = torch.from_numpy(labels1.copy()).to(torch_device)\n\n#         #Voxel shuffle augmentation\n#         if CFG.TRAIN_AUGM_NEIGHB_VOXEL_SWAP_PROB>0.0:\n#             data1_t = swap_neighbors_LMAP(data1_t)\n\n#         return data1_t.unsqueeze(dim=0).float(), labels1_t.long()\n\n#     def get_classes(self):\n#         #return an array with the corresponding class for each id\n#         cls0 = [ loc0[1]  for loc0 in self.all_locs]\n#         return cls0","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:02:56.841306Z","iopub.execute_input":"2025-02-01T18:02:56.841583Z","iopub.status.idle":"2025-02-01T18:02:56.854166Z","shell.execute_reply.started":"2025-02-01T18:02:56.841556Z","shell.execute_reply":"2025-02-01T18:02:56.853476Z"},"trusted":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# class CZII_Dataset8(Dataset):\n#     \"\"\"\n#     Selects regions with particles of interested being at the center\n#     A total 2/3 of the data volumes in this datset are random\n#     locations across all training volumes\n\n#     Augmentations added:\n#      - add gaussian noise\n     \n#     \"\"\"\n#     def __init__(self, exp_ids,\n#                  cubesize=48,\n#                  augm_rots=True, augm_flips=True):\n        \n#         self.exp_ids = exp_ids\n\n#         self.data_expid_dict={}\n        \n#         # expid, class, particle_pos\n#         self.all_locs = []\n        \n#         for exp0 in exp_ids:\n                \n#             data0, mask0, upl0 = get_data_and_sphere_masks(exp0)\n\n#             #backup\n#             self.data_expid_dict[exp0]={ 'data': data0, 'mask': mask0}\n\n#             self.all_locs.extend(upl0)\n                \n#         self.cubesize = cubesize\n\n#         count_particles= len(self.all_locs)\n\n#         other_regions_count = int(1.5*count_particles) #2/3\n        \n#         self.datavols_shape = data0.shape\n\n#         #create a list of random locations\n#         i_exp_rand = np.random.randint(0,len(exp_ids), size=other_regions_count)\n#         i_z = np.random.randint(0,self.datavols_shape[0], size=other_regions_count)\n#         i_y = np.random.randint(0,self.datavols_shape[1], size=other_regions_count)\n#         i_x = np.random.randint(0,self.datavols_shape[2], size=other_regions_count)\n\n#         for i in range(other_regions_count):\n#             self.all_locs.append((\n#                 exp_ids[i_exp_rand[i]],\n#                 0,\n#                 (i_z[i], i_y[i], i_x[i] ) #zyx\n#             ))\n        \n#         self.augm_rots = augm_rots\n#         self.augm_flips = augm_flips\n\n#     def __len__(self):\n#         return len(self.all_locs) # including random volumes\n\n#     def _random_voxel_shift(self,p=0.2):\n#         r0 = random.random()\n#         if r0< p/2:\n#             return -1\n#         elif r0<p:\n#             return 1\n#         return 0\n        \n#     def __getitem__(self, idx):\n#         global CFG\n#         # for the particle location choosen, get the cube centered at that point\n#         # if cube goes out of limits, adjust cube location to limits\n\n#         exp0, cl0, loc_zyx = self.all_locs[idx]\n\n#         # Load data from the experiment\n#         data_and_labels_dict = self.data_expid_dict[exp0]\n#         data0 = data_and_labels_dict['data']\n#         lbls0 = data_and_labels_dict['mask']\n\n#         z,y,x = loc_zyx\n\n#         width2 = self.cubesize/2\n        \n#         z0 = int(z - width2)\n#         if z0<1:\n#             z0=1\n#         z1 = z0 + self.cubesize\n#         if z1>=self.datavols_shape[0]:\n#             z1 = self.datavols_shape[0]-1\n#             z0 = z1- self.cubesize\n\n#         y0 = int(y - width2)\n#         if y0<1:\n#             y0=1\n#         y1 = y0 + self.cubesize\n#         if y1>=self.datavols_shape[1]:\n#             y1 = self.datavols_shape[1]-1\n#             y0 = y1- self.cubesize\n\n#         x0 = int(x - width2)\n#         if x0<1:\n#             x0=1\n#         x1 = x0 + self.cubesize\n#         if x1>=self.datavols_shape[2]:\n#             x1 = self.datavols_shape[2]-1\n#             x0 = x1- self.cubesize\n        \n#         labelscrop = lbls0[z0:z1 , y0:y1, x0:x1]\n\n#         #shift augmentation\n#         z_shift = self._random_voxel_shift()\n#         y_shift = self._random_voxel_shift()\n#         x_shift = self._random_voxel_shift()\n#         datacrop = data0[z0+z_shift:z1+z_shift , y0+y_shift:y1+y_shift, x0+x_shift:x1+x_shift]\n\n#         data1 = datacrop\n#         labels1=labelscrop\n#         if norm_func is not None:\n#             data1 = norm_func(data1)\n\n#         # contrast scale augmentation\n#         if random.randint(0,9)==5: #10%\n#             data1 = data1 * random.random()/0.5\n            \n#         # signal shift augmentation\n#         if random.randint(0,9)==6: #10%\n#             data1 = data1 + (random.random()/5+0.1) # +- 0.1\n\n#         # Random rotate on the XY axis\n#         if self.augm_rots:\n#             r = random.randint(0,3)\n#             if r>0:\n#                 data1 = np.rot90(data1,r,axes=(1,2))\n#                 labels1 = np.rot90(labels1,r,axes=(1,2))\n\n#         # Random flips\n#         if self.augm_flips:\n#             for ax in range(3):\n#                 #random flip y 25%\n#                 if random.randint(0,3)==3:\n#                     data1 = np.flip(data1,ax)\n#                     labels1 = np.flip(labels1,ax)\n\n\n#         data1_t = torch.from_numpy(data1.copy()).to(torch_device)\n#         labels1_t = torch.from_numpy(labels1.copy()).to(torch_device)\n\n#         #Voxel shuffle augmentation\n#         if CFG.TRAIN_AUGM_NEIGHB_VOXEL_SWAP_PROB>0.0:\n#             data1_t = swap_neighbors_LMAP(data1_t)\n\n#         if CFG.TRAIN_AUGM_GAUSS_NOISE_STD>0.0:\n#             #mean=0.0\n#             random_tensor = torch.normal(0.0, CFG.TRAIN_AUGM_GAUSS_NOISE_STD, data1_t.size(), device=data1_t.device.type)\n#             data1_t += random_tensor\n            \n#         return data1_t.unsqueeze(dim=0).float(), labels1_t.long()\n\n#     def get_classes(self):\n#         #return an array with the corresponding class for each id\n#         cls0 = [ loc0[1]  for loc0 in self.all_locs]\n#         return cls0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:02:56.854874Z","iopub.execute_input":"2025-02-01T18:02:56.855061Z","iopub.status.idle":"2025-02-01T18:02:56.870301Z","shell.execute_reply.started":"2025-02-01T18:02:56.855045Z","shell.execute_reply":"2025-02-01T18:02:56.869572Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"AUGMENT_ENABLED = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:02:56.871223Z","iopub.execute_input":"2025-02-01T18:02:56.871500Z","iopub.status.idle":"2025-02-01T18:02:56.886312Z","shell.execute_reply.started":"2025-02-01T18:02:56.871472Z","shell.execute_reply":"2025-02-01T18:02:56.885525Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CZII_Dataset9(Dataset):\n    \"\"\"\n    Selects regions with particles of interested being at the center\n    A total 2/3 of the data volumes in this datset are random\n    locations across all training volumes\n\n    Augmentations added:\n     - add gaussian noise\n     \n    \"\"\"\n    def __init__(self, exp_ids,\n                 cubesize=48,\n                 augm_rots=True, augm_flips=True):\n        \n        self.exp_ids = exp_ids\n\n        self.data_expid_dict={}\n        \n        # expid, class, particle_pos\n        self.all_locs = []\n        \n        for exp0 in exp_ids:\n                \n            data0, mask0, upl0 = get_data_and_sphere_masks(exp0)\n\n            #backup\n            self.data_expid_dict[exp0]={ 'data': data0, 'mask': mask0}\n\n            self.all_locs.extend(upl0)\n                \n        self.cubesize = cubesize\n\n        count_particles= len(self.all_locs)\n\n        other_regions_count = int(1.5*count_particles) #2/3\n        \n        self.datavols_shape = data0.shape\n\n        #create a list of random locations\n        i_exp_rand = np.random.randint(0,len(exp_ids), size=other_regions_count)\n        i_z = np.random.randint(0,self.datavols_shape[0], size=other_regions_count)\n        i_y = np.random.randint(0,self.datavols_shape[1], size=other_regions_count)\n        i_x = np.random.randint(0,self.datavols_shape[2], size=other_regions_count)\n\n        for i in range(other_regions_count):\n            self.all_locs.append((\n                exp_ids[i_exp_rand[i]],\n                0,\n                (i_z[i], i_y[i], i_x[i] ) #zyx\n            ))\n        \n        self.augm_rots = augm_rots\n        self.augm_flips = augm_flips\n\n    def __len__(self):\n        return len(self.all_locs) # including random volumes\n\n    def _random_voxel_shift(self):\n        r0 = random.random()\n        p= CFG.TRAIN_RANDOM_VOXEL_SHIFT_PROB\n        if r0< p/2:\n            return -1\n        elif r0<p:\n            return 1\n        return 0\n        \n    def __getitem__(self, idx):\n        global CFG\n        global AUGMENT_ENABLED\n\n        #print(f\"{idx=}, {AUGMENT_ENABLED=}\") # debug\n        \n        # for the particle location choosen, get the cube centered at that point\n        # if cube goes out of limits, adjust cube location to limits\n\n        exp0, cl0, loc_zyx = self.all_locs[idx]\n\n        # Load data from the experiment\n        data_and_labels_dict = self.data_expid_dict[exp0]\n        data0 = data_and_labels_dict['data']\n        lbls0 = data_and_labels_dict['mask']\n\n        z,y,x = loc_zyx\n\n        width2 = self.cubesize/2\n        \n        z0 = int(z - width2)\n        if z0<1:\n            z0=1\n        z1 = z0 + self.cubesize\n        if z1>=self.datavols_shape[0]:\n            z1 = self.datavols_shape[0]-1\n            z0 = z1- self.cubesize\n\n        y0 = int(y - width2)\n        if y0<1:\n            y0=1\n        y1 = y0 + self.cubesize\n        if y1>=self.datavols_shape[1]:\n            y1 = self.datavols_shape[1]-1\n            y0 = y1- self.cubesize\n\n        x0 = int(x - width2)\n        if x0<1:\n            x0=1\n        x1 = x0 + self.cubesize\n        if x1>=self.datavols_shape[2]:\n            x1 = self.datavols_shape[2]-1\n            x0 = x1- self.cubesize\n\n        #shift augmentation\n        z_shift = 0\n        y_shift = 0\n        x_shift = 0\n        if AUGMENT_ENABLED and CFG.TRAIN_RANDOM_VOXEL_SHIFT_PROB>0:\n            z_shift = self._random_voxel_shift()\n            y_shift = self._random_voxel_shift()\n            x_shift = self._random_voxel_shift()\n            \n        data1 = data0[z0+z_shift:z1+z_shift , y0+y_shift:y1+y_shift, x0+x_shift:x1+x_shift]\n        labels1 = lbls0[z0:z1 , y0:y1, x0:x1]\n\n        if norm_func is not None:\n            data1 = norm_func(data1)\n\n        if AUGMENT_ENABLED:    \n            # contrast scale augmentation\n            if random.randint(0,9)==5: #10%\n                data1 = data1 * random.random()/0.5\n                \n            # signal shift augmentation\n            if random.randint(0,9)==6: #10%\n                data1 = data1 + (random.random()/5+0.1) # +- 0.1\n    \n            # Random rotate on the XY axis\n            if self.augm_rots:\n                r = random.randint(0,3)\n                if r>0:\n                    data1 = np.rot90(data1,r,axes=(1,2))\n                    labels1 = np.rot90(labels1,r,axes=(1,2))\n    \n            # Random flips\n            if self.augm_flips:\n                for ax in range(3):\n                    #random flip y 25%\n                    if random.randint(0,3)==3:\n                        data1 = np.flip(data1,ax)\n                        labels1 = np.flip(labels1,ax)\n\n\n        data1_t = torch.from_numpy(data1.copy()).to(torch_device)\n        labels1_t = torch.from_numpy(labels1.copy()).to(torch_device)\n\n        if AUGMENT_ENABLED:\n            #Voxel shuffle augmentation\n            if CFG.TRAIN_AUGM_NEIGHB_VOXEL_SWAP_PROB>0.0:\n                data1_t = swap_neighbors_LMAP(data1_t)\n    \n            if CFG.TRAIN_AUGM_GAUSS_NOISE_STD>0.0:\n                #mean=0.0\n                random_tensor = torch.normal(0.0, CFG.TRAIN_AUGM_GAUSS_NOISE_STD, data1_t.size(), device=data1_t.device.type)\n                data1_t += random_tensor\n                \n        return data1_t.unsqueeze(dim=0).float(), labels1_t.long()\n\n    def get_classes(self):\n        #return an array with the corresponding class for each id\n        cls0 = [ loc0[1]  for loc0 in self.all_locs]\n        return cls0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:02:56.887568Z","iopub.execute_input":"2025-02-01T18:02:56.887852Z","iopub.status.idle":"2025-02-01T18:02:56.903659Z","shell.execute_reply.started":"2025-02-01T18:02:56.887833Z","shell.execute_reply":"2025-02-01T18:02:56.902856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"my_dataset = CZII_Dataset9(\n    CFG.IDS_TO_TRAIN,\n    cubesize=CFG.TRAIN_CUBE_SIZE,\n    augm_rots=CFG.TRAIN_AUGM_ROTS,\n    augm_flips=CFG.TRAIN_AUGM_FLIPS\n)","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:02.489047Z","iopub.execute_input":"2025-02-01T18:03:02.489357Z","iopub.status.idle":"2025-02-01T18:03:31.944182Z","shell.execute_reply.started":"2025-02-01T18:03:02.489329Z","shell.execute_reply":"2025-02-01T18:03:31.943212Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#my_dataset.all_locs","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:31.945420Z","iopub.execute_input":"2025-02-01T18:03:31.945776Z","iopub.status.idle":"2025-02-01T18:03:31.949036Z","shell.execute_reply.started":"2025-02-01T18:03:31.945741Z","shell.execute_reply":"2025-02-01T18:03:31.948230Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(my_dataset)","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:31.950269Z","iopub.execute_input":"2025-02-01T18:03:31.950553Z","iopub.status.idle":"2025-02-01T18:03:31.964729Z","shell.execute_reply.started":"2025-02-01T18:03:31.950533Z","shell.execute_reply":"2025-02-01T18:03:31.963911Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"idx0 = random.randint(0, len(my_dataset)-1)\nprint(idx0)\n#idx0=3\nzslice = 31\ndata0, mask0 = my_dataset[idx0]\nprint(f\"shapes data, mask : {data0.shape}, {mask0.shape}\")\n\nplt.figure(figsize=(5,5))\nplt.imshow(data0[0,zslice,...].cpu().numpy(), cmap='gray')\nplt.imshow(mask0[zslice,...].cpu().numpy(), cmap='tab10',vmin=0,vmax=10, alpha=0.5)\nplt.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:35.108911Z","iopub.execute_input":"2025-02-01T18:03:35.109330Z","iopub.status.idle":"2025-02-01T18:03:35.637258Z","shell.execute_reply.started":"2025-02-01T18:03:35.109296Z","shell.execute_reply":"2025-02-01T18:03:35.636148Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Setup model","metadata":{}},{"cell_type":"code","source":"model3d = create_WMSCSFB_grn_mednext_v1(\n    num_input_channels = CFG.WMCSFB_V1_INP_CHN,\n    num_classes = CFG.WMCSFB_V1_NUM_CLASSES,\n    model_id = CFG.WMCSFBL_V1_ID,\n    kernel_size = CFG.WMCSFB_V1_KERNEL_SIZE\n).to(torch_device)","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:38.542389Z","iopub.execute_input":"2025-02-01T18:03:38.542739Z","iopub.status.idle":"2025-02-01T18:03:40.483295Z","shell.execute_reply.started":"2025-02-01T18:03:38.542702Z","shell.execute_reply":"2025-02-01T18:03:40.482432Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# load model if pretrained but don't load its config\n\nif CFG.WMCSFB_USE_PRETRAINED_MODEL is not None:\n    model_dict = torch.load(CFG.WMCSFB_USE_PRETRAINED_MODEL, weights_only=False)\n    model3d.load_state_dict(model_dict['model_state_dict'])","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:42.191232Z","iopub.execute_input":"2025-02-01T18:03:42.191528Z","iopub.status.idle":"2025-02-01T18:03:42.195527Z","shell.execute_reply.started":"2025-02-01T18:03:42.191505Z","shell.execute_reply":"2025-02-01T18:03:42.194675Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loss","metadata":{}},{"cell_type":"code","source":"import segmentation_models_pytorch.utils","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:44.191848Z","iopub.execute_input":"2025-02-01T18:03:44.192133Z","iopub.status.idle":"2025-02-01T18:03:44.198688Z","shell.execute_reply.started":"2025-02-01T18:03:44.192110Z","shell.execute_reply":"2025-02-01T18:03:44.197916Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Custom Tverski-Loss with class weighting, chatgpt with SMP example help\n\nfrom typing import Optional\nimport torch\nfrom segmentation_models_pytorch.losses.tversky import TverskyLoss  # Import your existing TverskyLoss implementation\n\n# class TverskyLoss_with_class_weights(TverskyLoss):\n#     \"\"\"\n#     Tversky loss with support for class weights.\n    \n#     Args:\n#         class_weights: Optional tensor of shape (C,) where C is the number of classes,\n#         specifying the weight for each class.\n#     \"\"\"\n#     def __init__(self, *args, class_weights: Optional[torch.Tensor] = None, **kwargs):\n#         super().__init__(*args, **kwargs)\n#         self.class_weights = torch.tensor(class_weights)\n\n#     def compute_score(self, output, target, smooth=0.0, eps=1e-7, dims=None) -> torch.Tensor:\n#         # Compute the base Tversky score\n#         tversky_score = super().compute_score(output, target, smooth, eps, dims)\n        \n#         # Apply class weights if provided\n#         if self.class_weights is not None:\n#             if self.mode == \"multiclass\":  # MULTICLASS_MODE\n#                 tversky_score *= self.class_weights.view(1, -1, 1, 1)\n#             else:  # MULTILABEL_MODE or BINARY_MODE\n#                 tversky_score *= self.class_weights.view(1, -1, *[1] * (output.ndim - 2))\n\n#         return tversky_score\n\n\nclass TverskyLoss_with_class_weights_3d(TverskyLoss):\n    \"\"\"\n    Tversky loss with support for class weights, compatible with 2D and 3D data.\n    \n    Args:\n        class_weights: Optional tensor of shape (C,) where C is the number of classes,\n        specifying the weight for each class.\n    \"\"\"\n    def __init__(self, *args, class_weights: Optional[torch.Tensor] = None, **kwargs):\n        super().__init__(*args, **kwargs)\n        self.class_weights = class_weights\n\n    def compute_score(self, output, target, smooth=0.0, eps=1e-7, dims=None) -> torch.Tensor:\n        # Compute the base Tversky score\n        tversky_score = super().compute_score(output, target, smooth, eps, dims)\n\n        # Apply class weights if provided\n        if self.class_weights is not None:\n            # Directly apply the weights to the 1D tversky_score tensor\n            tversky_score *= self.class_weights\n\n        return tversky_score\n","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:45.857846Z","iopub.execute_input":"2025-02-01T18:03:45.858131Z","iopub.status.idle":"2025-02-01T18:03:45.863777Z","shell.execute_reply.started":"2025-02-01T18:03:45.858111Z","shell.execute_reply":"2025-02-01T18:03:45.862840Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#loss_func = smp.losses.TverskyLoss(mode='multiclass', from_logits=True,alpha=0.3, beta=0.7, eps=1e-4).to(torch_device)\n#loss_func = smp.losses.TverskyLoss(mode='multiclass', from_logits=True,alpha=0.3, beta=0.7, eps=1e-4, classes=[1]).to(torch_device)\n#loss_func = smp.losses.TverskyLoss(mode='multiclass', from_logits=True ,alpha=0.3, beta=0.7, eps=1e-6).to(torch_device)\n#eps=1e-9 leads to nan; 1e-7 leads to NaN but rarely, 1e-3 leads to no training\n\nloss_func = TverskyLoss_with_class_weights_3d(mode='multiclass', from_logits=True ,\n                                                      alpha=0.3, beta=0.7, eps=1e-6,\n                                                      class_weights= torch.tensor(CFG.CLASS_WEIGHTS).to(torch_device)/2\n                                             ).to(torch_device)\n\n#activ=torch.nn.functional.sigmoid\nactiv=None\nloss_func_and_activ= {\"func\":loss_func, \"activ\":activ}\nmetric_fn = smp.utils.metrics.IoU()","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:47.884958Z","iopub.execute_input":"2025-02-01T18:03:47.885270Z","iopub.status.idle":"2025-02-01T18:03:47.904350Z","shell.execute_reply.started":"2025-02-01T18:03:47.885243Z","shell.execute_reply":"2025-02-01T18:03:47.903759Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Checkpoint, DS, model and loss","metadata":{}},{"cell_type":"code","source":"len(my_dataset)","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:49.821667Z","iopub.execute_input":"2025-02-01T18:03:49.822005Z","iopub.status.idle":"2025-02-01T18:03:49.826853Z","shell.execute_reply.started":"2025-02-01T18:03:49.821978Z","shell.execute_reply":"2025-02-01T18:03:49.825914Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"idx1 = random.randint(0, len(my_dataset)-1)\nprint(idx0)\ndata1, mask1 = my_dataset[idx1]","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:50.081974Z","iopub.execute_input":"2025-02-01T18:03:50.082189Z","iopub.status.idle":"2025-02-01T18:03:50.092282Z","shell.execute_reply.started":"2025-02-01T18:03:50.082171Z","shell.execute_reply":"2025-02-01T18:03:50.091514Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data1.shape","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:51.622042Z","iopub.execute_input":"2025-02-01T18:03:51.622319Z","iopub.status.idle":"2025-02-01T18:03:51.627212Z","shell.execute_reply.started":"2025-02-01T18:03:51.622297Z","shell.execute_reply":"2025-02-01T18:03:51.626516Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# use previous data0, mask0\nmodel3d.eval()\nwith torch.no_grad():\n    y = model3d(data0.unsqueeze(dim=0))","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:53.609975Z","iopub.execute_input":"2025-02-01T18:03:53.610287Z","iopub.status.idle":"2025-02-01T18:03:54.353037Z","shell.execute_reply.started":"2025-02-01T18:03:53.610258Z","shell.execute_reply":"2025-02-01T18:03:54.352326Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"L = loss_func(y, mask0.unsqueeze(dim=0))\nL","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:54.354032Z","iopub.execute_input":"2025-02-01T18:03:54.354247Z","iopub.status.idle":"2025-02-01T18:03:54.563958Z","shell.execute_reply.started":"2025-02-01T18:03:54.354229Z","shell.execute_reply":"2025-02-01T18:03:54.563110Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# check loss_func zero\n#loss_func(  torch.nn.functional.one_hot(mask0).float().permute(dims=(3,0,1,2)).unsqueeze(dim=0), mask0.unsqueeze(dim=0))","metadata":{"execution":{"iopub.status.busy":"2025-02-01T18:03:56.396812Z","iopub.execute_input":"2025-02-01T18:03:56.397102Z","iopub.status.idle":"2025-02-01T18:03:56.400667Z","shell.execute_reply.started":"2025-02-01T18:03:56.397080Z","shell.execute_reply":"2025-02-01T18:03:56.399712Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"ok","metadata":{}},{"cell_type":"markdown","source":"# LR finder","metadata":{}},{"cell_type":"code","source":"ds_mocktrain, _ = torch.utils.data.random_split(my_dataset, [0.8,0.2])\ndl_train_mock = DataLoader(ds_mocktrain, batch_size=CFG.TRAIN_BATCH_SIZE)\noptimizer_lrfinder = torch.optim.AdamW(model3d.parameters(), lr=1e-8)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:05:00.656911Z","iopub.execute_input":"2025-02-01T18:05:00.657222Z","iopub.status.idle":"2025-02-01T18:05:00.663458Z","shell.execute_reply.started":"2025-02-01T18:05:00.657199Z","shell.execute_reply":"2025-02-01T18:05:00.662580Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lr_finder = LRFinder(model3d, optimizer_lrfinder, loss_func, device=torch_device)\n#lr_finder.range_test(dl_train_mock, start_lr=1e-6, end_lr=1e0, num_iter=50)\nlr_finder.range_test(dl_train_mock, start_lr=1e-6, end_lr=1e-1, num_iter=50)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:05:00.863035Z","iopub.execute_input":"2025-02-01T18:05:00.863277Z","iopub.status.idle":"2025-02-01T18:05:43.249425Z","shell.execute_reply.started":"2025-02-01T18:05:00.863258Z","shell.execute_reply":"2025-02-01T18:05:43.248466Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lr_finder.plot() # to inspect the loss-learning rate graph","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:05:43.250548Z","iopub.execute_input":"2025-02-01T18:05:43.250842Z","iopub.status.idle":"2025-02-01T18:05:43.653270Z","shell.execute_reply.started":"2025-02-01T18:05:43.250819Z","shell.execute_reply":"2025-02-01T18:05:43.652508Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lr_finder.reset() # to reset the model and optimizer to their initial state\ndel(dl_train_mock)\ndel(ds_mocktrain)\ntorch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:05:43.654582Z","iopub.execute_input":"2025-02-01T18:05:43.654921Z","iopub.status.idle":"2025-02-01T18:05:44.044732Z","shell.execute_reply.started":"2025-02-01T18:05:43.654894Z","shell.execute_reply":"2025-02-01T18:05:44.043839Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Best model preparation\n\nHave a way to keep track of best model from validation score","metadata":{}},{"cell_type":"code","source":"BEST_MODEL = {\"model_stdict\": None, \"best_metric\": 0}\n\ndef bm_model_max(model0, metric0):\n    global BEST_MODEL\n    if metric0>=BEST_MODEL[\"best_metric\"]:\n        BEST_MODEL[\"best_metric\"] = metric0\n        BEST_MODEL[\"model_stdict\"] = model0.state_dict().copy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:05:54.876239Z","iopub.execute_input":"2025-02-01T18:05:54.876578Z","iopub.status.idle":"2025-02-01T18:05:54.880770Z","shell.execute_reply.started":"2025-02-01T18:05:54.876548Z","shell.execute_reply":"2025-02-01T18:05:54.879873Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Setup training","metadata":{}},{"cell_type":"code","source":"print(f\"{CFG.TRAIN_LR=} , {CFG.TRAIN_MAX_LR=}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:05:55.699378Z","iopub.execute_input":"2025-02-01T18:05:55.699694Z","iopub.status.idle":"2025-02-01T18:05:55.704036Z","shell.execute_reply.started":"2025-02-01T18:05:55.699668Z","shell.execute_reply":"2025-02-01T18:05:55.703196Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# OVERRIDE train settings, after initial\n# CFG.TRAIN_EPOCHS = 8\n# #CFG.TRAIN_K_FOLDS_PER_EPOCH = 4\n# CFG.TRAIN_LR = 1.0e-3\n# CFG.TRAIN_MAX_LR = 1.0e-2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:06:06.495353Z","iopub.execute_input":"2025-02-01T18:06:06.495688Z","iopub.status.idle":"2025-02-01T18:06:06.499243Z","shell.execute_reply.started":"2025-02-01T18:06:06.495659Z","shell.execute_reply":"2025-02-01T18:06:06.498256Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"global_lr_log = []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:06:06.711049Z","iopub.execute_input":"2025-02-01T18:06:06.711339Z","iopub.status.idle":"2025-02-01T18:06:06.714875Z","shell.execute_reply.started":"2025-02-01T18:06:06.711313Z","shell.execute_reply":"2025-02-01T18:06:06.713932Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_loop(dataloader, model, loss_func_and_activ, optimizer, scaler, scheduler, log_every_batch_count=20, sched_do_step=True):\n    \"\"\"\n    Generic function for training a pytorch model\n    with provided model, loss_func_and_activ, optimizer, scaler, scheduler\n\n    This trains for a single epoch only.\n\n    Returns nothing, but the model weigths have certainly changed for better.\n    \"\"\"\n    global AUGMENT_ENABLED\n    \n    loss_fn = loss_func_and_activ[\"func\"]\n    activ_fn = loss_func_and_activ[\"activ\"]\n    \n    total_samples = len(dataloader.dataset)\n    count_samples = 0\n    \n    # Set the model to training mode - important for batch normalization and dropout layers\n    # Unnecessary in this situation but added for best practices\n    model.train()\n    AUGMENT_ENABLED=True # runs test with augmentations\n    \n    for batch, (X, y) in enumerate(dataloader): #problem if data is 184 size\n\n        optimizer.zero_grad()\n        \n        pred_nan_detected=False\n        \n        # Compute prediction and loss\n        with torch.amp.autocast(device_type=torch_device):\n            pred = model(X) \n            \n            if torch.isnan(pred).any():\n                pred_nan_detected=True\n                \n            if activ_fn is not None:\n                pred = activ_fn(pred)\n            \n            loss = loss_fn(pred, y)\n\n        if not pred_nan_detected:\n            # Backpropagation\n            #loss.backward()\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            \n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=3.0)\n        else:\n            print(\"pred has nan value\")\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.5)\n        \n        if sched_do_step:\n            scheduler.step() #in OneCycleLR this should be done. For others, the step should be handled elsewhere\n\n        count_samples+= len(X)\n        \n        if log_every_batch_count is not None:\n            if batch % log_every_batch_count == 0 or batch==len(dataloader)-1:\n                loss = loss.item()\n                cur_lr = scheduler.get_last_lr()[0]\n                global_lr_log.append(cur_lr)\n                print(f\"batch:{batch}/{(len(dataloader)-1)}  loss:{loss:>7f}  [{count_samples:>5d}/{total_samples:>5d}]. lr:{cur_lr:.3e}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:06:06.938709Z","iopub.execute_input":"2025-02-01T18:06:06.938919Z","iopub.status.idle":"2025-02-01T18:06:06.945859Z","shell.execute_reply.started":"2025-02-01T18:06:06.938902Z","shell.execute_reply":"2025-02-01T18:06:06.945053Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# cur_lr = 123456.789\n# formatted_number = f\". lr:{cur_lr:.3e}\"\n# print(formatted_number)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:06:09.948788Z","iopub.execute_input":"2025-02-01T18:06:09.949079Z","iopub.status.idle":"2025-02-01T18:06:09.952701Z","shell.execute_reply.started":"2025-02-01T18:06:09.949058Z","shell.execute_reply":"2025-02-01T18:06:09.951747Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def test_loop(dataloader, model, loss_func_and_activ, metric_fn=None):\n    \"\"\"\n    Generic function for runing test loop with provided parameters\n    Gets the average loss and average metric calcualted on the test data.\n    It can also calculate metric value\n\n    Note that the metric is calculated per batch and then averaged on all batches.\n    This may not be the most a appropriate way to calculate the metrics.\n\n    \"\"\"\n    global AUGMENT_ENABLED\n    \n    # Set the model to evaluation mode - important for batch normalization and dropout layers\n    # Unnecessary in this situation but added for best practices\n    model.eval()\n    #size = len(dataloader.dataset)\n    #num_batches = len(dataloader)\n    AUGMENT_ENABLED=False # runs test without augmentations\n\n    loss_fn = loss_func_and_activ[\"func\"]\n    activ_fn = loss_func_and_activ[\"activ\"]\n\n    test_losses=[]\n    test_metrics=[]\n\n    # Evaluating the model with torch.no_grad() ensures that no gradients are computed during test mode\n    # also serves to reduce unnecessary gradient computations and memory usage for tensors with requires_grad=True\n    with torch.no_grad():\n        for X, y in dataloader:\n            #X=X_parse(X)\n            #y=y_parse(y)\n            \n            pred = model(X)\n\n            if activ_fn is not None:\n                pred = activ_fn(pred)\n\n            loss = loss_fn(pred, y)\n\n            test_loss = loss.item()\n            test_losses.append(test_loss)\n            \n            if metric_fn is not None:\n                pred_argmax = torch.argmax(pred, dim=1)\n                metric = metric_fn(pred_argmax, y).item()\n                test_metrics.append(metric)\n            # #metric\n            # correct += (pred.argmax(1) == y).type(torch.float).sum().item()\n\n    avg_loss = np.mean(np.array(test_losses))\n    print(f\"Avg loss: {avg_loss:>8f}\")\n\n    avg_metric=None\n    if not metric_fn is None:\n        avg_metric = np.mean(np.array(test_metrics))\n        print(f\"Avg metric: {avg_metric:>8f}\")\n        bm_model_max(model, avg_metric) #New, saves best model\n        \n    return {\"avg_loss\":avg_loss, \"avg_metric\":avg_metric, \"test_metrics\":test_metrics, \"test_losses\":test_losses}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:06:10.350786Z","iopub.execute_input":"2025-02-01T18:06:10.351130Z","iopub.status.idle":"2025-02-01T18:06:10.357991Z","shell.execute_reply.started":"2025-02-01T18:06:10.351101Z","shell.execute_reply":"2025-02-01T18:06:10.356720Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Stratified k-fold","metadata":{}},{"cell_type":"code","source":" # stratifiedkfold\nidxs = range(len(my_dataset))\nidxs_old = list(idxs).copy()\nclasses = my_dataset.get_classes()\nclasses_old = classes.copy()\nsfk = StratifiedKFold(n_splits=CFG.TRAIN_K_FOLDS_PER_EPOCH, shuffle=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:06:13.072278Z","iopub.execute_input":"2025-02-01T18:06:13.072577Z","iopub.status.idle":"2025-02-01T18:06:13.076936Z","shell.execute_reply.started":"2025-02-01T18:06:13.072552Z","shell.execute_reply":"2025-02-01T18:06:13.075916Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = torch.optim.AdamW(model3d.parameters(), lr=CFG.TRAIN_LR)\nscaler=torch.amp.GradScaler(torch_device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:06:13.905000Z","iopub.execute_input":"2025-02-01T18:06:13.905287Z","iopub.status.idle":"2025-02-01T18:06:13.910547Z","shell.execute_reply.started":"2025-02-01T18:06:13.905265Z","shell.execute_reply":"2025-02-01T18:06:13.909433Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAIN_EPOCHS = CFG.TRAIN_EPOCHS\nTRAIN_BATCH_SIZE = CFG.TRAIN_BATCH_SIZE\nTRAIN_K_FOLDS_PER_EPOCH = CFG.TRAIN_K_FOLDS_PER_EPOCH\nTRAIN_MAX_LR = CFG.TRAIN_MAX_LR","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:06:14.631082Z","iopub.execute_input":"2025-02-01T18:06:14.631393Z","iopub.status.idle":"2025-02-01T18:06:14.635714Z","shell.execute_reply.started":"2025-02-01T18:06:14.631365Z","shell.execute_reply":"2025-02-01T18:06:14.634797Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_with_sfk_oneCycleLR(sfk0, optim0, scaler0):\n\n    #Uses Train configuration from global scope rather than from CFG\n    \n    global global_lr_log\n    global_lr_log.clear() #reset\n\n    global idxs\n    global classes\n    global TRAIN_EPOCHS\n    global TRAIN_BATCH_SIZE\n    global TRAIN_K_FOLDS_PER_EPOCH\n    global TRAIN_MAX_LR\n    #global CFG\n    \n    num_steps = 0\n    for x_ids, _ in sfk0.split(idxs,classes):\n        num_steps+= math.ceil(len(x_ids)/TRAIN_BATCH_SIZE)\n    num_steps*= TRAIN_EPOCHS\n    print(f\"{num_steps=}\")\n\n    # def calculate_total_steps(dataset_size, batch_size, epochs, n_splits):\n    #     batches_per_fold = math.ceil((dataset_size * (n_splits - 1) / n_splits) / batch_size)\n    #     steps_per_fold = batches_per_fold * epochs\n    #     total_steps = steps_per_fold * n_splits\n    #     return total_steps\n    \n    # #chatGPT calculation, is consistent\n    # num_steps1 = calculate_total_steps(len(my_dataset), CFG.TRAIN_BATCH_SIZE, CFG.TRAIN_EPOCHS, CFG.TRAIN_K_FOLDS_PER_EPOCH)\n    # num_steps1\n\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(\n        optim0,\n        max_lr= TRAIN_MAX_LR,\n        total_steps=num_steps\n        )\n\n    test_results=[]\n    \n    for ep in range(TRAIN_EPOCHS):\n        print(f\"-- Epoch {ep+1}/{TRAIN_EPOCHS} -------\")\n        \n        for k, (x_ids, val_ids) in enumerate(sfk0.split(idxs,classes)):\n            print(f\"---- k-fold {k+1}/{TRAIN_K_FOLDS_PER_EPOCH} ----\")\n            # create a new dataset and dataloader from \n            ds_train = Subset(my_dataset, x_ids)\n            ds_valid = Subset(my_dataset, val_ids)\n        \n            dl_train = DataLoader(ds_train, batch_size= TRAIN_BATCH_SIZE, shuffle=True)\n            dl_valid = DataLoader(ds_valid, batch_size= TRAIN_BATCH_SIZE, shuffle=False)\n        \n            train_loop(dl_train, model3d, loss_func_and_activ, optim0, scaler0, scheduler)\n        \n            test_res=None\n            if dl_valid is not None:\n                test_res= test_loop(dl_valid, model3d, loss_func_and_activ, metric_fn=metric_fn)\n                test_results.append(test_res)\n            \n            torch.cuda.empty_cache()\n        \n    print(f\"Done!\")\n    if dl_valid is not None:\n        print(f\"Final test loss is : {test_res['avg_loss']}, and metric is: {test_res['avg_metric']}\")\n\n    return test_results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:06:16.982392Z","iopub.execute_input":"2025-02-01T18:06:16.982746Z","iopub.status.idle":"2025-02-01T18:06:16.989782Z","shell.execute_reply.started":"2025-02-01T18:06:16.982715Z","shell.execute_reply":"2025-02-01T18:06:16.988851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_with_sfk_stepLR(sfk0, optim0, scaler0):\n    #train step2\n\n    global global_lr_log\n    global_lr_log.clear() #reset\n    global CFG\n    \n    scheduler = torch.optim.lr_scheduler.StepLR(\n        optim0,\n        step_size = 1,\n        gamma=CFG.TRAIN2_GAMMA\n        )\n    #    TRAIN3_COSWREST_INIT_LR = 1e-4,\n    # TRAIN3_COSWREST_EPOCHS = 6, # Note that each epoch runs a cosine cycle. Each epoch uses mini epochs, using sfk\n    \n    # TRAIN2_INIT_LR\n    # set LR manually\n    for param_group in optim0.param_groups:\n        param_group['lr'] = CFG.TRAIN2_INIT_LR\n    \n    scheduler.step() #ensure it is applied\n    \n    test_results=[]\n    \n    for ep in range(CFG.TRAIN2_EPOCHS):\n        print(f\"-- Epoch {ep+1}/{CFG.TRAIN2_EPOCHS} -------\")\n        \n        for k, (x_ids, val_ids) in enumerate(sfk0.split(idxs,classes)):\n            print(f\"---- k-fold {k+1}/{CFG.TRAIN_K_FOLDS_PER_EPOCH} ----\")\n            # create a new dataset and dataloader from \n            ds_train = Subset(my_dataset, x_ids)\n            ds_valid = Subset(my_dataset, val_ids)\n        \n            dl_train = DataLoader(ds_train, batch_size=CFG.TRAIN_BATCH_SIZE, shuffle=True)\n            dl_valid = DataLoader(ds_valid, batch_size=CFG.TRAIN_BATCH_SIZE, shuffle=False)\n\n            #def train_loop(dataloader, model, loss_func_and_activ, optimizer, scaler, scheduler, log_every_batch_count=20, sched_do_step=True):\n            train_loop(dl_train, model3d, loss_func_and_activ, optim0, scaler0, scheduler, sched_do_step=False)\n        \n            test_res=None\n            if dl_valid is not None:\n                test_res= test_loop(dl_valid, model3d, loss_func_and_activ, metric_fn=metric_fn)\n                test_results.append(test_res)\n            \n            torch.cuda.empty_cache()\n            \n        scheduler.step()\n        \n    print(f\"Done!\")\n    if dl_valid is not None:\n        print(f\"Final test loss is : {test_res['avg_loss']}, and metric is: {test_res['avg_metric']}\")\n\n    return test_results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:06:19.574761Z","iopub.execute_input":"2025-02-01T18:06:19.575080Z","iopub.status.idle":"2025-02-01T18:06:19.581942Z","shell.execute_reply.started":"2025-02-01T18:06:19.575058Z","shell.execute_reply":"2025-02-01T18:06:19.580981Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_with_sfk_CosAnnLR(sfk0, scaler0):\n    global global_lr_log\n    global_lr_log.clear() #reset\n    global CFG\n\n    test_results=[]\n    \n    # set LR manually\n    optim0 = torch.optim.AdamW(model3d.parameters(), lr=CFG.TRAIN3_COSWREST_INIT_LR)\n    \n    scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n        optim0,\n        T_0 = CFG.TRAIN_K_FOLDS_PER_EPOCH,\n    )\n    \n    for ep in range(CFG.TRAIN3_COSWREST_EPOCHS):\n        print(f\"-- Epoch {ep+1}/{CFG.TRAIN3_COSWREST_EPOCHS} -------\")\n        \n        for k, (x_ids, val_ids) in enumerate(sfk0.split(idxs,classes)):\n            print(f\"---- k-fold {k+1}/{CFG.TRAIN_K_FOLDS_PER_EPOCH} ----\")\n            # create a new dataset and dataloader from \n            ds_train = Subset(my_dataset, x_ids)\n            ds_valid = Subset(my_dataset, val_ids)\n        \n            dl_train = DataLoader(ds_train, batch_size=CFG.TRAIN_BATCH_SIZE, shuffle=True)\n            dl_valid = DataLoader(ds_valid, batch_size=CFG.TRAIN_BATCH_SIZE, shuffle=False)\n\n            #def train_loop(dataloader, model, loss_func_and_activ, optimizer, scaler, scheduler, log_every_batch_count=20, sched_do_step=True):\n            train_loop(dl_train, model3d, loss_func_and_activ, optim0, scaler0, scheduler, sched_do_step=False)\n        \n            test_res=None\n            if dl_valid is not None:\n                test_res= test_loop(dl_valid, model3d, loss_func_and_activ, metric_fn=metric_fn)\n                test_results.append(test_res)\n            \n            torch.cuda.empty_cache()\n            \n            scheduler.step()\n        \n    print(f\"Done!\")\n    if dl_valid is not None:\n        print(f\"Final test loss is : {test_res['avg_loss']}, and metric is: {test_res['avg_metric']}\")\n\n    return test_results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:06:19.739891Z","iopub.execute_input":"2025-02-01T18:06:19.740195Z","iopub.status.idle":"2025-02-01T18:06:19.747116Z","shell.execute_reply.started":"2025-02-01T18:06:19.740171Z","shell.execute_reply":"2025-02-01T18:06:19.746099Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Do training\n\nFirst with oneCycleLR","metadata":{}},{"cell_type":"code","source":"test_results = train_with_sfk_oneCycleLR(sfk, optimizer, scaler)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-01T18:06:23.211033Z","iopub.execute_input":"2025-02-01T18:06:23.211447Z","iopub.status.idle":"2025-02-01T18:06:59.555150Z","shell.execute_reply.started":"2025-02-01T18:06:23.211412Z","shell.execute_reply":"2025-02-01T18:06:59.553898Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"plot training progression","metadata":{}},{"cell_type":"code","source":"fig, axs= plt.subplots(3,1)\n\navg_losses = [tr0['avg_loss'] for tr0 in test_results]\navg_metrics = [tr0['avg_metric'] for tr0 in test_results]\n\nfor tr0 in test_results:\n    axs[0].plot( avg_losses)\n    axs[0].set_ylabel(\"avg_loss\", size=\"x-small\")\n    axs[0].set_xlabel(\"epoch\", size=\"x-small\")\n    axs[1].plot(avg_metrics)\n    axs[1].set_ylabel(\"avg_metrics\", size=\"x-small\")\n    axs[1].set_xlabel(\"epoch\", size=\"x-small\")\n    axs[2].plot(global_lr_log)\n    axs[2].set_ylabel(\"lr\", size=\"x-small\") ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"datetime_str= datetime.datetime.now().strftime(\"%Y-%m-%d-%H%M\")\nfname_stem= f\"{datetime_str}_wmcsfb_grn\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_fn = f\"{fname_stem}_step1.ptd\"\nprint(model_fn)\n\nmodel_dict = {'model_state_dict': model3d.state_dict(),\n             'CFG':CFG}\ntorch.save(model_dict, model_fn)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Part2 StepLR, if epochs>0","metadata":{}},{"cell_type":"code","source":"# if CFG.TRAIN2_EPOCHS>0:\n#     test_results = train_with_sfk_stepLR(sfk, optimizer, scaler)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# if CFG.TRAIN2_EPOCHS>0:\n    \n#     fig, axs= plt.subplots(3,1)\n    \n#     avg_losses = [tr0['avg_loss'] for tr0 in test_results]\n#     avg_metrics = [tr0['avg_metric'] for tr0 in test_results]\n    \n#     for tr0 in test_results:\n#         axs[0].plot( avg_losses)\n#         axs[0].set_ylabel(\"avg_loss\", size=\"x-small\")\n#         axs[0].set_xlabel(\"epoch\", size=\"x-small\")\n#         axs[1].plot(avg_metrics)\n#         axs[1].set_ylabel(\"avg_metrics\", size=\"x-small\")\n#         axs[1].set_xlabel(\"epoch\", size=\"x-small\")\n#         axs[2].plot(global_lr_log)\n#         axs[2].set_ylabel(\"lr\", size=\"x-small\") ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# if CFG.TRAIN2_EPOCHS>0:\n#     model_fn = f\"{fname_stem}_step2.ptd\"\n#     print(model_fn)\n    \n#     model_dict = {'model_state_dict': model3d.state_dict(),\n#                  'CFG':CFG}\n#     torch.save(model_dict, model_fn)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Part 3\n","metadata":{}},{"cell_type":"code","source":"if CFG.TRAIN3_COSWREST_EPOCHS>0:\n    test_results = train_with_sfk_CosAnnLR(sfk, scaler)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if CFG.TRAIN3_COSWREST_EPOCHS>0:\n    \n    fig, axs= plt.subplots(3,1)\n    \n    avg_losses = [tr0['avg_loss'] for tr0 in test_results]\n    avg_metrics = [tr0['avg_metric'] for tr0 in test_results]\n    \n    for tr0 in test_results:\n        axs[0].plot( avg_losses)\n        axs[0].set_ylabel(\"avg_loss\", size=\"x-small\")\n        axs[0].set_xlabel(\"epoch\", size=\"x-small\")\n        axs[1].plot(avg_metrics)\n        axs[1].set_ylabel(\"avg_metrics\", size=\"x-small\")\n        axs[1].set_xlabel(\"epoch\", size=\"x-small\")\n        axs[2].plot(global_lr_log)\n        axs[2].set_ylabel(\"lr\", size=\"x-small\") ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if CFG.TRAIN3_COSWREST_EPOCHS>0:\n    model_fn = f\"{fname_stem}_step3.ptd\"\n    print(model_fn)\n    \n    model_dict = {'model_state_dict': model3d.state_dict(),\n                 'CFG':CFG}\n    torch.save(model_dict, model_fn)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train4 - Find best 90% datavols indices\n\nExclude them from training and retrain using onecycleLR","metadata":{}},{"cell_type":"markdown","source":"Run onecycleLR again but excluding indices and adjusted LRs","metadata":{}},{"cell_type":"code","source":"if CFG.TRAIN4_EPOCHS>0:\n    dl_all = DataLoader(my_dataset, batch_size=1, shuffle=False)\n    test_res_all= test_loop(dl_all, model3d, loss_func_and_activ, metric_fn=metric_fn)\n    test_res_all_disc_perc = np.percentile(test_res_all['test_metrics'], CFG.TRAIN4_DISCRIMINATE_PERC)\n    #print(f\"{test_res_all_disc_perc=}\")\n    \n    indices_over_disc_perc = np.where(test_res_all['test_metrics'] > test_res_all_disc_perc)[0]\n    idxs = indices_over_disc_perc.tolist()\n    classes = np.array(classes_old)[idxs].tolist()\n\n    TRAIN_EPOCHS = CFG.TRAIN4_EPOCHS\n    TRAIN_K_FOLDS_PER_EPOCH = CFG.TRAIN4_K_FOLDS_PER_EPOCH\n    TRAIN_LR = CFG.TRAIN4_LR\n    TRAIN_MAX_LR = CFG.TRAIN4_MAX_LR\n\n    sfk4 = StratifiedKFold(n_splits=CFG.TRAIN4_K_FOLDS_PER_EPOCH, shuffle=True)\n    optimizer4 = torch.optim.AdamW(model3d.parameters(), lr=CFG.TRAIN4_LR)\n    scaler4 = torch.amp.GradScaler(torch_device)\n\n    test_results4 = train_with_sfk_oneCycleLR(sfk4, optimizer4, scaler4)\n    ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Plot training progression","metadata":{}},{"cell_type":"code","source":"if CFG.TRAIN4_EPOCHS>0:\n    fig, axs= plt.subplots(3,1)\n    \n    avg_losses = [tr0['avg_loss'] for tr0 in test_results4]\n    avg_metrics = [tr0['avg_metric'] for tr0 in test_results4]\n    \n    for tr0 in test_results:\n        axs[0].plot( avg_losses)\n        axs[0].set_ylabel(\"avg_loss\", size=\"x-small\")\n        axs[0].set_xlabel(\"epoch\", size=\"x-small\")\n        axs[1].plot(avg_metrics)\n        axs[1].set_ylabel(\"avg_metrics\", size=\"x-small\")\n        axs[1].set_xlabel(\"epoch\", size=\"x-small\")\n        axs[2].plot(global_lr_log)\n        axs[2].set_ylabel(\"lr\", size=\"x-small\") ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Save","metadata":{}},{"cell_type":"code","source":"if CFG.TRAIN4_EPOCHS>0:\n    model_fn = f\"{fname_stem}_step4.ptd\"\n    print(model_fn)\n    \n    model_dict = {'model_state_dict': model3d.state_dict(),\n                 'CFG':CFG}\n    torch.save(model_dict, model_fn)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# save best model","metadata":{}},{"cell_type":"code","source":"BEST_MODEL['best_metric']","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_fn = f\"{fname_stem}_best.ptd\"\nprint(model_fn)\n\nmodel_dict = {'model_state_dict': BEST_MODEL[\"model_stdict\"],\n             'CFG':CFG, \"best_metric\":  BEST_MODEL[\"best_metric\"]}\n\ntorch.save(model_dict, model_fn)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Do predictions on the remaining data volumes and get an estimate of score","metadata":{}},{"cell_type":"code","source":"def map_vol_function_by_blocking(func0, data3d, block_shape, margins_shape):\n    \"\"\"\n    Splits data3d to blocks of shape given, but with padding between them.\n    Then it applies function func0 to each of the blocks and collect data.\n    Resulting data is assumed to be the same shape as the input shape provided.\n    The resulting data is reassembled to a large volume, with padded regions discarded\n    except for margins.\n\n    func0 must be func0( numpy array with ndim=3 ) with no extra arguments.\n\n    If the function intended to be applied to the data volume requires arguments, please\n    use functools to generate a function that only requires a single 3-dim numpy array\n\n    An alternative to using this function is to use dask's map_overlap()\n    https://docs.dask.org/en/latest/array-overlap.html\n\n    Returns:\n        datares: result after applying func0 to the whole volume using the \n        blocking algorithm described\n    \"\"\"\n\n    shapedata=data3d.shape\n    #print(f\"map_vol_function_by_blocking() , data3d.shape:{data3d.shape} ,dtype:{data3d.dtype} block_shape:{block_shape}, margins_shape:{margins_shape}\")\n\n    #If for some reason the block shape is not big enough along one or more directions\n    # bl_step0 = np.array([ block_shape[i]-2*margins_shape[i] for i in range(3) ])\n    # bl_step = np.where( bl_step0<=0 , np.array(block_shape), bl_step0)\n\n    bl_step = np.array([ block_shape[i]-2*margins_shape[i] for i in range(3) ]) #default step\n    for i in range(3):\n        if bl_step[i]<0:\n            #bl_step[i]=block_shape[i]\n            raise ValueError(f\"margin with shape {margins_shape} too large  compared with block_shape {block_shape} at dim {i}. It should be block_shape[i]>2*margins_shape[i].\")\n        if block_shape[i]>=shapedata[i]:\n            bl_step[i]=shapedata[i]\n\n    #logging.debug(f\"bl_step:{bl_step}\")\n\n    datares = None #To collect results, it will be setup initially with correct dtype when first results arrive\n    b_continue=True\n\n    regions_plan = []\n\n    for iz0 in range(0,shapedata[0],bl_step[0]):\n\n        iz00=iz0\n        iz1 = iz0 + block_shape[0]\n        if iz1>shapedata[0]:\n            iz1 = shapedata[0]\n            iz00 = iz1 - block_shape[0]\n            if iz00<0: iz00=0\n        \n        for iy0 in range(0,shapedata[1], bl_step[1]):\n\n            iy00 = iy0\n            iy1 = iy0 + block_shape[1]\n            if iy1>shapedata[1]:\n                iy1 = shapedata[1]\n                iy00 = iy1 - block_shape[1]\n                if iy00<0: iy00=0\n\n            for ix0 in range(0,shapedata[2], bl_step[2]):\n                ix00 = ix0\n                ix1 = ix0 + block_shape[2]\n                if ix1>shapedata[2]:\n                    ix1 = shapedata[2]\n                    ix00 = ix1 - block_shape[2]\n                    if ix00<0: ix00=0\n\n                print(f\"BLOCK: New block, intended origin iz0,iy0,ix0 = {iz0},{iy0},{ix0} , use origin iz00,iy00,ix00 = {iz00},{iy00},{ix00} , end iz1,iy1,ix1 = {iz1},{iy1},{ix1}\")\n                \n                regions_plan.append((iz0,iy0,ix0, iz00,iy00,ix00, iz1,iy1,ix1))\n\n    print(\"Regions_plan (iz0,iy0,ix0, iz00,iy00,ix00, iz1,iy1,ix1)\")\n    print(f\"{regions_plan}\")\n\n    for reg0 in regions_plan:\n        iz0,iy0,ix0, iz00,iy00,ix00, iz1,iy1,ix1 = reg0\n\n        #Get the data block\n        datablock0 = data3d[iz00:iz1, iy00:iy1, ix00:ix1]\n        \n        print(f\"BLOCK: Start calculation with block [{iz00}:{iz1},{iy00}:{iy1},{ix00}:{ix1}]\")\n\n        #Do calculation with this datablock\n        data_res_block = func0(datablock0)\n        assert isinstance(data_res_block, np.ndarray)\n        assert data_res_block.ndim>=3\n\n        print(\"BLOCK: This block's calculation completed\")\n\n        if data_res_block is None:\n            raise ValueError( \"BLOCK: data_res_block is None. Check for errors. Stopping calculation\")\n        \n        #Store the datablock result, only the valid part\n        #unless it is the leftmost (first block) of the dimension given\n        jz0=0\n        jy0=0\n        jx0=0\n\n        #Crop the padded on the left side\n        if iz0 !=0 :\n            #jz0 += int( (block_shape[0] - bl_step[0]) / 2)\n            jz0 += margins_shape[0]\n        if iy0 !=0:\n            #jy0 += int( (block_shape[1] - bl_step[1]) / 2)\n            jy0 += margins_shape[1]\n        if ix0 !=0:\n            #jx0 += int( (block_shape[2] - bl_step[2]) / 2)\n            jx0 += margins_shape[2]\n        \n        if datares is None:\n            #Initialise\n            print(\"BLOCK: First block result initialises datares\")\n            if data_res_block.ndim==3:\n                datares = np.zeros( shapedata, dtype=data_res_block.dtype)\n            else:\n                datares = np.zeros( (data_res_block.shape[0],*shapedata), dtype=data_res_block.dtype)\n                \n\n        print(f\"BLOCK:Crop block result from origin jz0,jy0,jx0 = : {jz0},{jy0},{jx0} and copying to datares\")\n        if data_res_block.ndim==3:\n            datares[ iz00+jz0 : iz00+data_res_block.shape[0] , iy00+jy0 : iy00+data_res_block.shape[1] , ix00+jx0 : ix00+data_res_block.shape[2]] = data_res_block[jz0: , jy0: , jx0: ]\n        else:\n            datares[:, iz00+jz0 : iz00+data_res_block.shape[1] , iy00+jy0 : iy00+data_res_block.shape[2] , ix00+jx0 : ix00+data_res_block.shape[3]] = data_res_block[:, jz0: , jy0: , jx0: ]\n    print(\"BLOCK: Completed. Results should be in datares\")\n\n    return datares","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# helper function to infer volume with help of blocking function.\n# Function needs to accept single parameter bing volume\n\ndef infer_vol(data_vol):\n    data_norm = norm_func(data_vol)\n    data_t = torch.from_numpy(data_norm.copy()).cuda().unsqueeze(dim=0).unsqueeze(dim=0).float()\n    #print(data_t.shape)\n    model3d.eval()\n    with torch.no_grad():\n        Y = model3d(data_t)\n\n    torch.cuda.empty_cache()\n\n    # argmax to get label\n    y_argmax = torch.argmax(Y, dim=1)\n\n    res_np = y_argmax.squeeze().detach().cpu().numpy()\n    del(Y)\n    del(y_argmax)\n    del(data_t)\n    return res_np","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Get test data","metadata":{}},{"cell_type":"code","source":"# Get test data\n\ntest_data_mask_up = {}\n\nfor test_id0 in CFG.IDS_TO_TEST:\n    data0, mask0, up0 = get_data_and_sphere_masks(test_id0)\n    test_data_mask_up[test_id0] = (data0, mask0, up0)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Run predictions","metadata":{}},{"cell_type":"code","source":"preds_by_id = {}\n\nfor id, t in test_data_mask_up.items():\n    print(f\"*** id: {id} ***\")\n    data0,mask0,up0 = t\n\n    #pred = infer_vol(data0) #argmax is already applied\n    pred = map_vol_function_by_blocking(infer_vol, data0, (160,160,160), (14,14,14))\n    \n    print(f\"*** prediction complete of {id} , {pred.shape=} ***\")\n\n    preds_by_id[id]= {\n        \"data\":data0,\n        \"mask\":mask0,\n        \"used_particles\":up0,\n        \"prediction\": pred\n    }","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"preds_by_id.keys()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## scoring - create dataframes","metadata":{}},{"cell_type":"code","source":"# From https://www.kaggle.com/code/metric/czi-cryoet-84969\n# use distance_multiplier of 1 and beta=4\n# Make sure distances are multiplied by 10 (pixel to angstron)\n\nimport numpy as np\nimport pandas as pd\n\nfrom scipy.spatial import KDTree\n\n\nclass ParticipantVisibleError(Exception):\n    pass\n\n\ndef compute_metrics(reference_points, reference_radius, candidate_points):\n    num_reference_particles = len(reference_points)\n    num_candidate_particles = len(candidate_points)\n\n    if len(reference_points) == 0:\n        return 0, num_candidate_particles, 0\n\n    if len(candidate_points) == 0:\n        return 0, 0, num_reference_particles\n\n    ref_tree = KDTree(reference_points)\n    candidate_tree = KDTree(candidate_points)\n    raw_matches = candidate_tree.query_ball_tree(ref_tree, r=reference_radius)\n    matches_within_threshold = []\n    for match in raw_matches:\n        matches_within_threshold.extend(match)\n    # Prevent submitting multiple matches per particle.\n    # This won't be be strictly correct in the (extremely rare) case where true particles\n    # are very close to each other.\n    matches_within_threshold = set(matches_within_threshold)\n    tp = int(len(matches_within_threshold))\n    fp = int(num_candidate_particles - tp)\n    fn = int(num_reference_particles - tp)\n    return tp, fp, fn\n\n\ndef score(\n        solution: pd.DataFrame,\n        submission: pd.DataFrame,\n        row_id_column_name: str,\n        distance_multiplier: float,\n        beta: int) -> float:\n    '''\n    F_beta\n      - a true positive occurs when\n         - (a) the predicted location is within a threshold of the particle radius, and\n         - (b) the correct `particle_type` is specified\n      - raw results (TP, FP, FN) are aggregated across all experiments for each particle type\n      - f_beta is calculated for each particle type\n      - individual f_beta scores are weighted by particle type for final score\n    '''\n\n    particle_radius = {\n        'apo-ferritin': 60,\n        'beta-amylase': 65,\n        'beta-galactosidase': 90,\n        'ribosome': 150,\n        'thyroglobulin': 130,\n        'virus-like-particle': 135,\n    }\n\n    weights = {\n        'apo-ferritin': 1,\n        'beta-amylase': 0,\n        'beta-galactosidase': 2,\n        'ribosome': 1,\n        'thyroglobulin': 2,\n        'virus-like-particle': 1,\n    }\n\n    particle_radius = {k: v * distance_multiplier for k, v in particle_radius.items()}\n\n    # Filter submission to only contain experiments found in the solution split\n    split_experiments = set(solution['experiment'].unique())\n    submission = submission.loc[submission['experiment'].isin(split_experiments)]\n\n    # Only allow known particle types\n    if not set(submission['particle_type'].unique()).issubset(set(weights.keys())):\n        raise ParticipantVisibleError('Unrecognized `particle_type`.')\n\n    assert solution.duplicated(subset=['experiment', 'x', 'y', 'z']).sum() == 0\n    assert particle_radius.keys() == weights.keys()\n\n    results = {}\n    for particle_type in solution['particle_type'].unique():\n        results[particle_type] = {\n            'total_tp': 0,\n            'total_fp': 0,\n            'total_fn': 0,\n        }\n\n    for experiment in split_experiments:\n        for particle_type in solution['particle_type'].unique():\n            reference_radius = particle_radius[particle_type]\n            select = (solution['experiment'] == experiment) & (solution['particle_type'] == particle_type)\n            reference_points = solution.loc[select, ['x', 'y', 'z']].values\n\n            select = (submission['experiment'] == experiment) & (submission['particle_type'] == particle_type)\n            candidate_points = submission.loc[select, ['x', 'y', 'z']].values\n\n            if len(reference_points) == 0:\n                reference_points = np.array([])\n                reference_radius = 1\n\n            if len(candidate_points) == 0:\n                candidate_points = np.array([])\n\n            tp, fp, fn = compute_metrics(reference_points, reference_radius, candidate_points)\n\n            results[particle_type]['total_tp'] += tp\n            results[particle_type]['total_fp'] += fp\n            results[particle_type]['total_fn'] += fn\n\n    aggregate_fbeta = 0.0\n    for particle_type, totals in results.items():\n        tp = totals['total_tp']\n        fp = totals['total_fp']\n        fn = totals['total_fn']\n\n        precision = tp / (tp + fp) if tp + fp > 0 else 0\n        recall = tp / (tp + fn) if tp + fn > 0 else 0\n        fbeta = (1 + beta**2) * (precision * recall) / (beta**2 * precision + recall) if (precision + recall) > 0 else 0.0\n        aggregate_fbeta += fbeta * weights.get(particle_type, 1.0)\n\n    if weights:\n        aggregate_fbeta = aggregate_fbeta / sum(weights.values())\n    else:\n        aggregate_fbeta = aggregate_fbeta / len(results)\n    return aggregate_fbeta","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# get gnd truth dataframe\n\ngnd_experiments=[]\ngnd_particle_types=[]\ngnd_xs=[]\ngnd_ys=[]\ngnd_zs=[]\n\nfor id0 in CFG.IDS_TO_TEST:\n    print(f\"{id0=}\")\n    \n    gnd_dict_labels_id = get_dict_labels_by_id(id0)\n\n    for part_type, ilbl in PARTICLE_TO_CLASS.items():\n        \n        # ground truth particles\n        gnd_dict_labels_id_part = gnd_dict_labels_id[part_type]\n\n        for gz,gy,gx in gnd_dict_labels_id_part:\n            gnd_experiments.append(id0)\n            gnd_particle_types.append(part_type)\n            gnd_xs.append(gx * 10.0)\n            gnd_ys.append(gy * 10.0)\n            gnd_zs.append(gz * 10.0)\n            \ngnd_df = pandas.DataFrame( {\n    'experiment': gnd_experiments,\n    'particle_type': gnd_particle_types,\n    'z': gnd_zs,\n    'y': gnd_ys,\n    'x': gnd_xs,\n\n}, )\ngnd_df.index.name = 'id'","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"gnd_df.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_preds_pd_by_cc3d_params_shift(shift0=0.0):\n    experiments=[]\n    particle_types=[]\n    xs=[]\n    ys=[]\n    zs=[]\n    \n    for id0, v in preds_by_id.items():\n        print(f\"{id0=}\")\n        pred0 = v[\"prediction\"]\n    \n        for part_type, ilbl in PARTICLE_TO_CLASS.items():\n    \n            # For each label run conn components, regioprops and get centroids\n            pred_i_ilbl = (pred0==ilbl)\n    \n            if np.sum(pred_i_ilbl)>0:\n                pred_i_ilbl = cc3d.dust(pred_i_ilbl, threshold=10) # remove singel voxel predictions\n                pred_i_ilbl_cc = cc3d.connected_components(pred_i_ilbl)\n                \n                print(f\"{part_type} {ilbl=} has {pred_i_ilbl_cc.max()} regions\")\n    \n                #props_i_ilbl = regionprops( pred_i_ilbl_cc)\n                stats = cc3d.statistics(pred_i_ilbl_cc)\n                #print(stats.keys())\n                #centroids = stats['centroids'][1:,:] # Remove first one being the background\n                centroids = stats['centroids'][1:,:] *10.0\n    \n                #print(f\"centroids.shape : {centroids.shape} \") \n                size = centroids.shape[0]\n                \n                zs.extend(centroids[:,0]+ (shift0*10.0))\n                ys.extend(centroids[:,1]+ (shift0*10.0))\n                xs.extend(centroids[:,2]+ (shift0*10.0))\n                \n                experiments.extend([id0]*size)\n                particle_types.extend([part_type]*size)\n    \n    \n    subm_df = pandas.DataFrame( {\n        'experiment': experiments,\n        'particle_type': particle_types,\n        'z': zs,\n        'y': ys,\n        'x': xs,\n    }, )\n    subm_df.index.name = 'id'\n\n    return subm_df\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"subm_df = get_preds_pd_by_cc3d_params_shift(CFG.CC3D_CENTROID_SHIFT_PIX)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"subm_df.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## score from DFs","metadata":{}},{"cell_type":"code","source":"score0 = score(gnd_df, subm_df, \"id\",1.0, 4.0)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Test score: {score0}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}