{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"},{"sourceId":11438124,"sourceType":"datasetVersion","datasetId":7164801},{"sourceId":11516072,"sourceType":"datasetVersion","datasetId":7221951},{"sourceId":11545312,"sourceType":"datasetVersion","datasetId":7240234},{"sourceId":11610765,"sourceType":"datasetVersion","datasetId":7282747},{"sourceId":11822200,"sourceType":"datasetVersion","datasetId":7426148},{"sourceId":11946035,"sourceType":"datasetVersion","datasetId":7447807}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# DEIM Single Model Inference Notebook","metadata":{}},{"cell_type":"markdown","source":"- Training Data  \nonly official image data with num_motors>0 used (no external data, no negative sampling).  \n75% training, 25% validation  \n- Image Size  \n(384, 384, 3) (both training and inference)  \n- Model weight and DEIM code (including training config) are not public.  \n","metadata":{}},{"cell_type":"markdown","source":"```\n Average Precision  (AP) @[ IoU=0.50:0.95 | area=   all | maxDets=100 ] = 0.798\n Average Precision  (AP) @[ IoU=0.50      | area=   all | maxDets=100 ] = 0.958\n Average Precision  (AP) @[ IoU=0.75      | area=   all | maxDets=100 ] = 0.927\n Average Precision  (AP) @[ IoU=0.50:0.95 | area= small | maxDets=100 ] = -1.000\n Average Precision  (AP) @[ IoU=0.50:0.95 | area=medium | maxDets=100 ] = 0.798\n Average Precision  (AP) @[ IoU=0.50:0.95 | area= large | maxDets=100 ] = -1.000\n Average Recall     (AR) @[ IoU=0.50:0.95 | area=   all | maxDets=  1 ] = 0.785\n Average Recall     (AR) @[ IoU=0.50:0.95 | area=   all | maxDets= 10 ] = 0.854\n Average Recall     (AR) @[ IoU=0.50:0.95 | area=   all | maxDets=100 ] = 0.883\n Average Recall     (AR) @[ IoU=0.50:0.95 | area= small | maxDets=100 ] = -1.000\n Average Recall     (AR) @[ IoU=0.50:0.95 | area=medium | maxDets=100 ] = 0.883\n Average Recall     (AR) @[ IoU=0.50:0.95 | area= large | maxDets=100 ] = -1.000\n Average Recall     (AR) @[ IoU=0.50      | area=   all | maxDets=100 ] = 1.000\n Average Recall     (AR) @[ IoU=0.75      | area=   all | maxDets=100 ] = 0.971\n```","metadata":{}},{"cell_type":"code","source":"#!pip install -q /kaggle/input/byu-private-dataset/faster_coco_eval-1.6.5-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n#!pip install -q /kaggle/input/byu-private-dataset/calflops-0.3.2-py3-none-any.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:36:29.136515Z","iopub.execute_input":"2025-05-25T13:36:29.136741Z","iopub.status.idle":"2025-05-25T13:36:29.140504Z","shell.execute_reply.started":"2025-05-25T13:36:29.136723Z","shell.execute_reply":"2025-05-25T13:36:29.139833Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp -r /kaggle/input/deim422-offline-packages/packages /kaggle/working/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:36:29.141449Z","iopub.execute_input":"2025-05-25T13:36:29.141617Z","iopub.status.idle":"2025-05-25T13:36:48.483911Z","shell.execute_reply.started":"2025-05-25T13:36:29.141604Z","shell.execute_reply":"2025-05-25T13:36:48.483148Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp -r /kaggle/input/my-deim-train-wts-demo/DEIM /kaggle/working/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:36:48.485340Z","iopub.execute_input":"2025-05-25T13:36:48.485578Z","iopub.status.idle":"2025-05-25T13:36:49.132226Z","shell.execute_reply.started":"2025-05-25T13:36:48.485557Z","shell.execute_reply":"2025-05-25T13:36:49.131103Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/working/packages/')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:36:49.133598Z","iopub.execute_input":"2025-05-25T13:36:49.134334Z","iopub.status.idle":"2025-05-25T13:36:50.211373Z","shell.execute_reply.started":"2025-05-25T13:36:49.134307Z","shell.execute_reply":"2025-05-25T13:36:50.210621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sys.path.append('/kaggle/working/DEIM/')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:36:50.213068Z","iopub.execute_input":"2025-05-25T13:36:50.213289Z","iopub.status.idle":"2025-05-25T13:36:50.265660Z","shell.execute_reply.started":"2025-05-25T13:36:50.213271Z","shell.execute_reply":"2025-05-25T13:36:50.265108Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp /kaggle/input/pretrained-pphgnetv2-wts/* /kaggle/working/ ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:36:50.266263Z","iopub.execute_input":"2025-05-25T13:36:50.266422Z","iopub.status.idle":"2025-05-25T13:36:51.900852Z","shell.execute_reply.started":"2025-05-25T13:36:50.266408Z","shell.execute_reply":"2025-05-25T13:36:51.900115Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n%%writefile  /kaggle/working/DEIM/engine/backbone/hgnetv2.py\n\"\"\"\nreference\n- https://github.com/PaddlePaddle/PaddleDetection/blob/develop/ppdet/modeling/backbones/hgnet_v2.py\n\nCopyright (c) 2024 The D-FINE Authors. All Rights Reserved.\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport os\nfrom .common import FrozenBatchNorm2d\nfrom ..core import register\nimport logging\n\n# Constants for initialization\nkaiming_normal_ = nn.init.kaiming_normal_\nzeros_ = nn.init.zeros_\nones_ = nn.init.ones_\n\n__all__ = ['HGNetv2']\n\n\nclass LearnableAffineBlock(nn.Module):\n    def __init__(\n            self,\n            scale_value=1.0,\n            bias_value=0.0\n    ):\n        super().__init__()\n        self.scale = nn.Parameter(torch.tensor([scale_value]), requires_grad=True)\n        self.bias = nn.Parameter(torch.tensor([bias_value]), requires_grad=True)\n\n    def forward(self, x):\n        return self.scale * x + self.bias\n\n\nclass ConvBNAct(nn.Module):\n    def __init__(\n            self,\n            in_chs,\n            out_chs,\n            kernel_size,\n            stride=1,\n            groups=1,\n            padding='',\n            use_act=True,\n            use_lab=False\n    ):\n        super().__init__()\n        self.use_act = use_act\n        self.use_lab = use_lab\n        if padding == 'same':\n            self.conv = nn.Sequential(\n                nn.ZeroPad2d([0, 1, 0, 1]),\n                nn.Conv2d(\n                    in_chs,\n                    out_chs,\n                    kernel_size,\n                    stride,\n                    groups=groups,\n                    bias=False\n                )\n            )\n        else:\n            self.conv = nn.Conv2d(\n                in_chs,\n                out_chs,\n                kernel_size,\n                stride,\n                padding=(kernel_size - 1) // 2,\n                groups=groups,\n                bias=False\n            )\n        self.bn = nn.BatchNorm2d(out_chs)\n        if self.use_act:\n            self.act = nn.ReLU()\n        else:\n            self.act = nn.Identity()\n        if self.use_act and self.use_lab:\n            self.lab = LearnableAffineBlock()\n        else:\n            self.lab = nn.Identity()\n\n    def forward(self, x):\n        x = self.conv(x)\n        x = self.bn(x)\n        x = self.act(x)\n        x = self.lab(x)\n        return x\n\n\nclass LightConvBNAct(nn.Module):\n    def __init__(\n            self,\n            in_chs,\n            out_chs,\n            kernel_size,\n            groups=1,\n            use_lab=False,\n    ):\n        super().__init__()\n        self.conv1 = ConvBNAct(\n            in_chs,\n            out_chs,\n            kernel_size=1,\n            use_act=False,\n            use_lab=use_lab,\n        )\n        self.conv2 = ConvBNAct(\n            out_chs,\n            out_chs,\n            kernel_size=kernel_size,\n            groups=out_chs,\n            use_act=True,\n            use_lab=use_lab,\n        )\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.conv2(x)\n        return x\n\n\nclass StemBlock(nn.Module):\n    # for HGNetv2\n    def __init__(self, in_chs, mid_chs, out_chs, use_lab=False):\n        super().__init__()\n        self.stem1 = ConvBNAct(\n            in_chs,\n            mid_chs,\n            kernel_size=3,\n            stride=2,\n            use_lab=use_lab,\n        )\n        self.stem2a = ConvBNAct(\n            mid_chs,\n            mid_chs // 2,\n            kernel_size=2,\n            stride=1,\n            use_lab=use_lab,\n        )\n        self.stem2b = ConvBNAct(\n            mid_chs // 2,\n            mid_chs,\n            kernel_size=2,\n            stride=1,\n            use_lab=use_lab,\n        )\n        self.stem3 = ConvBNAct(\n            mid_chs * 2,\n            mid_chs,\n            kernel_size=3,\n            stride=2,\n            use_lab=use_lab,\n        )\n        self.stem4 = ConvBNAct(\n            mid_chs,\n            out_chs,\n            kernel_size=1,\n            stride=1,\n            use_lab=use_lab,\n        )\n        self.pool = nn.MaxPool2d(kernel_size=2, stride=1, ceil_mode=True)\n\n    def forward(self, x):\n        x = self.stem1(x)\n        x = F.pad(x, (0, 1, 0, 1))\n        x2 = self.stem2a(x)\n        x2 = F.pad(x2, (0, 1, 0, 1))\n        x2 = self.stem2b(x2)\n        x1 = self.pool(x)\n        x = torch.cat([x1, x2], dim=1)\n        x = self.stem3(x)\n        x = self.stem4(x)\n        return x\n\n\nclass EseModule(nn.Module):\n    def __init__(self, chs):\n        super().__init__()\n        self.conv = nn.Conv2d(\n            chs,\n            chs,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n        )\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        identity = x\n        x = x.mean((2, 3), keepdim=True)\n        x = self.conv(x)\n        x = self.sigmoid(x)\n        return torch.mul(identity, x)\n\n\nclass HG_Block(nn.Module):\n    def __init__(\n            self,\n            in_chs,\n            mid_chs,\n            out_chs,\n            layer_num,\n            kernel_size=3,\n            residual=False,\n            light_block=False,\n            use_lab=False,\n            agg='ese',\n            drop_path=0.,\n    ):\n        super().__init__()\n        self.residual = residual\n\n        self.layers = nn.ModuleList()\n        for i in range(layer_num):\n            if light_block:\n                self.layers.append(\n                    LightConvBNAct(\n                        in_chs if i == 0 else mid_chs,\n                        mid_chs,\n                        kernel_size=kernel_size,\n                        use_lab=use_lab,\n                    )\n                )\n            else:\n                self.layers.append(\n                    ConvBNAct(\n                        in_chs if i == 0 else mid_chs,\n                        mid_chs,\n                        kernel_size=kernel_size,\n                        stride=1,\n                        use_lab=use_lab,\n                    )\n                )\n\n        # feature aggregation\n        total_chs = in_chs + layer_num * mid_chs\n        if agg == 'se':\n            aggregation_squeeze_conv = ConvBNAct(\n                total_chs,\n                out_chs // 2,\n                kernel_size=1,\n                stride=1,\n                use_lab=use_lab,\n            )\n            aggregation_excitation_conv = ConvBNAct(\n                out_chs // 2,\n                out_chs,\n                kernel_size=1,\n                stride=1,\n                use_lab=use_lab,\n            )\n            self.aggregation = nn.Sequential(\n                aggregation_squeeze_conv,\n                aggregation_excitation_conv,\n            )\n        else:\n            aggregation_conv = ConvBNAct(\n                total_chs,\n                out_chs,\n                kernel_size=1,\n                stride=1,\n                use_lab=use_lab,\n            )\n            att = EseModule(out_chs)\n            self.aggregation = nn.Sequential(\n                aggregation_conv,\n                att,\n            )\n\n        self.drop_path = nn.Dropout(drop_path) if drop_path else nn.Identity()\n\n    def forward(self, x):\n        identity = x\n        output = [x]\n        for layer in self.layers:\n            x = layer(x)\n            output.append(x)\n        x = torch.cat(output, dim=1)\n        x = self.aggregation(x)\n        if self.residual:\n            x = self.drop_path(x) + identity\n        return x\n\n\nclass HG_Stage(nn.Module):\n    def __init__(\n            self,\n            in_chs,\n            mid_chs,\n            out_chs,\n            block_num,\n            layer_num,\n            downsample=True,\n            light_block=False,\n            kernel_size=3,\n            use_lab=False,\n            agg='se',\n            drop_path=0.,\n    ):\n        super().__init__()\n        self.downsample = downsample\n        if downsample:\n            self.downsample = ConvBNAct(\n                in_chs,\n                in_chs,\n                kernel_size=3,\n                stride=2,\n                groups=in_chs,\n                use_act=False,\n                use_lab=use_lab,\n            )\n        else:\n            self.downsample = nn.Identity()\n\n        blocks_list = []\n        for i in range(block_num):\n            blocks_list.append(\n                HG_Block(\n                    in_chs if i == 0 else out_chs,\n                    mid_chs,\n                    out_chs,\n                    layer_num,\n                    residual=False if i == 0 else True,\n                    kernel_size=kernel_size,\n                    light_block=light_block,\n                    use_lab=use_lab,\n                    agg=agg,\n                    drop_path=drop_path[i] if isinstance(drop_path, (list, tuple)) else drop_path,\n                )\n            )\n        self.blocks = nn.Sequential(*blocks_list)\n\n    def forward(self, x):\n        x = self.downsample(x)\n        x = self.blocks(x)\n        return x\n\n\n\n@register()\nclass HGNetv2(nn.Module):\n    \"\"\"\n    HGNetV2\n    Args:\n        stem_channels: list. Number of channels for the stem block.\n        stage_type: str. The stage configuration of HGNet. such as the number of channels, stride, etc.\n        use_lab: boolean. Whether to use LearnableAffineBlock in network.\n        lr_mult_list: list. Control the learning rate of different stages.\n    Returns:\n        model: nn.Layer. Specific HGNetV2 model depends on args.\n    \"\"\"\n\n    arch_configs = {\n        'B0': {\n            'stem_channels': [3, 16, 16],\n            'stage_config': {\n                # in_channels, mid_channels, out_channels, num_blocks, downsample, light_block, kernel_size, layer_num\n                \"stage1\": [16, 16, 64, 1, False, False, 3, 3],\n                \"stage2\": [64, 32, 256, 1, True, False, 3, 3],\n                \"stage3\": [256, 64, 512, 2, True, True, 5, 3],\n                \"stage4\": [512, 128, 1024, 1, True, True, 5, 3],\n            },\n            'url': 'file:///kaggle/working/PPHGNetV2_B0_stage1.pth'\n        },\n        'B1': {\n            'stem_channels': [3, 24, 32],\n            'stage_config': {\n                # in_channels, mid_channels, out_channels, num_blocks, downsample, light_block, kernel_size, layer_num\n                \"stage1\": [32, 32, 64, 1, False, False, 3, 3],\n                \"stage2\": [64, 48, 256, 1, True, False, 3, 3],\n                \"stage3\": [256, 96, 512, 2, True, True, 5, 3],\n                \"stage4\": [512, 192, 1024, 1, True, True, 5, 3],\n            },\n            'url': 'file:///kaggle/working/PPHGNetV2_B1_stage1.pth'\n        },\n        'B2': {\n            'stem_channels': [3, 24, 32],\n            'stage_config': {\n                # in_channels, mid_channels, out_channels, num_blocks, downsample, light_block, kernel_size, layer_num\n                \"stage1\": [32, 32, 96, 1, False, False, 3, 4],\n                \"stage2\": [96, 64, 384, 1, True, False, 3, 4],\n                \"stage3\": [384, 128, 768, 3, True, True, 5, 4],\n                \"stage4\": [768, 256, 1536, 1, True, True, 5, 4],\n            },\n            'url': 'file:///kaggle/working/PPHGNetV2_B2_stage1.pth'\n        },\n        'B3': {\n            'stem_channels': [3, 24, 32],\n            'stage_config': {\n                # in_channels, mid_channels, out_channels, num_blocks, downsample, light_block, kernel_size, layer_num\n                \"stage1\": [32, 32, 128, 1, False, False, 3, 5],\n                \"stage2\": [128, 64, 512, 1, True, False, 3, 5],\n                \"stage3\": [512, 128, 1024, 3, True, True, 5, 5],\n                \"stage4\": [1024, 256, 2048, 1, True, True, 5, 5],\n            },\n            'url': 'file:///kaggle/working/PPHGNetV2_B3_stage1.pth'\n        },\n        'B4': {\n            'stem_channels': [3, 32, 48],\n            'stage_config': {\n                # in_channels, mid_channels, out_channels, num_blocks, downsample, light_block, kernel_size, layer_num\n                \"stage1\": [48, 48, 128, 1, False, False, 3, 6],\n                \"stage2\": [128, 96, 512, 1, True, False, 3, 6],\n                \"stage3\": [512, 192, 1024, 3, True, True, 5, 6],\n                \"stage4\": [1024, 384, 2048, 1, True, True, 5, 6],\n            },\n            'url': 'file:///kaggle/working/PPHGNetV2_B4_stage1.pth'\n        },\n        'B5': {\n            'stem_channels': [3, 32, 64],\n            'stage_config': {\n                # in_channels, mid_channels, out_channels, num_blocks, downsample, light_block, kernel_size, layer_num\n                \"stage1\": [64, 64, 128, 1, False, False, 3, 6],\n                \"stage2\": [128, 128, 512, 2, True, False, 3, 6],\n                \"stage3\": [512, 256, 1024, 5, True, True, 5, 6],\n                \"stage4\": [1024, 512, 2048, 2, True, True, 5, 6],\n            },\n            'url': 'file:///kaggle/working/PPHGNetV2_B5_stage1.pth'\n        },\n        'B6': {\n            'stem_channels': [3, 48, 96],\n            'stage_config': {\n                # in_channels, mid_channels, out_channels, num_blocks, downsample, light_block, kernel_size, layer_num\n                \"stage1\": [96, 96, 192, 2, False, False, 3, 6],\n                \"stage2\": [192, 192, 512, 3, True, False, 3, 6],\n                \"stage3\": [512, 384, 1024, 6, True, True, 5, 6],\n                \"stage4\": [1024, 768, 2048, 3, True, True, 5, 6],\n            },\n            'url': 'file:///kaggle/working/PPHGNetV2_B6_stage1.pth'\n        },\n    }\n\n    def __init__(self,\n                 name,\n                 use_lab=False,\n                 return_idx=[1, 2, 3],\n                 freeze_stem_only=True,\n                 freeze_at=0,\n                 freeze_norm=True,\n                  pretrained=False,#pretrained=True,\n                 local_model_dir='./'):\n        super().__init__()\n        self.use_lab = use_lab\n        self.return_idx = return_idx\n\n        stem_channels = self.arch_configs[name]['stem_channels']\n        stage_config = self.arch_configs[name]['stage_config']\n        download_url = self.arch_configs[name]['url']\n\n        self._out_strides = [4, 8, 16, 32]\n        self._out_channels = [stage_config[k][2] for k in stage_config]\n\n        # stem\n        self.stem = StemBlock(\n                in_chs=stem_channels[0],\n                mid_chs=stem_channels[1],\n                out_chs=stem_channels[2],\n                use_lab=use_lab)\n\n        # stages\n        self.stages = nn.ModuleList()\n        for i, k in enumerate(stage_config):\n            in_channels, mid_channels, out_channels, block_num, downsample, light_block, kernel_size, layer_num = stage_config[\n                k]\n            self.stages.append(\n                HG_Stage(\n                    in_channels,\n                    mid_channels,\n                    out_channels,\n                    block_num,\n                    layer_num,\n                    downsample,\n                    light_block,\n                    kernel_size,\n                    use_lab))\n\n        if freeze_at >= 0:\n            self._freeze_parameters(self.stem)\n            if not freeze_stem_only:\n                for i in range(min(freeze_at + 1, len(self.stages))):\n                    self._freeze_parameters(self.stages[i])\n\n        if freeze_norm:\n            self._freeze_norm(self)\n\n        if pretrained:\n            RED, GREEN, RESET = \"\\033[91m\", \"\\033[92m\", \"\\033[0m\"\n            try:\n                model_path = local_model_dir + 'PPHGNetV2_' + name + '_stage1.pth'\n                if os.path.exists(model_path):\n                    state = torch.load(model_path, map_location='cpu')\n                    print(f\"Loaded stage1 {name} HGNetV2 from local file.\")\n                else:\n                    # If the file doesn't exist locally, download from the URL\n                    if torch.distributed.get_rank() == 0:\n                        print(GREEN + \"If the pretrained HGNetV2 can't be downloaded automatically. Please check your network connection.\" + RESET)\n                        print(GREEN + \"Please check your network connection. Or download the model manually from \" + RESET + f\"{download_url}\" + GREEN + \" to \" + RESET + f\"{local_model_dir}.\" + RESET)\n                        state = torch.hub.load_state_dict_from_url(download_url, map_location='cpu', model_dir=local_model_dir)\n                        torch.distributed.barrier()\n                    else:\n                        torch.distributed.barrier()\n                        state = torch.load(local_model_dir)\n\n                    print(f\"Loaded stage1 {name} HGNetV2 from URL.\")\n\n                self.load_state_dict(state)\n\n            except (Exception, KeyboardInterrupt) as e:\n                if torch.distributed.get_rank() == 0:\n                    print(f\"{str(e)}\")\n                    logging.error(RED + \"CRITICAL WARNING: Failed to load pretrained HGNetV2 model\" + RESET)\n                    logging.error(GREEN + \"Please check your network connection. Or download the model manually from \" \\\n                                + RESET + f\"{download_url}\" + GREEN + \" to \" + RESET + f\"{local_model_dir}.\" + RESET)\n                exit()\n\n\n\n\n    def _freeze_norm(self, m: nn.Module):\n        if isinstance(m, nn.BatchNorm2d):\n            m = FrozenBatchNorm2d(m.num_features)\n        else:\n            for name, child in m.named_children():\n                _child = self._freeze_norm(child)\n                if _child is not child:\n                    setattr(m, name, _child)\n        return m\n\n    def _freeze_parameters(self, m: nn.Module):\n        for p in m.parameters():\n            p.requires_grad = False\n\n    def forward(self, x):\n        x = self.stem(x)\n        outs = []\n        for idx, stage in enumerate(self.stages):\n            x = stage(x)\n            if idx in self.return_idx:\n                outs.append(x)\n        return outs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:36:51.902054Z","iopub.execute_input":"2025-05-25T13:36:51.902296Z","iopub.status.idle":"2025-05-25T13:36:51.915804Z","shell.execute_reply.started":"2025-05-25T13:36:51.902272Z","shell.execute_reply":"2025-05-25T13:36:51.915081Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install  -r /kaggle/working/DEIM/requirements.txt  --no-index --find-links=\"/kaggle/working/packages\" ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:36:51.916659Z","iopub.execute_input":"2025-05-25T13:36:51.917092Z","iopub.status.idle":"2025-05-25T13:37:45.456711Z","shell.execute_reply.started":"2025-05-25T13:36:51.917072Z","shell.execute_reply":"2025-05-25T13:37:45.456074Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!pip install  -r /kaggle/input/deim422-offline-packages/DEIM/requirements.txt  --no-index --find-links=\"/kaggle/input/deim422-offline-packages/packages\" ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:37:45.457782Z","iopub.execute_input":"2025-05-25T13:37:45.458527Z","iopub.status.idle":"2025-05-25T13:37:45.462080Z","shell.execute_reply.started":"2025-05-25T13:37:45.458500Z","shell.execute_reply":"2025-05-25T13:37:45.461346Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from glob import glob\nimport sys\nimport os\nfrom typing import Literal\n\nimport torch\nimport torch.nn as nn\nimport torchvision\nimport torchvision.transforms as T\nimport numpy as np\nfrom PIL import Image, ImageDraw\nimport pandas as pd\nimport cv2 \nfrom fastprogress import progress_bar as pb\nfrom tqdm import tqdm\nfrom scipy.spatial import distance\nfrom scipy.optimize import linear_sum_assignment\nimport networkx as nx\ntqdm.pandas()\n\nsys.path.append('/kaggle/working/DEIM')\nfrom engine.core import YAMLConfig","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:37:45.463049Z","iopub.execute_input":"2025-05-25T13:37:45.463623Z","iopub.status.idle":"2025-05-25T13:38:15.124984Z","shell.execute_reply.started":"2025-05-25T13:37:45.463603Z","shell.execute_reply":"2025-05-25T13:38:15.124434Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 0. Config","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\n#DEIM_CONFIG_FILEPATH = '/kaggle/input/my-deim-train-wts-demo/DEIM/configs/deim_dfine/deim_hgnetv2_l_coco_byu.yml'\n#DEIM_MODEL_FILEPATH = '/kaggle/input/byu-deim-toy-wts/best.pth'\n#IMAGE_SIZE = (384, 384)\n#DEIM_MODEL_FILEPATH = '/kaggle/input/my-deim-train-wts-demo/outputs/dfine_hgnetv2_l_coco/best_stg2.pth'\n#DEIM_MODEL_FILEPATH = '/kaggle/input/my-deim-train-wts-demo/outputs/dfine_hgnetv2_l_coco/best_stg1.pth'\n\nDEIM_CONFIG_FILEPATH = '/kaggle/input/my-deim-train-wts-demo/DEIM/configs/deim_dfine/deim_hgnetv2_x_coco_byu.yml'\nDEIM_MODEL_FILEPATH = '/kaggle/input/my-deim-train-wts-demo/output/dfine_hgnetv2_x_coco/best_stg1.pth'\nIMAGE_SIZE = (640, 640)\nimgsize=640\n\n#IMAGE_SIZE = (384,384)\n#imgsize=384\n\n#SCORE_TH_PRE = SCORE_TH_AGG = 0.825\n#score 0\n\n#SCORE_TH_PRE = SCORE_TH_AGG = 0.8\n\nSCORE_TH_PRE = SCORE_TH_AGG = 0.6\n\nGROUP_DIST_TH = 20.0\nMIN_DET_PER_GROUP = 1\nAGG_METHOD = 'score_highest'\n#AGG_METHOD = 'score_weighted_mean'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:56.646660Z","iopub.execute_input":"2025-05-25T13:39:56.646979Z","iopub.status.idle":"2025-05-25T13:39:56.651645Z","shell.execute_reply.started":"2025-05-25T13:39:56.646950Z","shell.execute_reply":"2025-05-25T13:39:56.650900Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 1. Data Preparation","metadata":{}},{"cell_type":"code","source":"BASE_IMAGE_DIR = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/\"\nTEST_IMAGE_DIR = os.path.join(BASE_IMAGE_DIR, \"test\")\ntest_tomo_dir_list = glob(f'{TEST_IMAGE_DIR}/*')\ntest_tomo_id_list = [d.split('/')[-1] for d in test_tomo_dir_list]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:56.652631Z","iopub.execute_input":"2025-05-25T13:39:56.652952Z","iopub.status.idle":"2025-05-25T13:39:56.677665Z","shell.execute_reply.started":"2025-05-25T13:39:56.652915Z","shell.execute_reply":"2025-05-25T13:39:56.677128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_images(tomo_id, train_or_test='test', resize_size=IMAGE_SIZE, loader='torchvision'):\n    assert loader in ['pil', 'torchvision']\n    image_dir = f'{BASE_IMAGE_DIR}/{train_or_test}/{tomo_id}'\n    image_files = sorted(glob(f'{image_dir}/*.*'))\n    df_image_files = pd.DataFrame({'filepath': image_files})\n    df_image_files['no'] = df_image_files['filepath'].map(lambda x: int(x.split('_')[-1].split('.')[0]))\n    df_image_files = df_image_files.sort_values(by='no', ascending=True)\n    # None : pil/torchvision resize results in slightly different values.\n    if loader == 'pil':\n        images = [Image.open(f).convert('L') for f in df_image_files['filepath']]\n        org_image_size = images[0].size  # (w, h)\n        if resize_size is not None:\n            images = [image.resize(resize_size) for image in images]\n        images = np.stack([np.asarray(image) for image in images])  # (n_frames, h, w)\n    elif loader == 'torchvision':\n        trainsforms = T.Resize(resize_size) if resize_size is not None else T.Compose([])\n        images = [torchvision.io.read_image(f) for f in df_image_files['filepath']]\n        org_image_size = (images[0].shape[2], images[0].shape[1])  # (w, h)\n        images = [trainsforms(image) for image in images]\n        images = torch.concatenate(images, dim=0)  # (n_frames, h, w)\n        images = images.numpy()\n    return images, df_image_files, org_image_size","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:57.005674Z","iopub.execute_input":"2025-05-25T13:39:57.005873Z","iopub.status.idle":"2025-05-25T13:39:57.012248Z","shell.execute_reply.started":"2025-05-25T13:39:57.005858Z","shell.execute_reply":"2025-05-25T13:39:57.011654Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %%time\n# images, df_image_files, org_image_size = load_images(tomo_id='tomo_003acc', loader='torchvision', resize_size=IMAGE_SIZE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:57.118859Z","iopub.execute_input":"2025-05-25T13:39:57.119321Z","iopub.status.idle":"2025-05-25T13:39:57.122162Z","shell.execute_reply.started":"2025-05-25T13:39:57.119304Z","shell.execute_reply":"2025-05-25T13:39:57.121497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %%time\n# images2, df_image_files2, org_image_size2 = load_images(tomo_id='tomo_003acc', loader='pil', resize_size=IMAGE_SIZE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:57.123351Z","iopub.execute_input":"2025-05-25T13:39:57.123602Z","iopub.status.idle":"2025-05-25T13:39:57.136499Z","shell.execute_reply.started":"2025-05-25T13:39:57.123588Z","shell.execute_reply":"2025-05-25T13:39:57.135906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Prepare DEIM Model","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nimport tempfile\nimport torch\nimport torch.distributed as dist\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.multiprocessing as mp\n\nfrom torch.nn.parallel import DistributedDataParallel as DDP\n\n# On Windows platform, the torch.distributed package only\n# supports Gloo backend, FileStore and TcpStore.\n# For FileStore, set init_method parameter in init_process_group\n# to a local file. Example as follow:\n# init_method=\"file:///f:/libtmp/some_file\"\n# dist.init_process_group(\n#    \"gloo\",\n#    rank=rank,\n#    init_method=init_method,\n#    world_size=world_size)\n# For TcpStore, same way as on Linux.\n\ndef setup(rank, world_size):\n    os.environ['MASTER_ADDR'] = 'localhost'\n    os.environ['MASTER_PORT'] = '12355'\n\n    # initialize the process group\n    dist.init_process_group(\"gloo\", rank=rank, world_size=world_size)\n\ndef cleanup():\n    dist.destroy_process_group()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:57.137064Z","iopub.execute_input":"2025-05-25T13:39:57.137253Z","iopub.status.idle":"2025-05-25T13:39:57.151581Z","shell.execute_reply.started":"2025-05-25T13:39:57.137239Z","shell.execute_reply":"2025-05-25T13:39:57.151092Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self, cfg):\n        super().__init__()\n        self.model = cfg.model.deploy()\n        self.postprocessor = cfg.postprocessor.deploy()\n\n    def forward(self, images, orig_target_sizes):\n        outputs = self.model(images)\n        outputs = self.postprocessor(outputs, orig_target_sizes)\n        return outputs\n\n\ndef prepare_deim_model(cfg_filepath: str, weight_filepath: str, device=device):\n    cfg = YAMLConfig(cfg_filepath, resume=weight_filepath)\n    checkpoint = torch.load(weight_filepath, map_location=device,weights_only=False)\n\n    #state = torch.load(weight_filepath, map_location=device,weights_only=True)\n    if 'ema' in checkpoint:\n         state = checkpoint['ema']['module']\n         \n    else:\n         state = checkpoint['model']\n\n    setup(0, 1)\n    # Load train mode state and convert to deploy mode\n    cfg.model.load_state_dict(state)\n    model = Model(cfg).to(device)\n    return model.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:57.152649Z","iopub.execute_input":"2025-05-25T13:39:57.152859Z","iopub.status.idle":"2025-05-25T13:39:57.172068Z","shell.execute_reply.started":"2025-05-25T13:39:57.152844Z","shell.execute_reply":"2025-05-25T13:39:57.171472Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(torch.__version__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:57.172623Z","iopub.execute_input":"2025-05-25T13:39:57.172865Z","iopub.status.idle":"2025-05-25T13:39:57.186302Z","shell.execute_reply.started":"2025-05-25T13:39:57.172851Z","shell.execute_reply":"2025-05-25T13:39:57.185578Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp /kaggle/input/pretrained-pphgnetv2-wts/* /kaggle/working/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:57.190159Z","iopub.execute_input":"2025-05-25T13:39:57.190398Z","iopub.status.idle":"2025-05-25T13:39:57.773429Z","shell.execute_reply.started":"2025-05-25T13:39:57.190383Z","shell.execute_reply":"2025-05-25T13:39:57.772599Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pwd","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:57.774976Z","iopub.execute_input":"2025-05-25T13:39:57.775223Z","iopub.status.idle":"2025-05-25T13:39:57.781768Z","shell.execute_reply.started":"2025-05-25T13:39:57.775203Z","shell.execute_reply":"2025-05-25T13:39:57.781069Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"deim_weights_dir='/kaggle/working/RT-DETR-main/D-FINE/weight/hgnetv2'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:57.782695Z","iopub.execute_input":"2025-05-25T13:39:57.782980Z","iopub.status.idle":"2025-05-25T13:39:57.793904Z","shell.execute_reply.started":"2025-05-25T13:39:57.782955Z","shell.execute_reply":"2025-05-25T13:39:57.793421Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.makedirs(deim_weights_dir, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:57.794691Z","iopub.execute_input":"2025-05-25T13:39:57.794900Z","iopub.status.idle":"2025-05-25T13:39:57.813003Z","shell.execute_reply.started":"2025-05-25T13:39:57.794878Z","shell.execute_reply":"2025-05-25T13:39:57.812273Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp /kaggle/input/pretrained-pphgnetv2-wts/*  /kaggle/working/RT-DETR-main/D-FINE/weight/hgnetv2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:57.814835Z","iopub.execute_input":"2025-05-25T13:39:57.815130Z","iopub.status.idle":"2025-05-25T13:39:58.162608Z","shell.execute_reply.started":"2025-05-25T13:39:57.815115Z","shell.execute_reply":"2025-05-25T13:39:58.161816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\nRANK = int(os.getenv('RANK', 0))\nLOCAL_RANK = int(os.getenv('LOCAL_RANK', -1))\nWORLD_SIZE = int(os.getenv('WORLD_SIZE', 1))\n\n# torch.distributed.init_process_group(backend=backend, init_method='env://')\ntorch.distributed.init_process_group( backend=\"gloo\",\n                        init_method=\"tcp://127.0.0.1:29500\",\n                        world_size=1,\n                        rank=0,)\ntorch.distributed.barrier()\n\nrank = torch.distributed.get_rank()\ntorch.cuda.set_device(rank)\ntorch.cuda.empty_cache()\nenabled_dist = True\n'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:58.163497Z","iopub.execute_input":"2025-05-25T13:39:58.163700Z","iopub.status.idle":"2025-05-25T13:39:58.169151Z","shell.execute_reply.started":"2025-05-25T13:39:58.163679Z","shell.execute_reply":"2025-05-25T13:39:58.168432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n'''\n\nlogger.info(\"Trying alternative initialization approach\")\n                        import tempfile\n\n                        temp_dir = tempfile.mkdtemp()\n                        file_path = os.path.join(temp_dir, \"shared_file\")\n\n                        store = torch.distributed.FileStore(file_path, 1)\n                        torch.distributed.init_process_group(\n                            backend=\"gloo\", store=store, rank=0, world_size=1\n                        )\n                        logger.info(\n                            \"Process group initialized successfully with FileStore\"\n                        )\n\n'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:58.169894Z","iopub.execute_input":"2025-05-25T13:39:58.170408Z","iopub.status.idle":"2025-05-25T13:39:58.182028Z","shell.execute_reply.started":"2025-05-25T13:39:58.170381Z","shell.execute_reply":"2025-05-25T13:39:58.181321Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"det_model = prepare_deim_model(\n    cfg_filepath=DEIM_CONFIG_FILEPATH,\n    weight_filepath=DEIM_MODEL_FILEPATH,\n\n   \n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:39:58.182750Z","iopub.execute_input":"2025-05-25T13:39:58.183011Z","iopub.status.idle":"2025-05-25T13:40:22.551880Z","shell.execute_reply.started":"2025-05-25T13:39:58.182983Z","shell.execute_reply":"2025-05-25T13:40:22.551100Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Run Detection","metadata":{}},{"cell_type":"code","source":"def rolling_mean_image(image: torch.Tensor, dim: int, window: int) -> torch.Tensor:\n    assert window % 2 == 1, \"Window size must be odd\"\n    n_dim = image.ndim\n\n    if dim != (n_dim - 1):\n        image = image.transpose(n_dim - 1, dim)  # move target dim to last\n\n    n_padding = (window - 1) // 2\n    pad_image_head = image[..., [0]].repeat([1] * (n_dim - 1) + [n_padding]).to(image)\n    pad_image_tail = image[..., [-1]].repeat([1] * (n_dim - 1) + [n_padding]).to(image)\n    image_padded = torch.cat([pad_image_head, image, pad_image_tail], dim=-1)\n\n    image_rolling_mean = image_padded.unfold(dimension=-1, size=window, step=1).mean(dim=-1)\n\n    if dim != (n_dim - 1):\n        image_rolling_mean = image_rolling_mean.transpose(n_dim - 1, dim)  # revert to original shape\n\n    assert image.shape == image_rolling_mean.shape\n    return image_rolling_mean","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:40:22.552921Z","iopub.execute_input":"2025-05-25T13:40:22.553184Z","iopub.status.idle":"2025-05-25T13:40:22.559111Z","shell.execute_reply.started":"2025-05-25T13:40:22.553158Z","shell.execute_reply":"2025-05-25T13:40:22.558504Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.inference_mode()\ndef inference_np_batch(model, np_image: np.ndarray, resize_size=IMAGE_SIZE):\n    if np_image.dtype == np.uint8:\n        np_image = np_image.astype(float) / 255\n    tensor_image = torch.tensor(np_image).permute(0, 3, 1, 2).float()  # (bs, h, w, ch) => (bs, ch, h, w)\n\n    transforms = T.Compose([\n        T.Resize(IMAGE_SIZE),\n    ])\n    im_data = transforms(tensor_image).to(device)\n    bs, ch, h, w = im_data.shape\n    orig_size = torch.tensor([[w, h]]).to(device)\n\n    output = model(im_data, orig_size)\n    labels, boxes, scores = output\n    return labels, boxes, scores\n\n\ndef filter_detection(labels, boxes, scores, thrh=0.4):\n    n_query1, =  labels.shape\n    n_query2, bbox_dim =  boxes.shape\n    n_query3, =  scores.shape\n    assert n_query1 == n_query2 == n_query3, (n_query1, n_query2, n_query3)\n    assert bbox_dim == 4\n    lab = labels[scores > thrh]\n    box = boxes[scores > thrh]\n    scrs = scores[scores > thrh]\n    return lab, box, scrs\n\n\n@torch.inference_mode()\ndef inference_tomo(model, tomo_id: str, batch_size: int = 4, th: float = 0.4) -> pd.DataFrame:\n    # 1. Load images for target tomo_id\n    images, df_image_files, org_image_size = load_images(tomo_id=tomo_id, loader='pil', resize_size=IMAGE_SIZE)\n    images = images.transpose(1, 2, 0)  # (n_frames, h, w) => (h, w, n_frames)\n    z_max = images.shape[-1] - 1\n    w_org, h_org = org_image_size\n\n    experimental = False\n    if experimental:\n        # calculate rolling mean along z-axis\n        images = rolling_mean_image(torch.tensor(images).float(), dim=2, window=21).numpy() / 255\n\n    # 2. Run detection on sliced 3ch images along z axis.\n    image_sliced_list = []\n    df_detection_list = []\n    z_center_list = list(range(1, z_max-1))\n    for z in pb(z_center_list):\n        image_sliced = images[:, :, z-1:z+1+1]  # (h, w, 3)\n        image_sliced_list.append(image_sliced)\n        if (len(image_sliced_list) >= batch_size) or (z == z_center_list[-1]):\n            image_sliced_batch = np.stack(image_sliced_list)  # (bs, h, w, 3)\n            labels, boxes, scores = inference_np_batch(model, image_sliced_batch)\n            for i in range(labels.shape[0]):\n                lab, box, scrs = filter_detection(labels[i], boxes[i], scores[i], th)\n                if len(lab) > 0:\n                    df_det = pd.DataFrame(data=box.cpu().numpy(), columns=['x1', 'y1', 'x2', 'y2'])\n                    df_det['z'] = z\n                    df_det['x_384'] = 0.5 * (df_det['x1'] + df_det['x2'])\n                    df_det['y_384'] = 0.5 * (df_det['y1'] + df_det['y2'])\n                    df_det['x_normed'] = df_det['x_384'] / imgsize\n                    df_det['y_normed'] = df_det['y_384'] / imgsize\n                    df_det['x'] = w_org * df_det['x_normed']\n                    df_det['y'] = h_org * df_det['y_normed']\n                    df_det['label'] = lab.cpu().tolist()\n                    df_det['score'] = scrs.cpu().tolist()\n                    df_det['tomo_id'] = tomo_id\n                    df_detection_list.append(df_det)\n            image_sliced_list = []\n    return pd.concat(df_detection_list) if len(df_detection_list) > 0 else pd.DataFrame([])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:40:22.559762Z","iopub.execute_input":"2025-05-25T13:40:22.560012Z","iopub.status.idle":"2025-05-25T13:40:22.577287Z","shell.execute_reply.started":"2025-05-25T13:40:22.559996Z","shell.execute_reply":"2025-05-25T13:40:22.576696Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Run detection for each tomo_id\ndf_det_list = []\n\nfor tomo_id in pb(test_tomo_id_list):\n    df_det_list.append(inference_tomo(det_model, tomo_id, th=SCORE_TH_PRE))\n\ndf_det_all = pd.concat(df_det_list)\ndf_det_all = df_det_all.reset_index(drop=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:40:22.578313Z","iopub.execute_input":"2025-05-25T13:40:22.578485Z","iopub.status.idle":"2025-05-25T13:42:16.970788Z","shell.execute_reply.started":"2025-05-25T13:40:22.578468Z","shell.execute_reply":"2025-05-25T13:42:16.970243Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Aggregate Detections","metadata":{}},{"cell_type":"code","source":"def aggregate_detection(\n    df_det: pd.DataFrame,\n    score_th: float,\n    group_dist_th: float = 5.0,\n    min_det_per_group: int = None,\n    agg_method: Literal['score_weighted_mean', 'score_highest'] = 'score_weighted_mean',\n) -> pd.DataFrame:\n    assert agg_method in ['score_weighted_mean', 'score_highest']\n    # pre-filter by score threshold\n    df_det = df_det[df_det['score'] >= score_th]\n    df_agg_det_tomo_list = []\n    for tomo_id, df_det_tomo in df_det.groupby('tomo_id'):\n        # calculate euclidean distance matrix (in voxel space) between each detections in this tomo_id\n        dist_mat = distance.cdist(df_det_tomo[['x', 'y', 'z']], df_det_tomo[['x', 'y', 'z']], metric='euclidean')\n        # calculate adjacency matrix based on distance matrix and threshold distance\n        adj_mat = (dist_mat <= group_dist_th).astype(int)\n        np.fill_diagonal(adj_mat, 0)\n        # group detections into connected graphs based on adjacency matrix\n        G = nx.from_numpy_array(adj_mat)\n        connected_components = list(nx.connected_components(G))\n        agg_det_dict_list = []\n        # Aggregate detections in each connected groups\n        for group_idx_set in connected_components:\n            df_det_grp = df_det_tomo.iloc[list(group_idx_set)]  # detections belonging to this group\n            if agg_method == 'score_weighted_mean':\n                z = (df_det_grp['z'] * df_det_grp['score']).sum() / df_det_grp['score'].sum()  # score weighted mean\n                y = (df_det_grp['y'] * df_det_grp['score']).sum() / df_det_grp['score'].sum()  # score weighted mean\n                x = (df_det_grp['x'] * df_det_grp['score']).sum() / df_det_grp['score'].sum()  # score weighted mean\n                score_mean = df_det_grp['score'].mean()\n            elif agg_method == 'score_highest':\n                z = df_det_grp.sort_values(by='score', ascending=False).iloc[0].z\n                y = df_det_grp.sort_values(by='score', ascending=False).iloc[0].y\n                x = df_det_grp.sort_values(by='score', ascending=False).iloc[0].x\n                score_mean = df_det_grp.sort_values(by='score', ascending=False).iloc[0].score\n            else:\n                raise ValueError(agg_method)\n            agg_det = {\n                'tomo_id': tomo_id,\n                'x': x,\n                'y': y,\n                'z': z,\n                'score_mean': score_mean,\n                'group_det_count': len(group_idx_set),  # detection count in this group\n            }\n            agg_det_dict_list.append(agg_det)\n        df_agg_det_tomo = pd.DataFrame(agg_det_dict_list)\n        if min_det_per_group is not None:\n            # delete the detections belonging to the groups that has detection count less than min_det_per_group\n            df_agg_det_tomo = df_agg_det_tomo[df_agg_det_tomo['group_det_count'] >= min_det_per_group]\n        # select highest (group_det_count, score_mean) group's aggregated detection as final detection for this tomo_id\n        if agg_method == 'score_weighted_mean':\n            order_by = ['group_det_count', 'score_mean']\n        elif agg_method == 'score_highest':\n            order_by = ['score_mean', 'group_det_count']\n        else:\n            raise ValueError(agg_method)\n        df_agg_det_tomo = df_agg_det_tomo.sort_values(by=order_by, ascending=False).iloc[:1]\n        df_agg_det_tomo_list.append(df_agg_det_tomo)\n    return pd.concat(df_agg_det_tomo_list)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:42:16.971635Z","iopub.execute_input":"2025-05-25T13:42:16.971919Z","iopub.status.idle":"2025-05-25T13:42:16.981560Z","shell.execute_reply.started":"2025-05-25T13:42:16.971892Z","shell.execute_reply":"2025-05-25T13:42:16.980976Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_det_agg = aggregate_detection(\n    df_det_all,\n    score_th=SCORE_TH_AGG,\n    group_dist_th=GROUP_DIST_TH,\n    min_det_per_group=MIN_DET_PER_GROUP,\n    agg_method=AGG_METHOD,\n)\nassert not df_det_agg['tomo_id'].duplicated().any()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:42:16.982274Z","iopub.execute_input":"2025-05-25T13:42:16.982476Z","iopub.status.idle":"2025-05-25T13:42:17.101884Z","shell.execute_reply.started":"2025-05-25T13:42:16.982461Z","shell.execute_reply":"2025-05-25T13:42:17.101219Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_det_agg","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:42:17.103989Z","iopub.execute_input":"2025-05-25T13:42:17.104170Z","iopub.status.idle":"2025-05-25T13:42:17.121499Z","shell.execute_reply.started":"2025-05-25T13:42:17.104155Z","shell.execute_reply":"2025-05-25T13:42:17.120773Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Submit","metadata":{}},{"cell_type":"code","source":"# no motor detected tomo_id list \nno_motor_tomo_id_list = list(set(test_tomo_id_list) - set(df_det_agg['tomo_id']))\nlen(no_motor_tomo_id_list)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:42:17.122393Z","iopub.execute_input":"2025-05-25T13:42:17.122588Z","iopub.status.idle":"2025-05-25T13:42:17.127337Z","shell.execute_reply.started":"2025-05-25T13:42:17.122572Z","shell.execute_reply":"2025-05-25T13:42:17.126637Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# no motor detected predictions\ndf_det_no_motor = pd.DataFrame({\n    'tomo_id': no_motor_tomo_id_list,\n    'Motor axis 0': [-1] * len(no_motor_tomo_id_list),\n    'Motor axis 1': [-1] * len(no_motor_tomo_id_list),\n    'Motor axis 2': [-1] * len(no_motor_tomo_id_list),\n})\n# motor detected predictions\ndf_det_agg = df_det_agg.rename(\n    columns={'z': 'Motor axis 0', 'y': 'Motor axis 1', 'x': 'Motor axis 2'}\n)[['tomo_id', 'Motor axis 0', 'Motor axis 1', 'Motor axis 2']]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:42:17.128057Z","iopub.execute_input":"2025-05-25T13:42:17.128292Z","iopub.status.idle":"2025-05-25T13:42:17.143244Z","shell.execute_reply.started":"2025-05-25T13:42:17.128272Z","shell.execute_reply":"2025-05-25T13:42:17.142684Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"display(df_det_no_motor)\ndisplay(df_det_agg)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:42:17.144017Z","iopub.execute_input":"2025-05-25T13:42:17.144258Z","iopub.status.idle":"2025-05-25T13:42:17.169702Z","shell.execute_reply.started":"2025-05-25T13:42:17.144237Z","shell.execute_reply":"2025-05-25T13:42:17.169071Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_submission = pd.concat([df_det_agg, df_det_no_motor])\nassert set(df_submission['tomo_id']) == set(test_tomo_id_list)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:42:17.170338Z","iopub.execute_input":"2025-05-25T13:42:17.170504Z","iopub.status.idle":"2025-05-25T13:42:17.179757Z","shell.execute_reply.started":"2025-05-25T13:42:17.170491Z","shell.execute_reply":"2025-05-25T13:42:17.179136Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_submission.to_csv('submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-25T13:42:17.180435Z","iopub.execute_input":"2025-05-25T13:42:17.180785Z","iopub.status.idle":"2025-05-25T13:42:17.200376Z","shell.execute_reply.started":"2025-05-25T13:42:17.180752Z","shell.execute_reply":"2025-05-25T13:42:17.199708Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}