{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":10446384,"sourceType":"datasetVersion","datasetId":6466271},{"sourceId":10664738,"sourceType":"datasetVersion","datasetId":6555362},{"sourceId":210375552,"sourceType":"kernelVersion"},{"sourceId":210540318,"sourceType":"kernelVersion"}],"dockerImageVersionId":30823,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Inference for wmcsfb_grn(mednext with fourier spherical bessel +grn)\n\ntrained in a seperate notebook\n\nModified to calculate patches but collecting probabilities and average them on the edges.","metadata":{}},{"cell_type":"code","source":"MODEL_FN=\"/kaggle/input/notebook-output-wmcsfb-grn-train/2025-02-04-1806_wmcsfb_grn_step1.ptd\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:53:24.919321Z","iopub.execute_input":"2025-01-30T15:53:24.919601Z","iopub.status.idle":"2025-01-30T15:53:24.933287Z","shell.execute_reply.started":"2025-01-30T15:53:24.919572Z","shell.execute_reply":"2025-01-30T15:53:24.931758Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"USE_DBSCAN_SOLUTION = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:53:24.933915Z","iopub.execute_input":"2025-01-30T15:53:24.934168Z","iopub.status.idle":"2025-01-30T15:53:24.953481Z","shell.execute_reply.started":"2025-01-30T15:53:24.934146Z","shell.execute_reply":"2025-01-30T15:53:24.952569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"USE_AGGREGATE_PATCHES_WITH_PROBS_MEAN = True","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Installs","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-01-30T15:53:27.749343Z","iopub.execute_input":"2025-01-30T15:53:27.749610Z","iopub.status.idle":"2025-01-30T15:53:34.603853Z","shell.execute_reply.started":"2025-01-30T15:53:27.749590Z","shell.execute_reply":"2025-01-30T15:53:34.602796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip -q install /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-01-30T15:53:34.605067Z","iopub.execute_input":"2025-01-30T15:53:34.605302Z","iopub.status.idle":"2025-01-30T15:53:37.887768Z","shell.execute_reply.started":"2025-01-30T15:53:34.605283Z","shell.execute_reply":"2025-01-30T15:53:37.886834Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# WMCSFB compact code","metadata":{}},{"cell_type":"code","source":"SPH_BESSEL_FILE_LOCATION = \"/kaggle/input/wmcsfb-spherical-bessel/spherical_bessel.npy\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:53:37.889210Z","iopub.execute_input":"2025-01-30T15:53:37.889440Z","iopub.status.idle":"2025-01-30T15:53:37.893088Z","shell.execute_reply.started":"2025-01-30T15:53:37.889421Z","shell.execute_reply":"2025-01-30T15:53:37.892307Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:53:40.844716Z","iopub.execute_input":"2025-01-30T15:53:40.845026Z","iopub.status.idle":"2025-01-30T15:53:44.271290Z","shell.execute_reply.started":"2025-01-30T15:53:40.844998Z","shell.execute_reply":"2025-01-30T15:53:44.270508Z"}},"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\n#import 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\n#from torch_lr_finder import LRFinder\n\nimport cc3d\nimport pandas\nfrom dataclasses import dataclass","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:53:50.852198Z","iopub.execute_input":"2025-01-30T15:53:50.852599Z","iopub.status.idle":"2025-01-30T15:53:51.703392Z","shell.execute_reply.started":"2025-01-30T15:53:50.852574Z","shell.execute_reply":"2025-01-30T15:53:51.702473Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"config","metadata":{}},{"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'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:53:52.372314Z","iopub.execute_input":"2025-01-30T15:53:52.373003Z","iopub.status.idle":"2025-01-30T15:53:52.376920Z","shell.execute_reply.started":"2025-01-30T15:53:52.372957Z","shell.execute_reply":"2025-01-30T15:53:52.375955Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:53:54.420762Z","iopub.execute_input":"2025-01-30T15:53:54.421076Z","iopub.status.idle":"2025-01-30T15:53:54.431041Z","shell.execute_reply.started":"2025-01-30T15:53:54.421050Z","shell.execute_reply":"2025-01-30T15:53:54.430290Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_folder_path = Path(TEST_BASE)\ntest_runs = [ d.parts[-1] for d in test_folder_path.iterdir()]\ntest_runs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:53:55.100447Z","iopub.execute_input":"2025-01-30T15:53:55.100773Z","iopub.status.idle":"2025-01-30T15:53:55.110367Z","shell.execute_reply.started":"2025-01-30T15:53:55.100739Z","shell.execute_reply":"2025-01-30T15:53:55.109479Z"}},"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    \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\n\ndef get_test_data_by_id(id):\n    \n    exp_webpzarr = Path(TEST_BASE)/id/ \"VoxelSpacing10.000/denoised.zarr\"\n    \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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:53:55.309258Z","iopub.execute_input":"2025-01-30T15:53:55.309565Z","iopub.status.idle":"2025-01-30T15:53:55.314926Z","shell.execute_reply.started":"2025-01-30T15:53:55.309540Z","shell.execute_reply":"2025-01-30T15:53:55.313947Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:53:56.076949Z","iopub.execute_input":"2025-01-30T15:53:56.077238Z","iopub.status.idle":"2025-01-30T15:53:56.083354Z","shell.execute_reply.started":"2025-01-30T15:53:56.077216Z","shell.execute_reply":"2025-01-30T15:53:56.082368Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Setup model","metadata":{}},{"cell_type":"markdown","source":"Use of CFG will be deprecated soon in favour of SimpleNamepace","metadata":{}},{"cell_type":"code","source":"# # dummy to load CFG. TOBE REMOVED\n# @dataclass\n# class CFG:\n    \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_6_4','TS_73_6', 'TS_99_9'] #5/7 vols, leaving others for testing\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#     GND_MASK_RADIUS_SCALE_FACTOR = 3/4 # new\n    \n#     TRAIN_AUGM_ROTS = False\n#     TRAIN_AUGM_FLIPS = True\n    \n#     #CLASS_WEIGHTS = [0.05, 1,1,1,1,1]\n#     TRAIN_CUBE_SIZE = 64\n#     TRAIN_BATCH_SIZE= 6\n#     TRAIN_EPOCHS = 8\n#     TRAIN_K_FOLDS_PER_EPOCH = 4\n#     TRAIN_LR = 1e-4\n#     TRAIN_MAX_LR = 1e-3\n#     #PARTICLE_TO_TRAIN = 'virus-like-particle' #test this particle first\n\n#     # second part of the training, using stepLR\n#     TRAIN2_EPOCHS = 4\n#     TRAIN2_INIT_LR = 1e-4\n#     TRAIN2_GAMMA = 0.31622 # 1/sqrt(10) every 2 steps to reduce by 10\n    \n#     CC3D_CENTROID_SHIFT_PIX = -0.5\n\n\n#     \"\"\"\n#     model3d = create_mednext_v1(\n#     num_input_channels = 1,\n#     num_classes = 6, # including background\n#     model_id = 'S',\n#     kernel_size = 3).to(torch_device)\n#     \"\"\"\n#     MEDNEXT_V1_INP_CHN = 1\n#     MEDNEXT_V1_NUM_CLASSES = 6 # including background\n#     MEDNEXT_V1_ID = 'S'\n#     MEDNEXT_V1_KERNEL_SIZE = 3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:53:58.096200Z","iopub.execute_input":"2025-01-30T15:53:58.096509Z","iopub.status.idle":"2025-01-30T15:53:58.100407Z","shell.execute_reply.started":"2025-01-30T15:53:58.096481Z","shell.execute_reply":"2025-01-30T15:53:58.099469Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# read file\n#model_dict = torch.load('/kaggle/input/2024-12-30-2134-mednext/2024-12-30-2134_mednext_step1.ptd')\n#model_dict = torch.load('/kaggle/input/2024-12-30-2134-mednext/2024-12-30-2134_mednext_step2.ptd') #0.470\n\n#model_dict = torch.load('/kaggle/input/mednext-v1-train/2025-01-01-1615_mednext_step1.ptd') # 0.677\n#model_dict = torch.load('/kaggle/input/2025-01-01-1615-mednext-step2/2025-01-01-1615_mednext_step2.ptd')\n\n# model_dict = torch.load('/kaggle/input/mednext-v1-train/2025-01-02-1349_mednext_step1.ptd')\n#model_dict = torch.load('/kaggle/input/mednext-v1-train/2025-01-02-1349_mednext_step2.ptd')\n\nmodel_dict = torch.load(MODEL_FN)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:53:58.326531Z","iopub.execute_input":"2025-01-30T15:53:58.326832Z","iopub.status.idle":"2025-01-30T15:53:58.882921Z","shell.execute_reply.started":"2025-01-30T15:53:58.326811Z","shell.execute_reply":"2025-01-30T15:53:58.882010Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_dict.keys()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:54:00.204409Z","iopub.execute_input":"2025-01-30T15:54:00.204759Z","iopub.status.idle":"2025-01-30T15:54:00.209381Z","shell.execute_reply.started":"2025-01-30T15:54:00.204728Z","shell.execute_reply":"2025-01-30T15:54:00.208649Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CFG = model_dict['CFG']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:54:00.899286Z","iopub.execute_input":"2025-01-30T15:54:00.899586Z","iopub.status.idle":"2025-01-30T15:54:00.903224Z","shell.execute_reply.started":"2025-01-30T15:54:00.899562Z","shell.execute_reply":"2025-01-30T15:54:00.902229Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CFG","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:54:01.852022Z","iopub.execute_input":"2025-01-30T15:54:01.852303Z","iopub.status.idle":"2025-01-30T15:54:01.858307Z","shell.execute_reply.started":"2025-01-30T15:54:01.852281Z","shell.execute_reply":"2025-01-30T15:54:01.857431Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:54:02.227862Z","iopub.execute_input":"2025-01-30T15:54:02.228171Z","iopub.status.idle":"2025-01-30T15:54:04.289211Z","shell.execute_reply.started":"2025-01-30T15:54:02.228148Z","shell.execute_reply":"2025-01-30T15:54:04.288261Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"load from last checkpoint","metadata":{}},{"cell_type":"code","source":"model3d.load_state_dict(model_dict['model_state_dict'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:54:10.076629Z","iopub.execute_input":"2025-01-30T15:54:10.077066Z","iopub.status.idle":"2025-01-30T15:54:10.091807Z","shell.execute_reply.started":"2025-01-30T15:54:10.077034Z","shell.execute_reply":"2025-01-30T15:54:10.090783Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Normalisation function","metadata":{}},{"cell_type":"code","source":"CFG.DATA_NORM_FUNC","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:54:11.627162Z","iopub.execute_input":"2025-01-30T15:54:11.627493Z","iopub.status.idle":"2025-01-30T15:54:11.633049Z","shell.execute_reply.started":"2025-01-30T15:54:11.627448Z","shell.execute_reply":"2025-01-30T15:54:11.632139Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:54:11.995043Z","iopub.execute_input":"2025-01-30T15:54:11.995323Z","iopub.status.idle":"2025-01-30T15:54:11.999609Z","shell.execute_reply.started":"2025-01-30T15:54:11.995302Z","shell.execute_reply":"2025-01-30T15:54:11.998392Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"norm_func","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:54:13.403192Z","iopub.execute_input":"2025-01-30T15:54:13.403513Z","iopub.status.idle":"2025-01-30T15:54:13.408492Z","shell.execute_reply.started":"2025-01-30T15:54:13.403483Z","shell.execute_reply":"2025-01-30T15:54:13.407629Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Particles","metadata":{}},{"cell_type":"code","source":"PARTICLE_TO_CLASS={ part0:i+1 for i,part0 in enumerate(CFG.PARTICLES)}\nPARTICLE_TO_CLASS","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:54:14.699171Z","iopub.execute_input":"2025-01-30T15:54:14.699493Z","iopub.status.idle":"2025-01-30T15:54:14.704803Z","shell.execute_reply.started":"2025-01-30T15:54:14.699466Z","shell.execute_reply":"2025-01-30T15:54:14.703997Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Do predictions","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        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(\"All blocks completed.\")\n\n    return datares","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:54:16.039431Z","iopub.execute_input":"2025-01-30T15:54:16.039783Z","iopub.status.idle":"2025-01-30T15:54:16.051215Z","shell.execute_reply.started":"2025-01-30T15:54:16.039753Z","shell.execute_reply":"2025-01-30T15:54:16.050076Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def map_vol_function_by_blocking_prob_mean(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    It assumes data is one-hot coded probabilities.\n    Resulting data is assumed to be the same shape as the input shape provided.\n    The resulting data is reassembled by taking the average probability.\n    In padded regions these can be probabilities of up to 3 regions.\n\n    func0 must be a function ( 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    res_probs=None\n    res_count=None\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        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 res_probs is None:\n            #Initialise\n            #print(\"BLOCK: First block result initialises datares\")\n            if data_res_block.ndim==3:\n                res_probs = np.zeros( shapedata, dtype=data_res_block.dtype)\n                res_count = np.zeros( shapedata, dtype=np.uint8)\n            else:\n                # assume dim 0 to be the one-hotted class\n                res_probs = np.zeros( (data_res_block.shape[0],*shapedata), dtype=data_res_block.dtype)\n                res_count = np.zeros( shapedata, dtype=np.uint8)\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            res_probs[ 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            res_count[ iz00+jz0 : iz00+data_res_block.shape[0] , iy00+jy0 : iy00+data_res_block.shape[1] , ix00+jx0 : ix00+data_res_block.shape[2] ] +=1\n        else:\n            res_probs[:, 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            res_count[ iz00+jz0 : iz00+data_res_block.shape[1] , iy00+jy0 : iy00+data_res_block.shape[2] , ix00+jx0 : ix00+data_res_block.shape[3] ] +=1\n    \n    print(\"All blocks completed. Calculating average and argmax\")\n\n    expanded_counts = res_count[np.newaxis, :, :]\n\n    res_mean = res_probs/ expanded_counts\n\n    res_argmax = np.argmax(res_mean, axis=0)\n\n    return res_argmax","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T16:03:32.975869Z","iopub.execute_input":"2025-01-30T16:03:32.976193Z","iopub.status.idle":"2025-01-30T16:03:32.989009Z","shell.execute_reply.started":"2025-01-30T16:03:32.976171Z","shell.execute_reply":"2025-01-30T16:03:32.988006Z"}},"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    global norm_func\n    \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,"execution":{"iopub.status.busy":"2025-01-30T15:59:35.750571Z","iopub.execute_input":"2025-01-30T15:59:35.750925Z","iopub.status.idle":"2025-01-30T15:59:35.755909Z","shell.execute_reply.started":"2025-01-30T15:59:35.750899Z","shell.execute_reply":"2025-01-30T15:59:35.755041Z"}},"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_probs(data_vol):\n    global norm_func\n    \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        Y_softmax = torch.nn.functional.softmax(Y,dim=1)\n\n    torch.cuda.empty_cache()\n\n    res_np = Y.squeeze().detach().cpu().numpy()\n    del(Y)\n    del(data_t)\n    return res_np","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:59:35.934166Z","iopub.execute_input":"2025-01-30T15:59:35.934444Z","iopub.status.idle":"2025-01-30T15:59:35.939386Z","shell.execute_reply.started":"2025-01-30T15:59:35.934424Z","shell.execute_reply":"2025-01-30T15:59:35.938472Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Routine to run the predictions","metadata":{}},{"cell_type":"code","source":"shift0=CFG.CC3D_CENTROID_SHIFT_PIX\n#shift0=-1.5 #override\nshift0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T15:59:36.879396Z","iopub.execute_input":"2025-01-30T15:59:36.879739Z","iopub.status.idle":"2025-01-30T15:59:36.884518Z","shell.execute_reply.started":"2025-01-30T15:59:36.879712Z","shell.execute_reply":"2025-01-30T15:59:36.883862Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"experiments=[]\nparticle_types=[]\nxs=[]\nys=[]\nzs=[]\n\nfor id0 in test_runs:\n    time0 = time.perf_counter()\n    \n    data0 = get_test_data_by_id(id0)\n    print(f\"*** {id0=} , {data0.shape=} ***\")\n\n    # mednext only works with sizes [64, 80, 96, 112, 128] , and 160 I found out\n    #pred = map_vol_function_by_blocking(infer_vol, data0, (128,128,128), (14,14,14))\n    #pred = map_vol_function_by_blocking(infer_vol, data0, (160,160,160), (14,14,14))\n    #pred = map_vol_function_by_blocking(infer_vol, data0, (170,170,170), (14,14,14)) # dont work 200, 192, 180, 170\n\n    if USE_AGGREGATE_PATCHES_WITH_PROBS_MEAN:\n        pred = map_vol_function_by_blocking_prob_mean(infer_vol_probs, data0, (160,160,160), (14,14,14))\n    else:\n        pred = map_vol_function_by_blocking(infer_vol, data0, (160,160,160), (14,14,14))\n    \n    print(f\"*** prediction complete of {id0} , {pred.shape=} ***\")\n\n    pred0 = pred\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    time1= time.perf_counter()\n    print(f\"id {id0} took {time1-time0} to complete\")\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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T16:03:37.700324Z","iopub.execute_input":"2025-01-30T16:03:37.700653Z","iopub.status.idle":"2025-01-30T16:05:58.989806Z","shell.execute_reply.started":"2025-01-30T16:03:37.700625Z","shell.execute_reply":"2025-01-30T16:05:58.989072Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"3 volumes in 1:53 minutes. Can be done","metadata":{}},{"cell_type":"code","source":"subm_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T16:06:04.924744Z","iopub.execute_input":"2025-01-30T16:06:04.925060Z","iopub.status.idle":"2025-01-30T16:06:04.940781Z","shell.execute_reply.started":"2025-01-30T16:06:04.925033Z","shell.execute_reply":"2025-01-30T16:06:04.939915Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Save","metadata":{}},{"cell_type":"code","source":"subm_df.to_csv(\"submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T16:06:07.588701Z","iopub.execute_input":"2025-01-30T16:06:07.588999Z","iopub.status.idle":"2025-01-30T16:06:07.601127Z","shell.execute_reply.started":"2025-01-30T16:06:07.588977Z","shell.execute_reply":"2025-01-30T16:06:07.600435Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## scoring - create dataframes","metadata":{}},{"cell_type":"code","source":"# IS_SUBMISSION = bool(os.getenv(\"KAGGLE_IS_COMPETITION_RERUN\"))\n# IS_SUBMISSION\n\nIS_SUBMISSION=False\nif len(test_runs)>3:\n    IS_SUBMISSION=True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T16:06:08.702210Z","iopub.execute_input":"2025-01-30T16:06:08.702496Z","iopub.status.idle":"2025-01-30T16:06:08.706191Z","shell.execute_reply.started":"2025-01-30T16:06:08.702474Z","shell.execute_reply":"2025-01-30T16:06:08.705318Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IS_SUBMISSION","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T16:06:10.572379Z","iopub.execute_input":"2025-01-30T16:06:10.572714Z","iopub.status.idle":"2025-01-30T16:06:10.577975Z","shell.execute_reply.started":"2025-01-30T16:06:10.572688Z","shell.execute_reply":"2025-01-30T16:06:10.577145Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if not IS_SUBMISSION:\n    \n    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\n\n\n    # 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    \n    import numpy as np\n    import pandas as pd\n    \n    from scipy.spatial import KDTree\n    \n    \n    class ParticipantVisibleError(Exception):\n        pass\n    \n    \n    def 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    \n    def 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\n\n    \n    # get gnd truth dataframe\n    gnd_experiments=[]\n    gnd_particle_types=[]\n    gnd_xs=[]\n    gnd_ys=[]\n    gnd_zs=[]\n    \n    for id0 in test_runs:\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                \n    gnd_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    }, )\n    gnd_df.index.name = 'id'\n\n    score0 = score(gnd_df, subm_df, \"id\",1.0, 4.0)\n    print(f\"Test score: {score0:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T16:06:10.737346Z","iopub.execute_input":"2025-01-30T16:06:10.737688Z","iopub.status.idle":"2025-01-30T16:06:10.878963Z","shell.execute_reply.started":"2025-01-30T16:06:10.737636Z","shell.execute_reply":"2025-01-30T16:06:10.878174Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# DBSCAN\n\nWith edits from https://www.kaggle.com/code/linheshen/esemble-2d-and-3d","metadata":{}},{"cell_type":"code","source":"from sklearn.cluster import DBSCAN\n\ndf = subm_df\n\nfinal = []\nfor pidx, p in enumerate(CFG.PARTICLES):\n    \n    pdf = df[df['particle_type'] == p].reset_index(drop=True)\n    p_rad = CFG.PARTICLE_RADIUS[p]\n    \n    grouped = pdf.groupby(['experiment'])\n\n    # print(grouped)\n    \n    for exp, group in grouped:\n        group = group.reset_index(drop=True)\n        \n        coords = group[['x', 'y', 'z']].values\n        db = DBSCAN(eps=p_rad, min_samples=2, metric='euclidean').fit(coords)\n        labels = db.labels_\n        \n        group['cluster'] = labels\n        \n        for cluster_id in np.unique(labels):\n            if cluster_id == -1:\n                continue  # 跳过噪声点\n            \n            cluster_points = group[group['cluster'] == cluster_id]\n            \n            avg_x = cluster_points['x'].mean()\n            avg_y = cluster_points['y'].mean()\n            avg_z = cluster_points['z'].mean()\n            \n            group.loc[group['cluster'] == cluster_id, ['x', 'y', 'z']] = avg_x, avg_y, avg_z\n            group = group.drop_duplicates(subset=['x', 'y', 'z'])\n        final.append(group)\n\ndf_dbscan = pd.concat(final, ignore_index=True)\ndf_dbscan = df_dbscan.drop(columns=['cluster'])\n\ndf_dbscan = df_dbscan.sort_values(by=['experiment', 'particle_type']).reset_index(drop=True)\n\ndf_dbscan['id'] = np.arange(0, len(df_dbscan))\n\nif not IS_SUBMISSION:\n    score_dbscan = score(gnd_df, df_dbscan, \"id\",1.0, 4.0)\n    print(f\"Test score: {score0:.4f}\")\n    print(f\"Test score_dbscan: {score_dbscan:.4f}\")\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T16:06:15.149417Z","iopub.execute_input":"2025-01-30T16:06:15.149731Z","iopub.status.idle":"2025-01-30T16:06:15.643714Z","shell.execute_reply.started":"2025-01-30T16:06:15.149707Z","shell.execute_reply":"2025-01-30T16:06:15.642764Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if USE_DBSCAN_SOLUTION:\n     df_dbscan.to_csv('submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T16:06:17.773011Z","iopub.execute_input":"2025-01-30T16:06:17.773344Z","iopub.status.idle":"2025-01-30T16:06:17.782194Z","shell.execute_reply.started":"2025-01-30T16:06:17.773311Z","shell.execute_reply":"2025-01-30T16:06:17.781562Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}