{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":9840041,"sourceType":"datasetVersion","datasetId":6036295},{"sourceId":9862305,"sourceType":"datasetVersion","datasetId":6052780},{"sourceId":9867543,"sourceType":"datasetVersion","datasetId":6040935},{"sourceId":11882450,"sourceType":"datasetVersion","datasetId":7468025},{"sourceId":12152328,"sourceType":"datasetVersion","datasetId":7653375},{"sourceId":224409375,"sourceType":"kernelVersion"},{"sourceId":389056,"sourceType":"modelInstanceVersion","modelInstanceId":320182,"modelId":339846},{"sourceId":389057,"sourceType":"modelInstanceVersion","modelInstanceId":320179,"modelId":339846},{"sourceId":400450,"sourceType":"modelInstanceVersion","modelInstanceId":327693,"modelId":348562},{"sourceId":400457,"sourceType":"modelInstanceVersion","modelInstanceId":327697,"modelId":348562},{"sourceId":465234,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":375372,"modelId":396196}],"dockerImageVersionId":30840,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Inference Configuration","metadata":{}},{"cell_type":"code","source":"# 3DUnet or 3DUnetAG or 3DUnetSE or 3DUnetCBAM\nmodel_architecture = \"3DUnetSE\"\nmodel_path = \"/kaggle/input/retrain-server/pytorch/default/1/best_metric_model_se_128_96_retrain3.pth\"\nmodel_channels = (16, 32, 64, 128)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:27.047585Z","iopub.execute_input":"2025-05-20T10:58:27.047988Z","iopub.status.idle":"2025-05-20T10:58:27.051855Z","shell.execute_reply.started":"2025-05-20T10:58:27.047955Z","shell.execute_reply":"2025-05-20T10:58:27.051155Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dependencies","metadata":{}},{"cell_type":"code","source":"! pip install -q --no-index --find-links /kaggle/input/copick-utils copick-utils","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:27.055189Z","iopub.execute_input":"2025-05-20T10:58:27.055398Z","iopub.status.idle":"2025-05-20T10:58:30.528837Z","shell.execute_reply.started":"2025-05-20T10:58:27.055379Z","shell.execute_reply":"2025-05-20T10:58:30.527619Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip install -q --no-index --find-links /kaggle/input/extra-dependency matplotlib tqdm copick","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:30.530286Z","iopub.execute_input":"2025-05-20T10:58:30.530613Z","iopub.status.idle":"2025-05-20T10:58:34.273453Z","shell.execute_reply.started":"2025-05-20T10:58:30.530578Z","shell.execute_reply":"2025-05-20T10:58:34.272501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip install -q --no-index --find-links /kaggle/input/extra-dependency monai-weekly[mlflow]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:34.275192Z","iopub.execute_input":"2025-05-20T10:58:34.275438Z","iopub.status.idle":"2025-05-20T10:58:37.901272Z","shell.execute_reply.started":"2025-05-20T10:58:34.275417Z","shell.execute_reply":"2025-05-20T10:58:37.900263Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip install -q --no-index --find-links /kaggle/input/extra-dependency connected-components-3d lightning","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:37.902719Z","iopub.execute_input":"2025-05-20T10:58:37.902978Z","iopub.status.idle":"2025-05-20T10:58:41.455602Z","shell.execute_reply.started":"2025-05-20T10:58:37.902956Z","shell.execute_reply":"2025-05-20T10:58:41.454451Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Make a copick project\nimport os\nimport shutil\n\nconfig_blob = \"\"\"{\n    \"name\": \"czii_cryoet_mlchallenge_2024\",\n    \"description\": \"2024 CZII CryoET ML Challenge training data.\",\n    \"version\": \"1.0.0\",\n\n    \"pickable_objects\": [\n        {\n            \"name\": \"apo-ferritin\",\n            \"is_particle\": true,\n            \"pdb_id\": \"4V1W\",\n            \"label\": 1,\n            \"color\": [  0, 117, 220, 128],\n            \"radius\": 60,\n            \"map_threshold\": 0.0418\n        },\n        {\n            \"name\": \"beta-amylase\",\n            \"is_particle\": true,\n            \"pdb_id\": \"1FA2\",\n            \"label\": 2,\n            \"color\": [153,  63,   0, 128],\n            \"radius\": 65,\n            \"map_threshold\": 0.035\n        },\n        {\n            \"name\": \"beta-galactosidase\",\n            \"is_particle\": true,\n            \"pdb_id\": \"6X1Q\",\n            \"label\": 3,\n            \"color\": [ 76,   0,  92, 128],\n            \"radius\": 90,\n            \"map_threshold\": 0.0578\n        },\n        {\n            \"name\": \"ribosome\",\n            \"is_particle\": true,\n            \"pdb_id\": \"6EK0\",\n            \"label\": 4,\n            \"color\": [  0,  92,  49, 128],\n            \"radius\": 150,\n            \"map_threshold\": 0.0374\n        },\n        {\n            \"name\": \"thyroglobulin\",\n            \"is_particle\": true,\n            \"pdb_id\": \"6SCJ\",\n            \"label\": 5,\n            \"color\": [ 43, 206,  72, 128],\n            \"radius\": 130,\n            \"map_threshold\": 0.0278\n        },\n        {\n            \"name\": \"virus-like-particle\",\n            \"is_particle\": true,\n            \"label\": 6,\n            \"color\": [255, 204, 153, 128],\n            \"radius\": 135,\n            \"map_threshold\": 0.201\n        }\n    ],\n\n    \"overlay_root\": \"/kaggle/working/overlay\",\n\n    \"overlay_fs_args\": {\n        \"auto_mkdir\": true\n    },\n\n    \"static_root\": \"/kaggle/input/czii-test-9-update/test/static\"\n}\"\"\"\n\ncopick_config_path = \"/kaggle/working/copick.config\"\noutput_overlay = \"/kaggle/working/overlay\"\n\nwith open(copick_config_path, \"w\") as f:\n    f.write(config_blob)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:41.456846Z","iopub.execute_input":"2025-05-20T10:58:41.457158Z","iopub.status.idle":"2025-05-20T10:58:41.46223Z","shell.execute_reply.started":"2025-05-20T10:58:41.457134Z","shell.execute_reply":"2025-05-20T10:58:41.461509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nfrom pathlib import Path\nimport torch\nimport torchinfo\nimport zarr, copick\nfrom tqdm import tqdm\nfrom monai.data import DataLoader, Dataset, CacheDataset, decollate_batch\nfrom monai.transforms import (\n    Compose,\n    EnsureChannelFirstd, \n    Orientationd,  \n    NormalizeIntensityd,\n)\nfrom monai.networks.nets import UNet\nfrom monai.metrics import DiceMetric, ConfusionMatrixMetric\nimport mlflow\nimport mlflow.pytorch\nimport torch.nn as nn","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:41.463023Z","iopub.execute_input":"2025-05-20T10:58:41.463218Z","iopub.status.idle":"2025-05-20T10:58:41.478157Z","shell.execute_reply.started":"2025-05-20T10:58:41.463194Z","shell.execute_reply":"2025-05-20T10:58:41.477322Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset Preparation","metadata":{}},{"cell_type":"markdown","source":"## Model setup","metadata":{}},{"cell_type":"code","source":"from __future__ import annotations\n\nfrom collections.abc import Sequence\n\nimport torch\nimport torch.nn as nn\n\nfrom monai.networks.blocks.convolutions import Convolution\nfrom monai.networks.layers.factories import Norm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:41.479119Z","iopub.execute_input":"2025-05-20T10:58:41.479409Z","iopub.status.idle":"2025-05-20T10:58:41.489656Z","shell.execute_reply.started":"2025-05-20T10:58:41.479381Z","shell.execute_reply":"2025-05-20T10:58:41.48896Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Attention Gate","metadata":{}},{"cell_type":"code","source":"if model_architecture == \"3DUnetAG\":\n    # Copyright (c) MONAI Consortium\n    # Licensed under the Apache License, Version 2.0 (the \"License\");\n    # you may not use this file except in compliance with the License.\n    # You may obtain a copy of the License at\n    #     http://www.apache.org/licenses/LICENSE-2.0\n    # Unless required by applicable law or agreed to in writing, software\n    # distributed under the License is distributed on an \"AS IS\" BASIS,\n    # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n    # See the License for the specific language governing permissions and\n    # limitations under the License.\n    \n    __all__ = [\"AttentionUnet\"]\n    \n    \n    class ConvBlock(nn.Module):\n    \n        def __init__(\n            self,\n            spatial_dims: int,\n            in_channels: int,\n            out_channels: int,\n            kernel_size: Sequence[int] | int = 3,\n            strides: int = 1,\n            dropout=0.0,\n        ):\n            super().__init__()\n            layers = [\n                Convolution(\n                    spatial_dims=spatial_dims,\n                    in_channels=in_channels,\n                    out_channels=out_channels,\n                    kernel_size=kernel_size,\n                    strides=strides,\n                    padding=None,\n                    adn_ordering=\"NDA\",\n                    act=\"relu\",\n                    norm=Norm.BATCH,\n                    dropout=dropout,\n                ),\n                Convolution(\n                    spatial_dims=spatial_dims,\n                    in_channels=out_channels,\n                    out_channels=out_channels,\n                    kernel_size=kernel_size,\n                    strides=1,\n                    padding=None,\n                    adn_ordering=\"NDA\",\n                    act=\"relu\",\n                    norm=Norm.BATCH,\n                    dropout=dropout,\n                ),\n            ]\n            self.conv = nn.Sequential(*layers)\n    \n        def forward(self, x: torch.Tensor) -> torch.Tensor:\n            x_c: torch.Tensor = self.conv(x)\n            return x_c\n    \n    \n    class UpConv(nn.Module):\n    \n        def __init__(self, spatial_dims: int, in_channels: int, out_channels: int, kernel_size=3, strides=2, dropout=0.0):\n            super().__init__()\n            self.up = Convolution(\n                spatial_dims,\n                in_channels,\n                out_channels,\n                strides=strides,\n                kernel_size=kernel_size,\n                act=\"relu\",\n                adn_ordering=\"NDA\",\n                norm=Norm.BATCH,\n                dropout=dropout,\n                is_transposed=True,\n            )\n    \n        def forward(self, x: torch.Tensor) -> torch.Tensor:\n            x_u: torch.Tensor = self.up(x)\n            return x_u\n    \n    \n    class AttentionBlock(nn.Module):\n    \n        def __init__(self, spatial_dims: int, f_int: int, f_g: int, f_l: int, dropout=0.0):\n            super().__init__()\n            self.W_g = nn.Sequential(\n                Convolution(\n                    spatial_dims=spatial_dims,\n                    in_channels=f_g,\n                    out_channels=f_int,\n                    kernel_size=1,\n                    strides=1,\n                    padding=0,\n                    dropout=dropout,\n                    conv_only=True,\n                ),\n                Norm[Norm.BATCH, spatial_dims](f_int),\n            )\n    \n            self.W_x = nn.Sequential(\n                Convolution(\n                    spatial_dims=spatial_dims,\n                    in_channels=f_l,\n                    out_channels=f_int,\n                    kernel_size=1,\n                    strides=1,\n                    padding=0,\n                    dropout=dropout,\n                    conv_only=True,\n                ),\n                Norm[Norm.BATCH, spatial_dims](f_int),\n            )\n    \n            self.psi = nn.Sequential(\n                Convolution(\n                    spatial_dims=spatial_dims,\n                    in_channels=f_int,\n                    out_channels=1,\n                    kernel_size=1,\n                    strides=1,\n                    padding=0,\n                    dropout=dropout,\n                    conv_only=True,\n                ),\n                Norm[Norm.BATCH, spatial_dims](1),\n                nn.Sigmoid(),\n            )\n    \n            self.relu = nn.ReLU()\n    \n        def forward(self, g: torch.Tensor, x: torch.Tensor) -> torch.Tensor:\n            g1 = self.W_g(g)\n            x1 = self.W_x(x)\n            psi: torch.Tensor = self.relu(g1 + x1)\n            psi = self.psi(psi)\n    \n            return x * psi\n    \n    \n    class AttentionLayer(nn.Module):\n    \n        def __init__(\n            self,\n            spatial_dims: int,\n            in_channels: int,\n            out_channels: int,\n            submodule: nn.Module,\n            up_kernel_size=3,\n            strides=2,\n            dropout=0.0,\n        ):\n            super().__init__()\n            self.attention = AttentionBlock(\n                spatial_dims=spatial_dims, f_g=in_channels, f_l=in_channels, f_int=in_channels // 2\n            )\n            self.upconv = UpConv(\n                spatial_dims=spatial_dims,\n                in_channels=out_channels,\n                out_channels=in_channels,\n                strides=strides,\n                kernel_size=up_kernel_size,\n            )\n            self.merge = Convolution(\n                spatial_dims=spatial_dims, in_channels=2 * in_channels, out_channels=in_channels, dropout=dropout\n            )\n            self.submodule = submodule\n    \n        def forward(self, x: torch.Tensor) -> torch.Tensor:\n            fromlower = self.upconv(self.submodule(x))\n            att = self.attention(g=fromlower, x=x)\n            att_m: torch.Tensor = self.merge(torch.cat((att, fromlower), dim=1))\n            return att_m\n    \n    \n    class AttentionUnet(nn.Module):\n        \"\"\"\n        Attention Unet based on\n        Otkay et al. \"Attention U-Net: Learning Where to Look for the Pancreas\"\n        https://arxiv.org/abs/1804.03999\n    \n        Args:\n            spatial_dims: number of spatial dimensions of the input image.\n            in_channels: number of the input channel.\n            out_channels: number of the output classes.\n            channels (Sequence[int]): sequence of channels. Top block first. The length of `channels` should be no less than 2.\n            strides (Sequence[int]): stride to use for convolutions.\n            kernel_size: convolution kernel size.\n            up_kernel_size: convolution kernel size for transposed convolution layers.\n            dropout: dropout ratio. Defaults to no dropout.\n        \"\"\"\n    \n        def __init__(\n            self,\n            spatial_dims: int,\n            in_channels: int,\n            out_channels: int,\n            channels: Sequence[int],\n            strides: Sequence[int],\n            kernel_size: Sequence[int] | int = 3,\n            up_kernel_size: Sequence[int] | int = 3,\n            dropout: float = 0.0,\n        ):\n            super().__init__()\n            self.dimensions = spatial_dims\n            self.in_channels = in_channels\n            self.out_channels = out_channels\n            self.channels = channels\n            self.strides = strides\n            self.kernel_size = kernel_size\n            self.dropout = dropout\n    \n            head = ConvBlock(\n                spatial_dims=spatial_dims,\n                in_channels=in_channels,\n                out_channels=channels[0],\n                dropout=dropout,\n                kernel_size=self.kernel_size,\n            )\n            reduce_channels = Convolution(\n                spatial_dims=spatial_dims,\n                in_channels=channels[0],\n                out_channels=out_channels,\n                kernel_size=1,\n                strides=1,\n                padding=0,\n                conv_only=True,\n            )\n            self.up_kernel_size = up_kernel_size\n    \n            def _create_block(channels: Sequence[int], strides: Sequence[int]) -> nn.Module:\n                if len(channels) > 2:\n                    subblock = _create_block(channels[1:], strides[1:])\n                    return AttentionLayer(\n                        spatial_dims=spatial_dims,\n                        in_channels=channels[0],\n                        out_channels=channels[1],\n                        submodule=nn.Sequential(\n                            ConvBlock(\n                                spatial_dims=spatial_dims,\n                                in_channels=channels[0],\n                                out_channels=channels[1],\n                                strides=strides[0],\n                                dropout=self.dropout,\n                                kernel_size=self.kernel_size,\n                            ),\n                            subblock,\n                        ),\n                        up_kernel_size=self.up_kernel_size,\n                        strides=strides[0],\n                        dropout=dropout,\n                    )\n                else:\n                    # the next layer is the bottom so stop recursion,\n                    # create the bottom layer as the subblock for this layer\n                    return self._get_bottom_layer(channels[0], channels[1], strides[0])\n    \n            encdec = _create_block(self.channels, self.strides)\n            self.model = nn.Sequential(head, encdec, reduce_channels)\n    \n        def _get_bottom_layer(self, in_channels: int, out_channels: int, strides: int) -> nn.Module:\n            return AttentionLayer(\n                spatial_dims=self.dimensions,\n                in_channels=in_channels,\n                out_channels=out_channels,\n                submodule=ConvBlock(\n                    spatial_dims=self.dimensions,\n                    in_channels=in_channels,\n                    out_channels=out_channels,\n                    strides=strides,\n                    dropout=self.dropout,\n                    kernel_size=self.kernel_size,\n                ),\n                up_kernel_size=self.up_kernel_size,\n                strides=strides,\n                dropout=self.dropout,\n            )\n    \n        def forward(self, x: torch.Tensor) -> torch.Tensor:\n            x_m: torch.Tensor = self.model(x)\n            return x_m","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-05-20T10:58:41.492115Z","iopub.execute_input":"2025-05-20T10:58:41.492328Z","iopub.status.idle":"2025-05-20T10:58:41.513519Z","shell.execute_reply.started":"2025-05-20T10:58:41.492309Z","shell.execute_reply":"2025-05-20T10:58:41.512706Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Squeeze and Excitation","metadata":{}},{"cell_type":"code","source":"if model_architecture == \"3DUnetSE\":\n    class DoubleConvolution(nn.Module):\n        \"\"\"\n        Auxiliary class to define a convolutional layer.\n        Each convolution block: 3x3 convolution, batch normalization, ReLU activation.\n    \n        Args:\n            nn.Module : receive the nn.Module properties\n        \"\"\"\n        def __init__(self, in_channels : int, out_channels : int) -> None:\n            \"\"\"\n            Args:\n                in_channels (int): amount of input channels (16 or 32 or 64 or 128)\n                out_channels (int): amount of output channels (16 or 32 or 64 or 128)\n            \"\"\" \n            super(DoubleConvolution, self).__init__()\n            \n            self.doubleConv = nn.Sequential(\n                nn.Conv3d(in_channels, out_channels, kernel_size = 3, padding=1),\n                nn.BatchNorm3d(out_channels), \n                nn.ReLU(inplace=True),\n                \n                nn.Conv3d(out_channels, out_channels, kernel_size = 3, padding=1),\n                nn.BatchNorm3d(out_channels), \n                nn.ReLU(inplace=True)\n            )\n            \n        def forward(self, x : torch.Tensor) -> torch.Tensor:\n            \"\"\"\n            Args:\n                x (torch.Tensor): input tensor\n    \n            Returns:\n                torch.Tensor: output tensor\n            \"\"\"\n            return self.doubleConv(x)\n    \n    class SE(nn.Module):\n        \"\"\"\n        Auxiliary class to define a squeeze and excitation layer.\n    \n        Args:\n            nn.Module: receive the nn.Module properties\n        \"\"\"\n        def __init__(self, in_channels : int) -> None:\n            \"\"\"\n            Args:\n                in_channels (int): amount of input channels (16 or 32 or 64 or 128)\n            \"\"\"        \n            super(SE, self).__init__()\n            \n            self.squeeze = nn.AdaptiveAvgPool3d(1) # Global Average Pooling\n            self.excitation = nn.Sequential(\n                nn.Linear(in_channels, in_channels // 8), # Reduction ratio = 8\n                nn.ReLU(inplace=True), # ReLU activation\n                nn.Linear(in_channels // 8, in_channels), # Increase ratio = 8\n                nn.Sigmoid() # Sigmoid activation\n            )\n            \n        def forward(self, x : torch.Tensor) -> torch.Tensor:\n            \"\"\"\n            Args:\n                x (torch.Tensor): input tensor\n    \n            Returns:\n                torch.Tensor: output tensor\n            \"\"\"        \n            batch_size, channels, _, _, _ = x.size()\n            y = self.squeeze(x).view(batch_size, channels)\n            y = self.excitation(y).view(batch_size, channels, 1, 1, 1)\n            return x * y.expand_as(x)\n    \n    class DownSampling(nn.Module):\n        \"\"\"\n        Auxiliary class to define a downsampling layer.\n        Each downsampling block: 2x2 max pooling, double convolution and squeeze and excitation.\n        input X output: [1, 16, 128, 128, 128] ->  [1, 32, 64, 64, 64] \n                        [1, 32, 64, 64, 64]    ->  [1, 64, 32, 32, 32]\n                        [1, 64, 32, 32, 32]    ->  [1, 128, 16, 16, 16]\n    \n        Args:\n            nn.Module: receive the nn.Module properties\n        \"\"\"\n        def __init__(self, in_channels : int, out_channels : int) -> None:\n            \"\"\"\n            Args:\n                in_channels (int): amount of input channels (16 or 32 or 64)\n                out_channels (int): amount of output channels (32 or 64 or 128)\n            \"\"\"        \n            super(DownSampling, self).__init__()\n            \n            self.maxpool = nn.MaxPool3d(2)\n            self.conv = DoubleConvolution(in_channels, out_channels)\n            self.attention = SE(out_channels)\n            \n        def forward(self, x : torch.Tensor) -> torch.Tensor:\n            \"\"\"\n            Args:\n                x (torch.Tensor): _description_\n    \n            Returns:\n                torch.Tensor: _description_\n            \"\"\"\n            out = self.maxpool(x) # 2x2 max pooling -> 1/2 the size but same amount of channels\n            out = self.conv(out) # double convolution -> same size but double the amount of channels\n            out = self.attention(out) # squeeze and excitation\n            return out\n    \n    class UpSampling(nn.Module):\n        \"\"\"\n        Auxiliary class to define a upsampling layer.\n        Each upsampling block: 2x2 upsampling, concatenation with skip connection, double convolution.\n        input X output: [1, 128, 16, 16, 16] ->  [1, 64, 32, 32, 32]\n                        [1, 64, 32, 32, 32]    ->  [1, 32, 64, 64, 64]\n                        [1, 32, 64, 64, 64]    ->  [1, 16, 128, 128, 128]\n                        \n        Args:\n            nn.Module: receive the nn.Module properties\n        \"\"\"\n        def __init__(self, in_channels: int, out_channels: int, bilinear: bool = False) -> None:\n            \"\"\"\n            Args:\n                in_channels (int): amount of input channels (128 or 64 or 32)\n                out_channels (int): amount of output channels (64 or 32 or 16)\n            \"\"\"\n            super(UpSampling, self).__init__()\n            \n            self.up = nn.ConvTranspose3d(in_channels, in_channels, kernel_size=2, stride=2)\n            self.conv = DoubleConvolution(int(in_channels + out_channels), out_channels)\n            \n        def forward(self, x : torch.Tensor, skip_connection : torch.Tensor) -> torch.Tensor:\n            \"\"\"\n            Args:\n                x (torch.Tensor): the input tensor\n                skip_connection (torch.Tensor): the skip connection from the downsampling path\n    \n            Returns:\n                torch.Tensor: the output tensor\n            \"\"\"\n            x = self.up(x) # 2x2 upsampling -> double the size but same amount of channels\n            x = torch.cat([skip_connection, x], dim=1) # concatenation with skip connection\n            out = self.conv(x) # double convolution -> same size but half the amount of channels\n            return out\n    \n    class SqueezeAndExcitation3DUnet(nn.Module):\n        def __init__(self, in_channels=1, out_channels=1, channels=(32, 64, 128, 128)) -> None:\n            super(SqueezeAndExcitation3DUnet, self).__init__()\n            layers = channels\n            \n            self.input = nn.Sequential(DoubleConvolution(in_channels, layers[0]), SE(layers[0])) # tranform the input to 16 channels and apply squeeze and excitation\n            # encoding path\n            self.down1 = DownSampling(layers[0], layers[1]) \n            self.down2 = DownSampling(layers[1], layers[2]) \n            self.down3 = DownSampling(layers[2], layers[3])\n            # decoding path\n            self.up1 = UpSampling(layers[3], layers[2])\n            self.up2 = UpSampling(layers[2], layers[1])\n            self.up3 = UpSampling(layers[1], layers[0])\n            self.output = nn.Sequential(nn.Conv3d(layers[0], out_channels, kernel_size=1)) # transform the output\n        \n        def forward(self, x : torch.Tensor) -> torch.Tensor:\n            \"\"\"\n            Args:\n                x (torch.Tensor): a tensor with shape [1, 1, 128, 128, 128]\n    \n            Returns:\n                torch.Tensor: a tensor with shape [1, 1, 128, 128, 128]\n            \"\"\"\n            input = self.input(x) # [1, 1, 128, 128, 128] -> [1, 16, 128, 128, 128]\n            down1_output = self.down1(input)# [1, 16, 128, 128, 128] ->[1, 32, 64, 64, 64]\n            down2_output = self.down2(down1_output) # [1, 32, 64, 64, 64] -> [1, 64, 32, 32, 32]\n            down3_output = self.down3(down2_output) # [1, 64, 32, 32, 32] -> [1, 128, 16, 16, 16]\n            out = self.up1(down3_output, down2_output) # [1, 128, 16, 16, 16] -> [1, 64, 32, 32, 32]\n            out = self.up2(out, down1_output) # [1, 64, 32, 32, 32] -> [1, 32, 64, 64, 64]\n            out = self.up3(out, input) # [1, 32, 64, 64, 64] -> [1, 16, 128, 128, 128]\n            out = self.output(out) # [1, 16, 128, 128, 128] -> [1, 1, 128, 128, 128]\n            return out","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-05-20T10:58:41.515294Z","iopub.execute_input":"2025-05-20T10:58:41.515513Z","iopub.status.idle":"2025-05-20T10:58:41.531145Z","shell.execute_reply.started":"2025-05-20T10:58:41.515495Z","shell.execute_reply":"2025-05-20T10:58:41.530431Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Convolution Block Attention Module","metadata":{}},{"cell_type":"code","source":"if model_architecture == \"3DUnetCBAM\":\n    from monai.networks.blocks.convolutions import Convolution\n    from monai.networks.layers.factories import Norm\n    class SingleConvolution(nn.Module):\n        \"\"\"\n        Auxiliary class to define a convolutional layer.\n        Each convolution block: convolution, batch normalization, ReLU activation.\n    \n        Args:\n            nn.Module : receive the nn.Module properties\n        \"\"\"\n        def __init__(self, in_channels : int, out_channels : int, kernel_size: int = 3, padding: int = 1, stride : int = 1, bias: bool = False) -> None:\n            \"\"\"\n            Args:\n                in_channels (int): amount of input channels\n                out_channels (int): amount of output channels\n            \"\"\" \n            super(SingleConvolution, self).__init__()\n            \n            self.kernel_size = kernel_size\n            self.padding = padding\n            self.stride = stride\n            self.bias = bias\n            \n            self.singleConv = nn.Sequential(\n                nn.Conv3d(in_channels, out_channels, kernel_size = self.kernel_size, padding = self.padding, stride = self.stride, bias = self.bias),\n                nn.BatchNorm3d(out_channels), \n                nn.ReLU(inplace=True),\n            )\n            \n        def forward(self, x : torch.Tensor) -> torch.Tensor:\n            \"\"\"\n            Args:\n                x (torch.Tensor): input tensor\n    \n            Returns:\n                torch.Tensor: output tensor\n            \"\"\"\n            return self.singleConv(x)\n    \n    class DoubleConvolution(nn.Module):\n        \"\"\"\n        Auxiliary class to define a convolutional layer.\n        Each convolution block: 3x3 convolution, batch normalization, ReLU activation.\n    \n        Args:\n            nn.Module : receive the nn.Module properties\n        \"\"\"\n        def __init__(self, in_channels : int, out_channels : int, kernel_size: int = 3, padding: int = 1, stride: int = 1, bias: bool = False) -> None:\n            \"\"\"\n            Args:\n                in_channels (int): amount of input channels\n                out_channels (int): amount of output channels\n            \"\"\" \n            super(DoubleConvolution, self).__init__()\n            \n            self.conv1 = SingleConvolution(in_channels, out_channels, \n                                          kernel_size=kernel_size, \n                                          padding=padding, \n                                          stride=stride, \n                                          bias=bias)\n            self.conv2 = SingleConvolution(out_channels, out_channels, \n                                          kernel_size=kernel_size, \n                                          padding=padding, \n                                          stride=stride, \n                                          bias=bias)\n            \n        def forward(self, x : torch.Tensor) -> torch.Tensor:\n            \"\"\"\n            Args:\n                x (torch.Tensor): input tensor\n    \n            Returns:\n                torch.Tensor: output tensor\n            \"\"\"\n            x = self.conv1(x)\n            x = self.conv2(x)\n            return x\n    \n    class ChannelGate(nn.Module): # C x 1 x 1\n        def __init__(self, in_channels: int, reduction_ratio: int = 16):\n            \n            super(ChannelGate, self).__init__()\n            \n            self.in_channels = in_channels\n            self.reduction_ratio = reduction_ratio\n            \n            self.squeezeAvg = nn.AdaptiveAvgPool3d(1) # Global Average Pooling\n            self.squeezeMax = nn.AdaptiveMaxPool3d(1) # Global Max Pooling\n            self.excitation = nn.Sequential(\n                nn.Linear(self.in_channels, self.in_channels // self.reduction_ratio),\n                nn.ReLU(inplace=True), # ReLU activation\n                nn.Linear(self.in_channels // self.reduction_ratio, self.in_channels),\n            )\n            self.sigActivation = nn.Sigmoid()\n            \n        def forward(self, x):\n            \n            batch_size, channels, _, _, _ = x.size()\n            \n            yAvg = self.squeezeAvg(x).view(batch_size, channels)\n            yAvg = self.excitation(yAvg).view(batch_size, channels, 1, 1, 1)\n            \n            yMax = self.squeezeMax(x).view(batch_size, channels)\n            yMax = self.excitation(yMax).view(batch_size, channels, 1, 1, 1)\n            \n            sum = yAvg + yMax\n            \n            return self.sigActivation(sum) * x\n    \n    class SpatialGate(nn.Module): # 1 x H x W\n        def __init__(self, in_channels: int, kernel_size: int = 7, padding: int = 3, bias: bool = False):\n            \n            super(SpatialGate, self).__init__()\n            \n            self.kernel_size = kernel_size\n            self.in_channels = in_channels\n            self.padding = padding\n            self.bias = bias\n            \n            self.squeezeAvg = nn.AdaptiveAvgPool3d(1)\n            self.squeezeMax = nn.AdaptiveMaxPool3d(1)\n            self.spatial = SingleConvolution(2 * self.in_channels, self.in_channels, kernel_size = self.kernel_size, padding = self.padding)\n            self.sigActivation = nn.Sigmoid()\n            \n        def forward(self, x):\n            \n            yAvg = self.squeezeAvg(x)\n            yMax = self.squeezeMax(x)\n            y = torch.cat([yAvg, yMax], dim=1)\n            y = self.spatial(y)\n            \n            return self.sigActivation(y) * x\n    \n    class CBAM(nn.Module):\n        def __init__(self, in_channels, reduction_ratio = 16, kernel_size: int = 7):\n            \n            super(CBAM, self).__init__()\n            \n            self.ChannelGate = ChannelGate(in_channels, reduction_ratio)\n            self.SpatialGate = SpatialGate(in_channels, kernel_size = kernel_size, padding = kernel_size // 2)\n            \n        def forward(self, x):\n            \n            x_out = self.ChannelGate(x)\n            x_out = self.SpatialGate(x_out)\n            \n            return x_out * x\n    \n    class DownSampling(nn.Module):\n        \"\"\"\n        Auxiliary class to define a downsampling layer.\n        Each downsampling block: 2x2 max pooling, double convolution and squeeze and excitation.\n        input X output: [1, 16, 128, 128, 128] ->  [1, 32, 64, 64, 64] \n                        [1, 32, 64, 64, 64]    ->  [1, 64, 32, 32, 32]\n                        [1, 64, 32, 32, 32]    ->  [1, 128, 16, 16, 16]\n    \n        Args:\n            nn.Module: receive the nn.Module properties\n        \"\"\"\n        def __init__(self, in_channels : int, out_channels : int, attention) -> None:\n            \"\"\"\n            Args:\n                in_channels (int): amount of input channels (16 or 32 or 64)\n                out_channels (int): amount of output channels (32 or 64 or 128)\n            \"\"\"        \n            super(DownSampling, self).__init__()\n            \n            self.attentionFunction = attention\n            \n            self.maxpool = nn.MaxPool3d(2)\n            self.conv = DoubleConvolution(in_channels, out_channels, kernel_size=3, padding=1, stride=1, bias=False)\n            self.attention = self.attentionFunction(out_channels)\n            \n        def forward(self, x : torch.Tensor) -> torch.Tensor:\n            \"\"\"\n            Args:\n                x (torch.Tensor): _description_\n    \n            Returns:\n                torch.Tensor: _description_\n            \"\"\"\n            out = self.maxpool(x) # 2x2 max pooling -> 1/2 the size but same amount of channels\n            out = self.conv(out) # double convolution -> same size but double the amount of channels\n            out = self.attention(out) # squeeze and excitation\n            return out\n    \n    class UpSampling(nn.Module):\n        \"\"\"\n        Auxiliary class to define a upsampling layer.\n        Each upsampling block: 2x2 upsampling, concatenation with skip connection, double convolution.\n        input X output: [1, 128, 16, 16, 16] ->  [1, 64, 32, 32, 32]\n                        [1, 64, 32, 32, 32]    ->  [1, 32, 64, 64, 64]\n                        [1, 32, 64, 64, 64]    ->  [1, 16, 128, 128, 128]\n                        \n        Args:\n            nn.Module: receive the nn.Module properties\n        \"\"\"\n        def __init__(self, in_channels: int, out_channels: int, bilinear: bool = False) -> None:\n            \"\"\"\n            Args:\n                in_channels (int): amount of input channels (128 or 64 or 32)\n                out_channels (int): amount of output channels (64 or 32 or 16)\n            \"\"\"\n            super(UpSampling, self).__init__()\n            \n            self.up = nn.ConvTranspose3d(in_channels, in_channels, kernel_size=2, stride=2)\n            self.conv = DoubleConvolution(int(in_channels + out_channels), out_channels, kernel_size=3, padding=1, stride=1, bias=False)\n            \n        def forward(self, x : torch.Tensor, skip_connection : torch.Tensor) -> torch.Tensor:\n            \"\"\"\n            Args:\n                x (torch.Tensor): the input tensor\n                skip_connection (torch.Tensor): the skip connection from the downsampling path\n    \n            Returns:\n                torch.Tensor: the output tensor\n            \"\"\"\n            x = self.up(x) # 2x2 upsampling -> double the size but same amount of channels\n            x = torch.cat([skip_connection, x], dim=1) # concatenation with skip connection\n            out = self.conv(x) # double convolution -> same size but half the amount of channels\n            return out\n    \n    class Attention_Unet(nn.Module):\n        def __init__(self, attentionFunction, in_channels=1, out_channels=1, channels=(32, 64, 128, 128)) -> None:\n            super(Attention_Unet, self).__init__()\n            self.attentionFunction = attentionFunction\n            layers = channels \n    \n            self.attentionFunction = attentionFunction\n            \n            self.input = nn.Sequential(DoubleConvolution(in_channels, layers[0], kernel_size=3, padding = 1, stride=1, bias=False), self.attentionFunction(layers[0])) # tranform the input to 16 channels and apply squeeze and excitation\n            # encoding path\n            self.down1 = DownSampling(layers[0], layers[1], self.attentionFunction) \n            self.down2 = DownSampling(layers[1], layers[2], self.attentionFunction) \n            self.down3 = DownSampling(layers[2], layers[3], self.attentionFunction)\n            # decoding path\n            self.up1 = UpSampling(layers[3], layers[2])\n            self.up2 = UpSampling(layers[2], layers[1])\n            self.up3 = UpSampling(layers[1], layers[0])\n            self.output = nn.Sequential(nn.Conv3d(layers[0], out_channels, kernel_size=1)) # transform the output to 7 channel \n        \n        def forward(self, x : torch.Tensor) -> torch.Tensor:\n            \"\"\"\n            Args:\n                x (torch.Tensor): a tensor with shape [1, 1, 128, 128, 128]\n    \n            Returns:\n                torch.Tensor: a tensor with shape [1, 1, 128, 128, 128]\n            \"\"\"\n            input = self.input(x) # [1, 1, 128, 128, 128] -> [1, 16, 128, 128, 128]\n            down1_output = self.down1(input)# [1, 16, 128, 128, 128] ->[1, 32, 64, 64, 64]\n            down2_output = self.down2(down1_output) # [1, 32, 64, 64, 64] -> [1, 64, 32, 32, 32]\n            down3_output = self.down3(down2_output) # [1, 64, 32, 32, 32] -> [1, 128, 16, 16, 16]\n            out = self.up1(down3_output, down2_output) # [1, 128, 16, 16, 16] -> [1, 64, 32, 32, 32]\n            out = self.up2(out, down1_output) # [1, 64, 32, 32, 32] -> [1, 32, 64, 64, 64]\n            out = self.up3(out, input) # [1, 32, 64, 64, 64] -> [1, 16, 128, 128, 128]\n            out = self.output(out) # [1, 16, 128, 128, 128] -> [1, 1, 128, 128, 128]\n            return out","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-05-20T10:58:41.532036Z","iopub.execute_input":"2025-05-20T10:58:41.532357Z","iopub.status.idle":"2025-05-20T10:58:41.553035Z","shell.execute_reply.started":"2025-05-20T10:58:41.532326Z","shell.execute_reply":"2025-05-20T10:58:41.552198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)\n\nif model_architecture == \"3DUnet\":\n    model = UNet(\n        spatial_dims=3,\n        in_channels=1,\n        out_channels=7,\n        channels=model_channels,\n        strides=(2, 2, 1),\n        num_res_units=1,\n    ).to(device)\n\nelif model_architecture == \"3DUnetAG\":\n    model = AttentionUnet(\n        spatial_dims=3,\n        in_channels=1,\n        out_channels=7,\n        channels=model_channels,\n        strides=(2, 2, 1),\n    ).to(device)\n    \nelif model_architecture == \"3DUnetSE\":\n    model = SqueezeAndExcitation3DUnet(\n        in_channels=1,\n        out_channels=7,\n        channels=model_channels,\n    ).to(device)\n    \nelif model_architecture == \"3DUnetCBAM\":\n    model = Attention_Unet(\n        attentionFunction=CBAM,\n        in_channels=1,\n        out_channels=7,\n        channels=model_channels,\n    ).to(device)\n    \n# if multiple GPUs are available, wrap the model with DataParallel\n# if torch.cuda.device_count() > 1:\n#     print(f\"Using {torch.cuda.device_count()} GPUs!\")\n#     model = nn.DataParallel(model) \n\nmodel = model.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:41.55368Z","iopub.execute_input":"2025-05-20T10:58:41.553898Z","iopub.status.idle":"2025-05-20T10:58:41.600408Z","shell.execute_reply.started":"2025-05-20T10:58:41.553879Z","shell.execute_reply":"2025-05-20T10:58:41.599603Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Predict","metadata":{}},{"cell_type":"code","source":"# Load the best model\nmodel.load_state_dict(torch.load(model_path, weights_only=True))  \nmodel.to(\"cuda\")\nmodel.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:41.601255Z","iopub.execute_input":"2025-05-20T10:58:41.601514Z","iopub.status.idle":"2025-05-20T10:58:41.647214Z","shell.execute_reply.started":"2025-05-20T10:58:41.601493Z","shell.execute_reply":"2025-05-20T10:58:41.646565Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\n\ncopick_config_path = \"/kaggle/working\" + \"/copick.config\"\n\nwith open(copick_config_path) as f:\n    copick_config = json.load(f)\n    print(copick_config)\n\ncopick_config['static_root'] = '/kaggle/input/server-datatest/test/static'\n# copick_config['static_root'] = '/kaggle/input/czii-cryo-et-object-identification/test/static'\n\ncopick_test_config_path = 'copick_test.config'\n\nwith open(copick_test_config_path, 'w') as outfile:\n    json.dump(copick_config, outfile)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:41.647872Z","iopub.execute_input":"2025-05-20T10:58:41.64806Z","iopub.status.idle":"2025-05-20T10:58:41.653669Z","shell.execute_reply.started":"2025-05-20T10:58:41.648043Z","shell.execute_reply":"2025-05-20T10:58:41.65297Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import copick\n\nroot = copick.from_file(copick_test_config_path)\n\ncopick_user_name = \"copickUtils\"\nvoxel_size = 10\ntomo_type = \"denoised\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:41.654432Z","iopub.execute_input":"2025-05-20T10:58:41.654626Z","iopub.status.idle":"2025-05-20T10:58:41.666436Z","shell.execute_reply.started":"2025-05-20T10:58:41.654604Z","shell.execute_reply":"2025-05-20T10:58:41.665763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"inference_transforms = Compose([\n    EnsureChannelFirstd(keys=[\"image\"], channel_dim=\"no_channel\"),\n    NormalizeIntensityd(keys=\"image\"),\n    Orientationd(keys=[\"image\"], axcodes=\"RAS\")\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:41.667237Z","iopub.execute_input":"2025-05-20T10:58:41.66746Z","iopub.status.idle":"2025-05-20T10:58:41.67818Z","shell.execute_reply.started":"2025-05-20T10:58:41.667437Z","shell.execute_reply":"2025-05-20T10:58:41.677533Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"voxel_size = 10\ntomo_type = \"denoised\"\ndevice = \"cuda\"\n\nclasses = [1, 2, 3, 4, 5, 6]\nid_to_name = {\n    1: \"apo-ferritin\", \n    2: \"beta-amylase\",\n    3: \"beta-galactosidase\", \n    4: \"ribosome\", \n    5: \"thyroglobulin\", \n    6: \"virus-like-particle\"\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:41.678975Z","iopub.execute_input":"2025-05-20T10:58:41.679228Z","iopub.status.idle":"2025-05-20T10:58:41.693421Z","shell.execute_reply.started":"2025-05-20T10:58:41.679197Z","shell.execute_reply":"2025-05-20T10:58:41.692648Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_thresholds = {\n    1: {\"blob\": 90,  \"certainty\": 0.5},\n    2: {\"blob\": 100, \"certainty\": 0.7},\n    3: {\"blob\": 200, \"certainty\": 0.95},\n    4: {\"blob\": 100,\"certainty\": 0.4},\n    5: {\"blob\": 500, \"certainty\": 0.95},\n    6: {\"blob\": 500,\"certainty\": 0.5},\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:41.694164Z","iopub.execute_input":"2025-05-20T10:58:41.694431Z","iopub.status.idle":"2025-05-20T10:58:41.705314Z","shell.execute_reply.started":"2025-05-20T10:58:41.694403Z","shell.execute_reply":"2025-05-20T10:58:41.704405Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\ndef dict_to_df(location_dict, run_name):\n    rows = []\n    for particle_type, coords in location_dict.items():\n        for x, y, z in coords:\n            rows.append({\n                \"experiment\": run_name,\n                \"particle_type\": particle_type,\n                \"x\": x,\n                \"y\": y,\n                \"z\": z\n            })\n    return pd.DataFrame(rows)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:41.706297Z","iopub.execute_input":"2025-05-20T10:58:41.706588Z","iopub.status.idle":"2025-05-20T10:58:41.716403Z","shell.execute_reply.started":"2025-05-20T10:58:41.706547Z","shell.execute_reply":"2025-05-20T10:58:41.715661Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from monai.inferers import sliding_window_inference\nfrom skimage.measure import label, regionprops\n\nlocation_df = []\n\nwith torch.no_grad():\n    for run in tqdm(root.runs):\n        print(f\"Processing run: {run.name}\")\n\n        # Get voxel spacing spec\n        voxel_spacing_spec = run.get_voxel_spacing(voxel_size)\n        if voxel_spacing_spec is None:\n            print(f\"Skipping {run.name}: No voxel spacing for voxel_size={voxel_size}\")\n            continue\n\n        voxel_spacing = voxel_spacing_spec.voxel_size\n        tomo = voxel_spacing_spec.get_tomogram(tomo_type).numpy()\n\n        sample = {\"image\": tomo}\n        sample = inference_transforms(sample)\n\n        input_tensor = sample[\"image\"].unsqueeze(0).to(\"cuda\")  # shape: (1, 1, Z, Y, X)\n\n        # Sliding window inference\n        output = sliding_window_inference(\n            inputs=input_tensor,\n            roi_size=(96, 96, 96),\n            sw_batch_size=8,\n            predictor=model,\n            overlap=0.25\n        )[0]  # shape: (num_classes, Z, Y, X)\n\n        probs = torch.softmax(output, dim=0)  # (num_classes, Z, Y, X)\n        pred = torch.argmax(probs, dim=0).cpu().numpy()  # (Z, Y, X)\n\n        # Extract centroids\n        location = {}\n        for c in classes:\n            thresholds = class_thresholds.get(c, {\"blob\": 500, \"certainty\": 0.5})\n            blob_thresh = thresholds[\"blob\"]\n            certainty_thresh = thresholds[\"certainty\"]\n        \n            binary_mask = (probs[c] > certainty_thresh).cpu().numpy().astype(np.uint8)\n            labeled = label(binary_mask)\n            regions = regionprops(labeled)\n        \n            centroids = []\n            for region in regions:\n                if region.area > blob_thresh:\n                    z, y, x = region.centroid\n                    centroids.append([\n                        x * 10.012444,\n                        y * 10.012444,\n                        z * 10.012444\n                    ])\n        \n            location[id_to_name[c]] = np.array(centroids)\n\n\n        df = dict_to_df(location, run.name)\n        location_df.append(df)\n\n# Save output\nlocation_df = pd.concat(location_df, ignore_index=True)\nlocation_df.index.name = \"id\"\nlocation_df.to_csv(\"submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T10:58:41.717121Z","iopub.execute_input":"2025-05-20T10:58:41.717311Z","iopub.status.idle":"2025-05-20T11:01:20.51213Z","shell.execute_reply.started":"2025-05-20T10:58:41.717295Z","shell.execute_reply":"2025-05-20T11:01:20.511372Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"location_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T11:01:20.512968Z","iopub.execute_input":"2025-05-20T11:01:20.513189Z","iopub.status.idle":"2025-05-20T11:01:20.541193Z","shell.execute_reply.started":"2025-05-20T11:01:20.513169Z","shell.execute_reply":"2025-05-20T11:01:20.54023Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp -r /kaggle/input/hengck-czii-cryo-et-01/* .","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T11:01:20.542316Z","iopub.execute_input":"2025-05-20T11:01:20.542721Z","iopub.status.idle":"2025-05-20T11:01:21.021882Z","shell.execute_reply.started":"2025-05-20T11:01:20.542682Z","shell.execute_reply":"2025-05-20T11:01:21.02093Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from czii_helper import *\nfrom dataset import *\nfrom scipy.optimize import linear_sum_assignment\nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T11:01:21.022899Z","iopub.execute_input":"2025-05-20T11:01:21.023129Z","iopub.status.idle":"2025-05-20T11:01:21.030474Z","shell.execute_reply.started":"2025-05-20T11:01:21.023109Z","shell.execute_reply":"2025-05-20T11:01:21.029802Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nif os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n    MODE = 'submit'\nelse:\n    MODE = 'local'\n\nvalid_dir ='/kaggle/input/server-datatest/test'\nvalid_id = ['TS_101_5']\n\ndef do_one_eval(truth, predict, threshold):\n    P=len(predict)\n    T=len(truth)\n\n    if P==0:\n        hit=[[],[]]\n        miss=np.arange(T).tolist()\n        fp=[]\n        metric = [P,T,len(hit[0]),len(miss),len(fp)]\n        return hit, fp, miss, metric\n\n    if T==0:\n        hit=[[],[]]\n        fp=np.arange(P).tolist()\n        miss=[]\n        metric = [P,T,len(hit[0]),len(miss),len(fp)]\n        return hit, fp, miss, metric\n\n    #---\n    distance = predict.reshape(P,1,3)-truth.reshape(1,T,3)\n    distance = distance**2\n    distance = distance.sum(axis=2)\n    distance = np.sqrt(distance)\n    p_index, t_index = linear_sum_assignment(distance)\n\n    valid = distance[p_index, t_index] <= threshold\n    p_index = p_index[valid]\n    t_index = t_index[valid]\n    hit = [p_index.tolist(), t_index.tolist()]\n    miss = np.arange(T)\n    miss = miss[~np.isin(miss,t_index)].tolist()\n    fp = np.arange(P)\n    fp = fp[~np.isin(fp,p_index)].tolist()\n\n    metric = [P,T,len(hit[0]),len(miss),len(fp)] #for lb metric F-beta copmutation\n    return hit, fp, miss, metric\n\n\ndef compute_lb(submit_df, overlay_dir):\n    valid_id = list(submit_df['experiment'].unique())\n    print(valid_id)\n\n    eval_df = []\n    for id in valid_id:\n        truth = read_one_truth(id, overlay_dir) #=f'{valid_dir}/overlay/ExperimentRuns')\n        id_df = submit_df[submit_df['experiment'] == id]\n        for p in PARTICLE:\n            p = dotdict(p)\n            print('\\r', id, p.name, end='', flush=True)\n            xyz_truth = truth[p.name]\n            xyz_predict = id_df[id_df['particle_type'] == p.name][['x', 'y', 'z']].values\n            hit, fp, miss, metric = do_one_eval(xyz_truth, xyz_predict, p.radius* 0.5)\n            eval_df.append(dotdict(\n                id=id, particle_type=p.name,\n                P=metric[0], T=metric[1], hit=metric[2], miss=metric[3], fp=metric[4],\n            ))\n    print('')\n    eval_df = pd.DataFrame(eval_df)\n    gb = eval_df.groupby('particle_type').agg('sum').drop(columns=['id'])\n    gb.loc[:, 'precision'] = gb['hit'] / gb['P']\n    gb.loc[:, 'precision'] = gb['precision'].fillna(0)\n    gb.loc[:, 'recall'] = gb['hit'] / gb['T']\n    gb.loc[:, 'recall'] = gb['recall'].fillna(0)\n    gb.loc[:, 'f-beta4'] = 17 * gb['precision'] * gb['recall'] / (16 * gb['precision'] + gb['recall'])\n    gb.loc[:, 'f-beta4'] = gb['f-beta4'].fillna(0)\n\n    gb = gb.sort_values('particle_type').reset_index(drop=False)\n    # https://www.kaggle.com/competitions/czii-cryo-et-object-identification/discussion/544895\n    gb.loc[:, 'weight'] = [1, 0, 2, 1, 2, 1]\n    lb_score = (gb['f-beta4'] * gb['weight']).sum() / gb['weight'].sum()\n    return gb, lb_score","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T11:03:02.829497Z","iopub.execute_input":"2025-05-20T11:03:02.829884Z","iopub.status.idle":"2025-05-20T11:03:02.843041Z","shell.execute_reply.started":"2025-05-20T11:03:02.829851Z","shell.execute_reply":"2025-05-20T11:03:02.842154Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nfrom mpl_toolkits.mplot3d import Axes3D\nimport seaborn as sns\n\nif MODE == 'local':\n    submit_df = pd.read_csv('submission.csv')\n    gb, lb_score = compute_lb(submit_df, f'{valid_dir}/overlay/ExperimentRuns')\n    \n    print(\"Details:\")\n    print(gb)\n    print()\n    print(\"Leaderboard Score:\")\n    print(lb_score)\n\n    # ---- Show one volume ----\n    for id in valid_id:\n        # id = valid_id[0]\n        truth = read_one_truth(id, overlay_dir=f'{valid_dir}/overlay/ExperimentRuns')\n        submit_df_id = submit_df[submit_df['experiment'] == id]\n    \n        fig = plt.figure(figsize=(20, 10))\n        particle_metrics = []\n    \n        for p in PARTICLE:\n            p = dotdict(p)\n            xyz_truth = truth[p.name]\n            xyz_predict = submit_df_id[submit_df_id['particle_type'] == p.name][['x', 'y', 'z']].values\n    \n            hit, fp, miss, _ = do_one_eval(xyz_truth, xyz_predict, p.radius)\n            particle_metrics.append({\n                'Particle': p.name,\n                'Truth Count': len(xyz_truth),\n                'Predicted Count': len(xyz_predict),\n                'Hit': len(hit[0]),\n                'False Positive': len(fp),\n                'Missed': len(miss),\n            })\n    \n            ax = fig.add_subplot(2, 3, p.label, projection='3d')\n            ax.set_title(f'{p.name} ({p.difficulty})', fontsize=12)\n    \n            # Hits: red\n            if hit[0]:\n                pt_pred = xyz_predict[hit[0]]\n                pt_gt = xyz_truth[hit[1]]\n                ax.scatter(pt_pred[:, 0], pt_pred[:, 1], pt_pred[:, 2], alpha=0.6, color='g', label='Hit (Pred)')\n                ax.scatter(pt_gt[:, 0], pt_gt[:, 1], pt_gt[:, 2], s=80, facecolors='none', edgecolors='g', label='Hit (GT)')\n    \n            # False Positives: black\n            if len(fp) > 0:\n                pt_fp = xyz_predict[fp]\n                ax.scatter(pt_fp[:, 0], pt_fp[:, 1], pt_fp[:, 2], alpha=0.8, color='k', label='False Positive')\n    \n            # Misses: blue hollow\n            if len(miss) > 0:\n                pt_miss = xyz_truth[miss]\n                ax.scatter(pt_miss[:, 0], pt_miss[:, 1], pt_miss[:, 2], s=160, facecolors='none', edgecolors='r', label='Missed')\n    \n            ax.legend(loc='upper right', fontsize=8)\n    \n        plt.suptitle(f\"Evaluation for Volume: {id}\", fontsize=16)\n        plt.tight_layout()\n        plt.show()\n    \n        # ---- Display Summary Table ----\n        metric_df = pd.DataFrame(particle_metrics)\n        print(\"\\nPer-Particle Evaluation Summary:\")\n        print(metric_df.to_string(index=False))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T11:03:04.79428Z","iopub.execute_input":"2025-05-20T11:03:04.794586Z","iopub.status.idle":"2025-05-20T11:03:16.775131Z","shell.execute_reply.started":"2025-05-20T11:03:04.794561Z","shell.execute_reply":"2025-05-20T11:03:16.774194Z"}},"outputs":[],"execution_count":null}]}