{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":10457442,"sourceType":"datasetVersion","datasetId":6473663},{"sourceId":10665419,"sourceType":"datasetVersion","datasetId":6472787},{"sourceId":214918197,"sourceType":"kernelVersion"}],"dockerImageVersionId":30823,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 22nd place solution\n\n* CV (TS_6_4): 0.8214\n* Public LB: 0.767\n* Private LB: 0.761\n\n\n```\n'TS_6_4'\n+----+---------------------+-----+-----+-------+--------+------+-------------+----------+-----------+----------+\n|    | particle_type       |   P |   T |   hit |   miss |   fp |   precision |   recall |   f-beta4 |   weight |\n|----+---------------------+-----+-----+-------+--------+------+-------------+----------+-----------+----------|\n|  0 | apo-ferritin        | 102 |  58 |    58 |      0 |   44 |    0.568627 | 1        |  0.957282 |        1 |\n|  1 | beta-amylase        |   0 |   9 |     0 |      9 |    0 |    0        | 0        |  0        |        0 |\n|  2 | beta-galactosidase  |  45 |  12 |    11 |      1 |   34 |    0.244444 | 0.916667 |  0.78903  |        2 |\n|  3 | ribosome            | 126 |  74 |    69 |      5 |   57 |    0.547619 | 0.932432 |  0.89542  |        1 |\n|  4 | thyroglobulin       | 102 |  30 |    26 |      4 |   76 |    0.254902 | 0.866667 |  0.75945  |        2 |\n|  5 | virus-like-particle |  10 |  10 |     8 |      2 |    2 |    0.8      | 0.8      |  0.8      |        1 |\n+----+---------------------+-----+-----+-------+--------+------+-------------+----------+-----------+----------+\n```","metadata":{}},{"cell_type":"markdown","source":"## Import libraries","metadata":{}},{"cell_type":"code","source":"from concurrent.futures import ProcessPoolExecutor\nimport copy\nimport multiprocessing\nimport math\nimport os\nfrom pathlib import Path\nimport  shutil\nimport time\nfrom typing import Optional\n\nimport albumentations as A\nimport cc3d\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport zarr\nfrom omegaconf import OmegaConf\nfrom torch.nn import Conv2d, Dropout, LayerNorm, Linear, Softmax\nfrom torch.nn.modules.utils import _pair\nfrom scipy.ndimage import labeled_comprehension","metadata":{"execution":{"iopub.status.busy":"2025-02-06T03:55:56.010487Z","iopub.execute_input":"2025-02-06T03:55:56.010792Z","iopub.status.idle":"2025-02-06T03:56:33.351901Z","shell.execute_reply.started":"2025-02-06T03:55:56.010761Z","shell.execute_reply":"2025-02-06T03:56:33.350985Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"class CONFIG:\n    ############## Data Config ##############\n    STATIC_DIR = \"/kaggle/input/czii-cryo-et-object-identification/test/static/ExperimentRuns\"\n    PERCENTILE = 99.5\n\n    ############## Model Config ##############\n    MODEL_3D_INFOS = [\n        # tf_efficientnet_b4.ns_jft_in1k\n        dict(\n            config_path=\"/kaggle/input/czii-2024-models/0139.yaml\",\n            ckpt_path=\"/kaggle/input/czii-2024-models/exp0139_fold0_best_val_lb_score.pth\",\n            is_half=True,\n            is_hflip=True,\n            is_vflip=False,\n            is_dflip=False,\n        ),\n        # tf_efficientnet_b4.ns_jft_in1k (pretrain)\n        dict(\n            config_path=\"/kaggle/input/czii-2024-models/0150.yaml\",\n            ckpt_path=\"/kaggle/input/czii-2024-models/exp0150_fold0_best_val_lb_score.pth\",\n            is_half=True,\n            is_hflip=True,\n            is_vflip=False,\n            is_dflip=False,\n        ),\n        # resnet50.fb_ssl_yfcc100m_ft_in1k (pretrain)\n        dict(\n            config_path=\"/kaggle/input/czii-2024-models/0156.yaml\",\n            ckpt_path=\"/kaggle/input/czii-2024-models/exp0156_fold0_latest.pth\",\n            is_half=True,\n            is_hflip=True,\n            is_vflip=False,\n            is_dflip=False,\n        ),\n        # tf_efficientnet_b3.ns_jft_in1k (pretrain)\n        dict(\n            config_path=\"/kaggle/input/czii-2024-models/0159.yaml\",\n            ckpt_path=\"/kaggle/input/czii-2024-models/exp0159_fold0_latest.pth\",\n            is_half=True,\n            is_hflip=True,\n            is_vflip=False,\n            is_dflip=False,\n        ),\n        # seresnet33ts.ra2_in1k (pretrain)\n        dict(\n            config_path=\"/kaggle/input/czii-2024-models/0182.yaml\",\n            ckpt_path=\"/kaggle/input/czii-2024-models/exp0182_fold1_best_val_lb_score.pth\",\n            is_half=True,\n            is_hflip=True,\n            is_vflip=False,\n            is_dflip=False,\n        ),\n        # eca_nfnet_l0.ra2_in1k\n        dict(\n            config_path=\"/kaggle/input/czii-2024-models/0185.yaml\",\n            ckpt_path=\"/kaggle/input/czii-2024-models/exp0185_fold1_best_val_lb_score.pth\",\n            is_half=False,\n            is_hflip=True,\n            is_vflip=False,\n            is_dflip=False,\n        ),\n        # tf_efficientnet_b5.ns_jft_in1k\n        dict(\n            config_path=\"/kaggle/input/czii-2024-models/0186.yaml\",\n            ckpt_path=\"/kaggle/input/czii-2024-models/exp0186_fold1_best_val_lb_score.pth\",\n            is_half=True,\n            is_hflip=True,\n            is_vflip=False,\n            is_dflip=False,\n        ),\n    ]\n\n    REMOVER_INFOS = [\n        dict(\n            particle_type=\"apo-ferritin\",\n            image_size=96,\n            volume_size=16,\n            threshold=0.15,\n            path_list=[\n                dict(\n                    config=\"/kaggle/input/czii-2024-models/0191.yaml\",\n                    ckpts=[\n                        \"/kaggle/input/czii-2024-models/exp0191_fold0_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0191_fold1_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0191_fold2_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0191_fold3_latest.pth\",\n                    ]\n                ),\n                dict(\n                    config=\"/kaggle/input/czii-2024-models/0196.yaml\",\n                    ckpts=[\n                        \"/kaggle/input/czii-2024-models/exp0196_fold0_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0196_fold1_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0196_fold2_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0196_fold3_latest.pth\",\n                    ]\n                ),\n                dict(\n                    config=\"/kaggle/input/czii-2024-models/0201.yaml\",\n                    ckpts=[\n                        \"/kaggle/input/czii-2024-models/exp0201_fold0_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0201_fold1_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0201_fold2_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0201_fold3_latest.pth\",\n                    ]\n                ),\n            ]\n        ),\n        dict(\n            particle_type=\"ribosome\",\n            image_size=96,\n            volume_size=38,\n            threshold=0.15,\n            path_list=[\n                dict(\n                    config=\"/kaggle/input/czii-2024-models/0193.yaml\",\n                    ckpts=[\n                        \"/kaggle/input/czii-2024-models/exp0193_fold0_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0193_fold1_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0193_fold2_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0193_fold3_latest.pth\",\n                    ]\n                ),\n                dict(\n                    config=\"/kaggle/input/czii-2024-models/0198.yaml\",\n                    ckpts=[\n                        \"/kaggle/input/czii-2024-models/exp0198_fold0_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0198_fold1_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0198_fold2_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0198_fold3_latest.pth\",\n                    ]\n                ),\n                dict(\n                    config=\"/kaggle/input/czii-2024-models/0203.yaml\",\n                    ckpts=[\n                        \"/kaggle/input/czii-2024-models/exp0203_fold0_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0203_fold1_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0203_fold2_latest.pth\",\n                        \"/kaggle/input/czii-2024-models/exp0203_fold3_latest.pth\",\n                    ]\n                ),\n            ]\n        ),\n    ]\n\n    ############## Inference Config ##############\n    MAX_WORKERS = 2\n    CLASS_INFOS = [\n        dict(id=1, name=\"apo-ferritin\", radius=60, color=(255, 0, 0)),\n        dict(id=2, name=\"beta-amylase\", radius=65, color=(0, 255, 0)),\n        dict(id=3, name=\"beta-galactosidase\", radius=90, color=(0, 0, 255)),\n        dict(id=4, name=\"ribosome\", radius=150, color=(255, 255, 0)),\n        dict(id=5, name=\"thyroglobulin\", radius=130, color=(255, 0, 255)),\n        dict(id=6, name=\"virus-like-particle\", radius=135, color=(0, 255, 255)),\n    ]\n    CENTROID_POST_PROCESS_SETTINGS = dict(\n        class_infos=CLASS_INFOS,\n        thresholds=[0.2, 1.0, 0.2, 0.2, 0.2, 0.2],\n        blob_thresholds=[0, 0, 0, 0, 0, 0],\n        connectivity=26,\n        voxel_size=10.012,\n    )\n    WPF_POST_PROCESS_SETTINGS = dict(\n        class_infos=CLASS_INFOS,\n        thresholds=[0.1, 1.0, 0.3, 0.1, 0.3, 0.5],\n    )\n\nCONFIG.ZARR_FILES = sorted(Path(CONFIG.STATIC_DIR).glob(\"**/denoised.zarr\"))\n\n# reference: https://qiita.com/rho-guy/items/1d31572a5cbf12c27e53\n# (1)Notebook, (2)Save&Run, (3)Submit\nif os.environ.get('KAGGLE_KERNEL_RUN_TYPE', 'Interactive') == 'Interactive':\n    print(\"Running on (1) Notebook...\")\n    CONFIG.ENV = \"Notebook\"\nelif os.environ.get('KAGGLE_KERNEL_RUN_TYPE','') == 'Batch':\n    if len(CONFIG.ZARR_FILES) == 3:\n        print(\"Running on (2) Save&Run\")\n        CONFIG.ENV = \"SaveRun\"\n    else:\n        print(\"Running on (3) Submit\")\n        CONFIG.ENV = \"Submit\"\n\n\nif CONFIG.ENV == \"SaveRun\":\n    # SaveRunのときは30件のTomogramを使用して耐久テストを行う\n    counter = 1\n    new_zarr_files = []\n    for _ in range(10):\n        for zarr_file in CONFIG.ZARR_FILES:\n            new_zarr_file = Path(\"./temp\") / f\"TS_77_{counter}\" / \"VoxelSpacing10.000\" / \"denoised.zarr\"\n            new_zarr_file.parent.mkdir(parents=True, exist_ok=True)\n            shutil.copytree(zarr_file, new_zarr_file)\n            new_zarr_files.append(new_zarr_file)\n            counter += 1\n    CONFIG.ZARR_FILES.extend(new_zarr_files)","metadata":{"execution":{"iopub.status.busy":"2025-02-06T03:56:33.352790Z","iopub.execute_input":"2025-02-06T03:56:33.353223Z","iopub.status.idle":"2025-02-06T03:56:33.404834Z","shell.execute_reply.started":"2025-02-06T03:56:33.353201Z","shell.execute_reply":"2025-02-06T03:56:33.404176Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_config(config_path: str, dot_list: list) -> dict:\n    config_omega_from_yaml = OmegaConf.load(config_path)\n    config_omega_from_args = OmegaConf.from_dotlist(dot_list)\n    config_omega = OmegaConf.merge(config_omega_from_yaml, config_omega_from_args)\n    config = OmegaConf.to_container(config_omega, resolve=True)  # DictConfig -> dict\n    return config","metadata":{"execution":{"iopub.status.busy":"2025-02-06T03:56:33.405686Z","iopub.execute_input":"2025-02-06T03:56:33.405986Z","iopub.status.idle":"2025-02-06T03:56:33.409861Z","shell.execute_reply.started":"2025-02-06T03:56:33.405954Z","shell.execute_reply":"2025-02-06T03:56:33.409063Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def normalize_parcentile(volume: np.ndarray, percentile: float) -> np.ndarray:\n    # https://www.kaggle.com/code/itsuki9180/czii-making-datasets-for-yolo\n    lower, upper = np.percentile(volume, [100 - percentile, percentile])\n    volume = np.clip(volume, lower, upper)\n    volume = (volume - volume.min()) / (volume.max() - volume.min() + 1e-12)\n    return volume\n\n\ndef load_tomogram(zarr_path: Path) -> np.ndarray:\n    z = zarr.open(zarr_path, mode='r')\n    volume = np.array(z[0])\n    d, h, w = volume.shape\n    if (d, h, w) != (184, 630, 630):\n        raise ValueError(f\"Shape is not (184, 630, 630): {zarr_path}\")\n    return normalize_parcentile(volume, CONFIG.PERCENTILE)","metadata":{"execution":{"iopub.status.busy":"2025-02-06T03:56:33.410825Z","iopub.execute_input":"2025-02-06T03:56:33.411115Z","iopub.status.idle":"2025-02-06T03:56:33.423418Z","shell.execute_reply.started":"2025-02-06T03:56:33.411084Z","shell.execute_reply":"2025-02-06T03:56:33.422733Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VolumePatchExtractor:\n    \"\"\"\n    This class tiles a 3D volume in (D, H, W) format into multiple overlapping\n    patches based on the specified tile size and overlap size.\n\n    Attributes:\n        tile_size (tuple): (tile_d, tile_h, tile_w)\n        overlap_size (tuple): (overlap_d, overlap_h, overlap_w)\n        weight_mode (str): \"constant\" or \"linear\"\n    \"\"\"\n\n    def __init__(\n        self,\n        tile_size: tuple[int, int, int] = (16, 256, 256),\n        overlap_size: tuple[int, int, int] = (4, 32, 32),\n        weight_mode: str = \"constant\",  # \"constant\" or \"linear\"\n    ):\n        \"\"\"\n        Constructor\n\n        Args:\n            tile_size (tuple): (tile_d, tile_h, tile_w)\n            overlap_size (tuple): (overlap_d, overlap_h, overlap_w)\n            weight_mode (str): \"constant\" or \"linear\"\n        \"\"\"\n        self.tile_size = tile_size\n        self.overlap_size = overlap_size\n        if weight_mode not in [\"constant\", \"linear\"]:\n            raise ValueError(\"weight_mode must be either 'constant' or 'linear'\")\n        self.weight_mode = weight_mode\n\n    def _create_1d_linear_weight(\n        self,\n        mode: str,\n        size: int,\n        overlap: int,\n        is_min_edge: bool,\n        is_max_edge: bool,\n    ) -> np.ndarray:\n        \"\"\"\n        Create a 1D linear or constant weight vector.\n\n        Args:\n            mode (str): \"constant\" or \"linear\"\n            size (int): Length of the dimension (e.g. tile_d, tile_h, or tile_w)\n            overlap (int): Overlap size in this dimension\n            is_min_edge (bool): Whether this patch is at the \"start\" of the dimension\n            is_max_edge (bool): Whether this patch is at the \"end\" of the dimension\n\n        Returns:\n            w (np.ndarray): 1D array of weights, shape=(size,)\n        \"\"\"\n        if mode == \"linear\":\n            w = np.zeros(size, dtype=np.float32)\n            for i in range(size):\n                dist_min = i\n                dist_max = size - 1 - i\n                # 左右または前後のうち小さい方の距離をとる\n                dist_edge = min(dist_min, dist_max)\n\n                # ボリューム端のタイルは減衰しないようにする\n                if is_min_edge and (dist_min < overlap):\n                    w[i] = 1.0\n                elif is_max_edge and (dist_max < overlap):\n                    w[i] = 1.0\n                else:\n                    # 端でなければ線形に減衰\n                    if dist_edge >= overlap:\n                        w[i] = 1.0\n                    else:\n                        w[i] = dist_edge / overlap\n        else:  # \"constant\"\n            w = np.ones(size, dtype=np.float32)\n\n        return w\n\n    def _create_3d_weight_map(\n        self,\n        mode: str,\n        tile_d: int,\n        tile_h: int,\n        tile_w: int,\n        overlap_d: int,\n        overlap_h: int,\n        overlap_w: int,\n        is_front_edge: bool,\n        is_back_edge: bool,\n        is_top_edge: bool,\n        is_bottom_edge: bool,\n        is_left_edge: bool,\n        is_right_edge: bool,\n    ) -> np.ndarray:\n        \"\"\"\n        Create a 3D weight map of shape (tile_d, tile_h, tile_w).\n\n        Args:\n            mode (str): \"constant\" or \"linear\"\n            tile_d, tile_h, tile_w (int): patch size along each dimension\n            overlap_d, overlap_h, overlap_w (int): overlap sizes\n            is_front_edge (bool): patch is at the front-most region (z=0 side)\n            is_back_edge (bool):  patch is at the back-most region  (z=max side)\n            is_top_edge (bool):   patch is at the top    (y=0 side)\n            is_bottom_edge (bool):patch is at the bottom (y=max side)\n            is_left_edge (bool):  patch is at the left   (x=0 side)\n            is_right_edge (bool): patch is at the right  (x=max side)\n\n        Returns:\n            weight (np.ndarray): 3D array of weights in shape=(tile_d, tile_h, tile_w)\n        \"\"\"\n        if self.weight_mode == \"constant\":\n            return np.ones((tile_d, tile_h, tile_w), dtype=np.float32)\n        else:\n            # Z方向ウェイト\n            weight_z = self._create_1d_linear_weight(\n                mode=mode,\n                size=tile_d,\n                overlap=overlap_d,\n                is_min_edge=is_front_edge,\n                is_max_edge=is_back_edge,\n            )\n            # Y方向ウェイト\n            weight_y = self._create_1d_linear_weight(\n                mode=mode,\n                size=tile_h,\n                overlap=overlap_h,\n                is_min_edge=is_top_edge,\n                is_max_edge=is_bottom_edge,\n            )\n            # X方向ウェイト\n            weight_x = self._create_1d_linear_weight(\n                mode=mode,\n                size=tile_w,\n                overlap=overlap_w,\n                is_min_edge=is_left_edge,\n                is_max_edge=is_right_edge,\n            )\n            # 3次元の外積\n            # weight_z: (tile_d,) -> shape=(tile_d, 1, 1)\n            # weight_y: (tile_h,) -> shape=(1, tile_h, 1)\n            # weight_x: (tile_w,) -> shape=(1, 1, tile_w)\n            weight_3d = weight_z[:, None, None] * weight_y[None, :, None] * weight_x[None, None, :]\n\n        return weight_3d\n\n    def extract(self, volume: np.ndarray) -> tuple[list[np.ndarray], list[np.ndarray], list[tuple[int]]]:\n        \"\"\"\n        Splits (tiles) a volume of shape (D, H, W) into patches of size (tile_d, tile_h, tile_w)\n        according to the tile size and overlap size.\n\n        Args:\n            volume (np.ndarray): Volume in (D, H, W) format.\n\n        Returns:\n            tiled_volumes (list of np.ndarray):\n                List of patches, each with shape (tile_d, tile_h, tile_w).\n            tile_weights (list of np.ndarray):\n                List of 3D weight maps for each patch, shape (tile_d, tile_h, tile_w).\n            tile_coords (list of tuple):\n                List of coordinates for each patch, in (z_min, y_min, x_min, z_max, y_max, x_max).\n        \"\"\"\n        # Get the depth, height, width\n        d, h, w = volume.shape\n\n        tile_d, tile_h, tile_w = self.tile_size\n        overlap_d, overlap_h, overlap_w = self.overlap_size\n\n        # Calculate the slide steps\n        step_d = tile_d - overlap_d\n        step_h = tile_h - overlap_h\n        step_w = tile_w - overlap_w\n\n        # Prepare lists to store results\n        tiles = []\n        tile_weights = []\n        tile_coords = []\n\n        # Slide along depth (z)\n        z_pos = 0\n        while z_pos < d:\n            if z_pos + tile_d > d:\n                z_pos = d - tile_d  # adjust for last patch if needed\n            # Slide along height (y)\n            y_pos = 0\n            while y_pos < h:\n                if y_pos + tile_h > h:\n                    y_pos = h - tile_h\n                # Slide along width (x)\n                x_pos = 0\n                while x_pos < w:\n                    if x_pos + tile_w > w:\n                        x_pos = w - tile_w\n\n                    # Define the bounding box\n                    z_min, y_min, x_min = z_pos, y_pos, x_pos\n                    z_max, y_max, x_max = z_pos + tile_d, y_pos + tile_h, x_pos + tile_w\n\n                    # Extract the patch\n                    patch = volume[z_min:z_max, y_min:y_max, x_min:x_max]\n\n                    # Determine edge flags\n                    is_front_edge = z_min == 0\n                    is_back_edge = z_max == d\n                    is_top_edge = y_min == 0\n                    is_bottom_edge = y_max == h\n                    is_left_edge = x_min == 0\n                    is_right_edge = x_max == w\n\n                    patch_weight = self._create_3d_weight_map(\n                        mode=self.weight_mode,\n                        tile_d=patch.shape[0],\n                        tile_h=patch.shape[1],\n                        tile_w=patch.shape[2],\n                        overlap_d=overlap_d,\n                        overlap_h=overlap_h,\n                        overlap_w=overlap_w,\n                        is_front_edge=is_front_edge,\n                        is_back_edge=is_back_edge,\n                        is_top_edge=is_top_edge,\n                        is_bottom_edge=is_bottom_edge,\n                        is_left_edge=is_left_edge,\n                        is_right_edge=is_right_edge,\n                    )\n\n                    tiles.append(patch)\n                    tile_weights.append(patch_weight)\n                    tile_coords.append((z_min, y_min, x_min, z_max, y_max, x_max))\n\n                    if x_pos + tile_w >= w:\n                        break\n                    else:\n                        x_pos += step_w\n                if y_pos + tile_h >= h:\n                    break\n                else:\n                    y_pos += step_h\n            if z_pos + tile_d >= d:\n                break\n            else:\n                z_pos += step_d\n\n        return tiles, tile_weights, tile_coords\n","metadata":{"execution":{"iopub.status.busy":"2025-02-06T03:56:33.425428Z","iopub.execute_input":"2025-02-06T03:56:33.425663Z","iopub.status.idle":"2025-02-06T03:56:33.440699Z","shell.execute_reply.started":"2025-02-06T03:56:33.425643Z","shell.execute_reply":"2025-02-06T03:56:33.439925Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3D Semantic Segmentation","metadata":{}},{"cell_type":"markdown","source":"### 2D-3D Semantic Segmentation Model\n\nreference: https://www.kaggle.com/code/hengck23/3d-unet-using-2d-image-encoder","metadata":{}},{"cell_type":"code","source":"class MyDecoderBlock3d(nn.Module):\n    \"\"\"\n    3D Decoder block implementing skip connections and progressive upsampling.\n    Extends the 2D decoder block concept to handle volumetric data.\n\n    Attributes:\n        conv1 (nn.Sequential): First 3D convolution block with batch norm and ReLU\n        attention1 (nn.Module): Optional attention mechanism after skip connection\n        conv2 (nn.Sequential): Second 3D convolution block with batch norm and ReLU\n        attention2 (nn.Module): Optional attention mechanism after second conv\n    \"\"\"\n\n    def __init__(self, in_channel: int, skip_channel: int, out_channel: int, spetial_scaling: int, depth_scaling: int):\n        \"\"\"\n        Initialize the 3D decoder block.\n\n        Args:\n            in_channel (int): Number of input channels\n            skip_channel (int): Number of channels in skip connection\n            out_channel (int): Number of output channels\n            spetial_scaling (int): Scaling factor for spatial dimensions upsampling\n            depth_scaling (int): Scaling factor for depth dimension upsampling\n        \"\"\"\n        super().__init__()\n        self.conv1 = nn.Sequential(\n            nn.Conv3d(in_channel + skip_channel, out_channel, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm3d(out_channel),\n            nn.ReLU(inplace=True),\n        )\n        self.attention1 = nn.Identity()\n        self.conv2 = nn.Sequential(\n            nn.Conv3d(out_channel, out_channel, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm3d(out_channel),\n            nn.ReLU(inplace=True),\n        )\n        self.attention2 = nn.Identity()\n        self.spetial_scaling = spetial_scaling\n        self.depth_scaling = depth_scaling\n\n    def forward(self, x: torch.Tensor, skip: Optional[torch.Tensor] = None) -> torch.Tensor:\n        \"\"\"\n        Forward pass of the 3D decoder block.\n\n        Args:\n            x (torch.Tensor): Input tensor\n            skip (torch.Tensor, optional): Skip connection tensor\n\n        Returns:\n            torch.Tensor: Processed feature volume\n        \"\"\"\n        # Upsample with different scaling for depth dimension\n        x = F.interpolate(x, scale_factor=(self.depth_scaling, self.spetial_scaling, self.spetial_scaling), mode=\"nearest\")\n        if skip is not None:\n            x = torch.cat([x, skip], dim=1)\n            x = self.attention1(x)\n        x = self.conv1(x)\n        x = self.conv2(x)\n        x = self.attention2(x)\n        return x\n\n\nclass MyUnetDecoder3d(nn.Module):\n    \"\"\"\n    3D UNet decoder implementation with skip connections.\n    Handles volumetric data with separate scaling for depth dimension.\n\n    Attributes:\n        center (nn.Module): Optional center processing block\n        block (nn.ModuleList): List of 3D decoder blocks\n    \"\"\"\n\n    def __init__(\n        self,\n        in_channel: int,\n        skip_channels: list[int],\n        out_channels: list[int],\n        spetial_scalings: list[int],\n        depth_scalings: list[int],\n    ):\n        \"\"\"\n        Initialize the 3D UNet decoder.\n\n        Args:\n            in_channel (int): Number of input channels\n            skip_channels (list): List of skip connection channels\n            out_channels (list): List of output channels for each block\n            spetial_scalings (list): List of spatial scaling factors for each block\n            depth_scalings (list): List of depth scaling factors for each block\n        \"\"\"\n        super().__init__()\n        self.center = nn.Identity()\n\n        i_channel = [\n            in_channel,\n        ] + out_channels[:-1]\n        s_channel = skip_channels\n        o_channel = out_channels\n        block = [MyDecoderBlock3d(i, s, o, ss, ds) for i, s, o, ss, ds in zip(i_channel, s_channel, o_channel, spetial_scalings, depth_scalings)]\n        self.block = nn.ModuleList(block)\n\n    def forward(self, feature: torch.Tensor, skip: list[torch.Tensor]):\n        \"\"\"\n        Forward pass of the 3D UNet decoder.\n\n        Args:\n            feature (torch.Tensor): Input feature tensor\n            skip (list): List of skip connection tensors\n            depth_scaling (list): List of depth scaling factors for each block\n\n        Returns:\n            tuple: (Final output tensor, List of intermediate decoder outputs)\n        \"\"\"\n        d = self.center(feature)\n        decode = []\n        for i, block in enumerate(self.block):\n            s = skip[i]\n            d = block(d, s)\n            decode.append(d)\n        last = d\n        return last, decode\n\n\nclass DepthDownsampleBlock(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int, depth_scaling: int):\n        super(DepthDownsampleBlock, self).__init__()\n        self.conv = nn.Conv3d(\n            in_channels=in_channels,\n            out_channels=out_channels,\n            kernel_size=(depth_scaling, 1, 1),\n            stride=(depth_scaling, 1, 1),\n            padding=0,\n            bias=False,\n        )\n        self.bn = nn.BatchNorm3d(out_channels)\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x = self.conv(x)\n        x = self.bn(x)\n        x = self.relu(x)\n        return x\n\n\nclass My2DEncoder(nn.Module):\n    def __init__(self, timm_model: dict, depth_scalings: list[int]):\n        super().__init__()\n        self.out_indices = timm_model.pop(\"out_indices\")\n        self.in_chans = timm_model[\"in_chans\"]\n        self.depth_scalings = depth_scalings\n        timm_encoder = timm.create_model(**timm_model)\n        self.out_channels = timm_encoder.feature_info.channels()\n        self.out_channels = [self.out_channels[i] for i in self.out_indices]\n\n        if len(self.out_channels) != len(self.depth_scalings):\n            raise ValueError(\"Length of out_indices and depth_scalings must be the same\")\n\n        if timm_model[\"model_name\"].startswith(\"resnet\") or timm_model[\"model_name\"].startswith(\"resnest\") or timm_model[\"model_name\"].startswith(\"seresnext\"):\n            self.spatial_donwsample_blocks = nn.ModuleList(\n                [\n                    # 1/2\n                    nn.Sequential(\n                        timm_encoder.conv1,\n                        timm_encoder.bn1,\n                        timm_encoder.act1,\n                    ),\n                    # 1/4\n                    nn.Sequential(\n                        timm_encoder.maxpool,\n                        timm_encoder.layer1,\n                    ),\n                    # 1/8\n                    timm_encoder.layer2,\n                    # 1/16\n                    timm_encoder.layer3,\n                    # 1/32\n                    timm_encoder.layer4,\n                ]\n            )\n            self.spacial_scalings = [2, 2, 2, 2, 2]\n        elif timm_model[\"model_name\"].startswith(\"seresnet\"):\n            self.spatial_donwsample_blocks = nn.ModuleList(\n                [\n                    # 1/2\n                    nn.Sequential(\n                        timm_encoder.stem_conv1,\n                        timm_encoder.stem_conv2,\n                    ),\n                    # 1/4\n                    nn.Sequential(\n                        timm_encoder.stem_conv3,\n                        timm_encoder.stages_0,\n                    ),\n                    # 1/8\n                    timm_encoder.stages_1,\n                    # 1/16\n                    timm_encoder.stages_2,\n                    # 1/32\n                    nn.Sequential(\n                        timm_encoder.stages_3,\n                        timm_encoder.final_conv,\n                    ),\n                ]\n            )\n            self.spacial_scalings = [2, 2, 2, 2, 2]\n        elif timm_model[\"model_name\"].startswith(\"tf_efficientnet_\"):\n            self.spatial_donwsample_blocks = nn.ModuleList(\n                [\n                    # 1/2\n                    nn.Sequential(\n                        timm_encoder.conv_stem,\n                        timm_encoder.bn1,\n                        timm_encoder.blocks[0],\n                    ),\n                    # 1/4\n                    timm_encoder.blocks[1],\n                    # 1/8\n                    timm_encoder.blocks[2],\n                    # 1/16\n                    nn.Sequential(\n                        timm_encoder.blocks[3],\n                        timm_encoder.blocks[4],\n                    ),\n                    nn.Sequential(\n                        # 1/32\n                        timm_encoder.blocks[5],\n                        timm_encoder.blocks[6],\n                    ),\n                ]\n            )\n            self.spacial_scalings = [2, 2, 2, 2, 2]\n        elif timm_model[\"model_name\"].startswith(\"tf_efficientnetv2_\"):\n            self.spatial_donwsample_blocks = nn.ModuleList(\n                [\n                    # 1/2\n                    nn.Sequential(\n                        timm_encoder.conv_stem,\n                        timm_encoder.bn1,\n                        timm_encoder.blocks[0],\n                    ),\n                    # 1/4\n                    timm_encoder.blocks[1],\n                    # 1/8\n                    timm_encoder.blocks[2],\n                    # 1/16\n                    nn.Sequential(\n                        timm_encoder.blocks[3],\n                        timm_encoder.blocks[4],\n                    ),\n                    timm_encoder.blocks[5],\n                ]\n            )\n            self.spacial_scalings = [2, 2, 2, 2, 2]\n        elif timm_model[\"model_name\"].startswith(\"mobilevitv2\"):\n            self.spatial_donwsample_blocks = nn.ModuleList(\n                [\n                    # 1/2\n                    nn.Sequential(\n                        timm_encoder.stem,\n                        timm_encoder.stages_0,\n                    ),\n                    # 1/4\n                    timm_encoder.stages_1,\n                    # 1/8\n                    timm_encoder.stages_2,\n                    # 1/16\n                    timm_encoder.stages_3,\n                    # 1/32\n                    timm_encoder.stages_4,\n                ]\n            )\n            self.spacial_scalings = [2, 2, 2, 2, 2]\n        elif timm_model[\"model_name\"].startswith(\"convnext\"):\n            self.spatial_donwsample_blocks = nn.ModuleList(\n                [\n                    # 1/4\n                    nn.Sequential(\n                        timm_encoder.stem_0,\n                        timm_encoder.stem_1,\n                        timm_encoder.stages_0,\n                    ),\n                    # 1/8\n                    timm_encoder.stages_1,\n                    # 1/16\n                    timm_encoder.stages_2,\n                    # 1/32\n                    timm_encoder.stages_3,\n                ]\n            )\n            self.spacial_scalings = [4, 2, 2, 2]\n        elif timm_model[\"model_name\"].startswith(\"eca_nfnet\"):\n            self.spatial_donwsample_blocks = nn.ModuleList(\n                [\n                    # 1/2\n                    nn.Sequential(\n                        timm_encoder.stem_conv1,\n                        timm_encoder.stem_act2,\n                        timm_encoder.stem_conv2,\n                        timm_encoder.stem_act3,\n                        timm_encoder.stem_conv3,\n                    ),\n                    # 1/4\n                    nn.Sequential(\n                        timm_encoder.stem_act4,\n                        timm_encoder.stem_conv4,\n                        timm_encoder.stages_0,\n                    ),\n                    # 1/8\n                    timm_encoder.stages_1,\n                    # 1/16\n                    timm_encoder.stages_2,\n                    # 1/32\n                    nn.Sequential(\n                        timm_encoder.stages_3,\n                        timm_encoder.final_conv,\n                    ),\n                ]\n            )\n            self.spacial_scalings = [2, 2, 2, 2, 2]\n        elif timm_model[\"model_name\"].startswith(\"maxvit\"):\n            self.spatial_donwsample_blocks = nn.ModuleList(\n                [\n                    # 1/2\n                    timm_encoder.stem,\n                    # 1/4\n                    timm_encoder.stages_0,\n                    # 1/8\n                    timm_encoder.stages_1,\n                    # 1/16\n                    timm_encoder.stages_2,\n                    # 1/32\n                    timm_encoder.stages_3,\n                ]\n            )\n            self.spacial_scalings = [2, 2, 2, 2, 2]\n        else:\n            raise ValueError(f\"Unsupported model: {timm_model['model_name']}\")\n\n        num_blocks = max(self.out_indices) + 1\n        if len(self.spatial_donwsample_blocks) > num_blocks:\n            self.spatial_donwsample_blocks = self.spatial_donwsample_blocks[:num_blocks]\n            self.spacial_scalings = self.spacial_scalings[:num_blocks]\n\n        if len(self.spatial_donwsample_blocks) != len(self.depth_scalings):\n            raise ValueError(\"Length of out_indices and depth_scalings must be the same\")\n\n        self.depth_donwsample_blocks = nn.ModuleList()\n        for out_channel, depth_scaling in zip(self.out_channels, self.depth_scalings, strict=False):\n            self.depth_donwsample_blocks.append(\n                DepthDownsampleBlock(\n                    in_channels=out_channel,\n                    out_channels=out_channel,\n                    depth_scaling=depth_scaling,\n                )\n            )\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        outputs = []\n\n        B, _, num_slices, h, w = x.shape\n        x = x.reshape(B * num_slices, 1, h, w)\n\n        # normalize\n        x = (x - 0.5) / 0.5\n        x = x.expand(-1, self.in_chans, -1, -1)\n\n        for i, (spatial_block, depth_block) in enumerate(zip(self.spatial_donwsample_blocks, self.depth_donwsample_blocks, strict=True)):\n            \n            x = spatial_block(x)\n            fc, fh, fw = x.shape[1:]\n\n            x1 = x.reshape(B, -1, fc, fh, fw).permute(0, 2, 1, 3, 4)\n            x1 = depth_block(x1)\n            x = x1.permute(0, 2, 1, 3, 4).reshape(-1, fc, fh, fw)\n            if i in self.out_indices:\n                outputs.append(x1)\n\n        return outputs\n\n\nclass CZII3DNet(nn.Module):\n    def __init__(\n        self,\n        timm_model: dict,\n        loss: dict,\n        depth_scalings: list[int] = [2, 2, 2, 2, 1],\n        decoder_dims: list[int] = [256, 128, 64, 32, 16],\n        num_classes: int = 6,\n        dropout: bool = False,\n        mode: str = \"multi_label\",\n        checkpoint: Optional[str] = None,\n    ):\n        super().__init__()\n\n        self.num_classes = num_classes\n        if mode == \"multi_label\":\n            self.output_channels = num_classes\n        elif mode == \"multi_class\":\n            self.output_channels = num_classes + 1\n        else:\n            raise ValueError(f\"{mode} is not supported.\")\n        self.mode = mode\n\n        self.encoder = My2DEncoder(timm_model, depth_scalings)\n        encoder_dims = self.encoder.out_channels\n        spetial_scalings = self.encoder.spacial_scalings\n        self.decoder = MyUnetDecoder3d(\n            in_channel=encoder_dims[-1],\n            skip_channels=encoder_dims[:-1][::-1] + [0],\n            out_channels=decoder_dims,\n            spetial_scalings=spetial_scalings[::-1],\n            depth_scalings=depth_scalings[::-1],\n        )\n        self.head = nn.Conv3d(decoder_dims[-1], self.output_channels, kernel_size=1)\n\n        if checkpoint is not None:\n            ckpt = torch.load(checkpoint)\n            if \"ema_model_state_dict\" in ckpt:\n                missing_keys, unexpected_keys = self.load_state_dict(ckpt[\"ema_model_state_dict\"], strict=False)\n            else:\n                missing_keys, unexpected_keys = self.load_state_dict(ckpt[\"state_dict\"], strict=False)\n            for k in missing_keys:\n                if k not in [\"target_loss.bce_loss.pos_weight\"]:\n                    raise ValueError(f\"Missing key: {k}\")\n            if len(unexpected_keys) > 0:\n                raise ValueError(f\"Unexpected keys: {unexpected_keys}\")\n\n    def forward(self, slices: torch.Tensor) -> torch.Tensor:\n        outputs = dict()\n\n        feats = self.encoder(slices)\n        last_feats, _ = self.decoder(feature=feats[-1], skip=feats[:-1][::-1] + [None])\n        logits = self.head(last_feats)\n        outputs[\"logits\"] = logits\n\n        if self.mode == \"multi_label\":\n            outputs[\"scores\"] = logits.detach().sigmoid()\n        else:\n            outputs[\"scores\"] = logits.detach().softmax(dim=1)\n\n        return outputs","metadata":{"execution":{"iopub.status.busy":"2025-02-06T03:56:33.441786Z","iopub.execute_input":"2025-02-06T03:56:33.442002Z","iopub.status.idle":"2025-02-06T03:56:33.470483Z","shell.execute_reply.started":"2025-02-06T03:56:33.441971Z","shell.execute_reply":"2025-02-06T03:56:33.469677Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3D Semantic Segmentation Detector","metadata":{}},{"cell_type":"code","source":"class Detector3D:\n    def __init__(self, config_path: str, ckpt_path: str, device: torch.device, is_half: bool = True, is_hflip: bool = True, is_vflip: bool = False, is_dflip: bool = False):\n        self.device = device\n        self.is_half = is_half\n        self.is_hflip = is_hflip\n        self.is_vflip = is_vflip\n        self.is_dflip = is_dflip\n        self.num_inference = 1 + int(is_hflip) + int(is_vflip) + int(is_dflip)\n\n        self.name = Path(config_path).stem\n\n        config = get_config(config_path, [])\n        config['model']['timm_model']['pretrained'] = False\n        config['model']['checkpoint'] = None\n        self.patch_extractor = VolumePatchExtractor(**config[\"dataset\"][\"validation\"][\"volume_patch_extractor_params\"])\n\n        model_name = config['model'].pop('name')\n        model_class = globals()[model_name]\n        model = model_class(**config['model'])\n        self.mode = model.mode\n        self.num_classes = model.num_classes\n\n        checkpoint = torch.load(ckpt_path, map_location=\"cpu\")\n        if \"ema_model_state_dict\" in checkpoint:\n            missing_keys, unexpected_keys = model.load_state_dict(checkpoint[\"ema_model_state_dict\"], strict=False)\n        else:\n            missing_keys, unexpected_keys = model.load_state_dict(checkpoint[\"state_dict\"], strict=False)\n        for k in missing_keys:\n            if k not in [\"loss.bce_loss.pos_weight\"]:\n                raise ValueError(f\"Missing key: {k}\")\n        for k in unexpected_keys:\n            if k not in [\"loss.bce_loss.pos_weight\"]:\n                raise ValueError(f\"Unexpected key: {k}\")\n\n        self.model = model\n        self.model.eval()\n        if self.is_half:\n            self.model = self.model.half()\n        self.model.to(device)\n\n    def _preprocess(self, volume_norm: np.ndarray) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:\n        patched_images, patched_weights, patched_coords = self.patch_extractor.extract(volume_norm)\n        patched_images = np.stack(patched_images).astype(np.float32)\n        patched_images = patched_images[:, None, ...]\n        patched_weights = np.stack(patched_weights).astype(np.float32)\n\n        # to torch\n        patched_images = torch.from_numpy(patched_images).to(torch.float32)\n        patched_weights = torch.from_numpy(patched_weights).to(torch.float32)\n        patched_coords = torch.tensor(patched_coords)\n\n        return patched_images, patched_weights, patched_coords\n\n    @ torch.no_grad()\n    def _forward(self, image: torch.Tensor) -> torch.Tensor:\n        score = self.model(image)[\"scores\"]\n        if self.is_half and torch.isnan(score).any():\n            self.model = self.model.float()\n            score = self.model(image.float())[\"scores\"]\n            self.model = self.model.half()\n        return score\n\n\n    @ torch.no_grad()\n    def detect(self, volume_norm: np.ndarray) -> np.ndarray:\n        patched_images, patched_weights, patched_coords = self._preprocess(volume_norm)\n        if self.is_half:\n            patched_images = patched_images.half()\n        patched_images = patched_images.to(self.device)\n        patched_weights = patched_weights.to(self.device)\n\n        d, h, w = volume_norm.shape\n        agg_score = torch.zeros((self.num_classes, d, h, w), device=self.device)\n        agg_weight = torch.zeros((self.num_classes, d, h, w), device=self.device)\n\n        for patch_image, patch_weight, patch_coord in zip(patched_images, patched_weights, patched_coords, strict=False):\n            patch_image = patch_image[None, ...]\n            patch_weight = patch_weight\n\n            # forward\n            patch_score = self._forward(patch_image)[0]\n\n            # TODO: Batch方向に統合したいね\n            if self.is_hflip:\n                hflip_patch_image = torch.flip(patch_image, dims=[4])\n                hflip_patch_score = self._forward(hflip_patch_image)\n                hflip_patch_score = torch.flip(hflip_patch_score, dims=[4])\n                patch_score += hflip_patch_score[0]\n\n            if self.is_vflip:\n                vflip_patch_image = torch.flip(patch_image, dims=[3])\n                vflip_patch_score = self._forward(vflip_patch_image)\n                vflip_patch_score = torch.flip(vflip_patch_score, dims=[3])\n                patch_score += vflip_patch_score[0]\n\n            if self.is_dflip:\n                dflip_patch_image = torch.flip(patch_image, dims=[2])\n                dflip_patch_score = self._forward(dflip_patch_image)\n                dflip_patch_score = torch.flip(dflip_patch_score, dims=[2])\n                patch_score += dflip_patch_score[0]\n\n            patch_score = patch_score / self.num_inference\n\n            z_min, y_min, x_min, z_max, y_max, x_max = patch_coord\n            if self.mode == \"multi_label\":\n                agg_score[:, z_min:z_max, y_min:y_max, x_min:x_max] += patch_score * patch_weight[None, ...]\n            else:\n                agg_score[:, z_min:z_max, y_min:y_max, x_min:x_max] += patch_score[1:] * patch_weight[None, ...]\n            agg_weight[:, z_min:z_max, y_min:y_max, x_min:x_max] += patch_weight\n\n        agg_score = agg_score / (agg_weight + 0.0001)\n        probability = agg_score.cpu().numpy()\n\n        return probability","metadata":{"execution":{"iopub.status.busy":"2025-02-06T03:56:33.471358Z","iopub.execute_input":"2025-02-06T03:56:33.471639Z","iopub.status.idle":"2025-02-06T03:56:33.486110Z","shell.execute_reply.started":"2025-02-06T03:56:33.471605Z","shell.execute_reply":"2025-02-06T03:56:33.485495Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Post Processor","metadata":{}},{"cell_type":"code","source":"class ClusteringPostProcessor:\n    \"\"\"\n    3D連結成分解析 (cc3d) により、しきい値以上の領域をクラスタリングし、\n    cc3d.statistics(..., overlap_array=...) の centroid 情報を用いて\n    加重重心 (x, y, z) を推定するポストプロセッサ。\n\n    出力は pandas.DataFrame:\n      列 = [\"particle_type\", \"x\", \"y\", \"z\"].\n    \"\"\"\n\n    def __init__(\n        self,\n        class_infos: list,\n        thresholds: float | list[float] = [0.5, 1.0, 0.5, 0.5, 0.5, 0.5],\n        blob_thresholds: int | list[int] = [0, 0, 0, 0, 0, 0],\n        connectivity: int = 26,\n        voxel_size: float = 10.012,\n        score_type: str = \"max\",\n    ) -> None:\n        \"\"\"\n        Parameters\n        ----------\n        class_infos : list\n            クラス情報\n        thresholds : float | list[float]\n            スコアしきい値 (0~1)。この値以上のボクセルのみ連結成分解析の対象にする。\n        blob_thresholds : float | list[float]\n            面積しきい値。この値未満のラベルは無視する。\n        connectivity : int\n            3D連結度 (6, 18, 26など)。cc3d.connected_components() の引数に渡す。\n        voxel_size : float\n            ボクセルサイズ (オングストローム)\n        score_type : str\n            領域スコアの計算方法: \"mean\", \"sum\", \"max\" など。\n        \"\"\"\n        self.class_infos = copy.deepcopy(class_infos)\n        self.thresholds = thresholds if isinstance(thresholds, list) else [thresholds] * len(class_infos)\n        self.blob_thresholds = blob_thresholds if isinstance(blob_thresholds, list) else [blob_thresholds] * len(class_infos)\n        if len(self.thresholds) != len(class_infos) or len(self.blob_thresholds) != len(class_infos):\n            raise ValueError(\"len(self.thresholds) != len(class_infos) or len(self.blob_thresholds) != len(class_infos)\")\n        self.connectivity = connectivity\n        self.voxel_size = voxel_size\n        self.score_type = score_type\n\n    def __call__(self, pred_heatmaps: np.ndarray) -> pd.DataFrame:\n        \"\"\"\n        shape = (C, D, H, W) のセグメンテーション予測 (0~1) を受け取り、\n        閾値以上の領域をラベル付けして、その重心・領域スコア・ボクセル数を計算。\n\n        Returns\n        -------\n        pd.DataFrame:\n            列 = [\"particle_type\", \"x\", \"y\", \"z\", \"score\", \"voxel_count\"] \n            - (x,y,z) は重心 [オングストローム]、score は領域スコア (計算方法は score_type 依存)、\n              voxel_count は領域ボクセル数\n        \"\"\"\n        df_list = []\n\n        for class_info, threshold, blob_threshold in zip(self.class_infos, self.thresholds, self.blob_thresholds):\n            c_idx = class_info[\"id\"] - 1\n            class_name = class_info[\"name\"]\n            volume_c = pred_heatmaps[c_idx]\n\n            mask = (volume_c >= threshold).astype(np.uint8)\n            if not np.any(mask):\n                # 1つも領域が無いならスキップ\n                continue\n\n            labels_out = cc3d.connected_components(mask, connectivity=self.connectivity, return_N=False)\n            stats = cc3d.statistics(labels_out)\n\n            centroids_zyx_all = stats[\"centroids\"][1:]\n            voxel_counts_all = stats[\"voxel_counts\"][1:]\n            valid_mask = (voxel_counts_all >= blob_threshold)\n            centroids_zyx = centroids_zyx_all[valid_mask]\n            voxel_counts = voxel_counts_all[valid_mask]\n\n            # 領域のスコアを手動で計算する\n            region_scores = []\n            # valid_label_ids = np.where(valid_mask)[0] + 1  # label番号(1～N)\n            valid_label_ids = np.flatnonzero(valid_mask) + 1\n            \n            # labeled_comprehensionは[func(input[labels == i]) for i in index]と同じで高速（っぽい、すごいっすね～～～～）\n            if self.score_type == \"mean\":\n                region_scores = labeled_comprehension(volume_c, labels_out, valid_label_ids, np.mean, float, 0)\n            elif self.score_type == \"sum\":\n                region_scores = labeled_comprehension(volume_c, labels_out, valid_label_ids, np.sum, float, 0)\n            elif self.score_type == \"max\":\n                region_scores = labeled_comprehension(volume_c, labels_out, valid_label_ids, np.max, float, 0)\n            else:\n                raise ValueError(f\"Invalid score_type: {self.score_type}\")\n\n            coords_xyz = centroids_zyx[:, ::-1] * self.voxel_size\n\n            df_c = pd.DataFrame(coords_xyz, columns=[\"x\", \"y\", \"z\"])\n            df_c[\"particle_type\"] = class_name\n            df_c[\"score\"] = region_scores\n            df_c[\"voxel_count\"] = voxel_counts\n\n            df_list.append(df_c)\n\n        if len(df_list) == 0:\n            return pd.DataFrame(columns=[\"particle_type\", \"x\", \"y\", \"z\", \"score\", \"voxel_count\"])\n        else:\n            df = pd.concat(df_list, ignore_index=True)\n            df = df[[\"particle_type\", \"x\", \"y\", \"z\", \"score\", \"voxel_count\"]]\n            return df\n\n\nclass WPFPostProcessor:\n    def __init__(\n        self,\n        class_infos: list,\n        thresholds: float = [0.5, 1.0, 0.5, 0.5, 0.5, 0.5],\n    ):\n        self.class_infos = copy.deepcopy(class_infos)\n        self.thresholds = thresholds\n        if len(self.thresholds) != len(class_infos):\n            raise ValueError(\"len(self.thresholds) != len(class_infos)\")\n\n    def _weighted_points_fusion(\n        self,\n        centroid_df :pd.DataFrame,\n        distance_threshold: float,\n    ):\n        \"\"\"\n        3Dの点 (x, y, z, score, model_id) をクラスタリングし、加重平均で座標を融合する WPF。\n\n        Args:\n            centroid_df (pd.DataFrame):\n                ClusteringPostProcessorの出力。列 = [\"experiment\", \"model_id\", \"particle_type\", \"x\", \"y\", \"z\", \"score\", \"voxel_count\"]\n            distance_threshold (float):\n                3D距離の閾値。この距離以内なら同一クラスターとみなす。\n\n        Returns:\n            pd.DataFrame:\n                列 = [\"experiment\", \"particle_type\", \"x\", \"y\", \"z\", \"score\"]\n        \"\"\"\n\n        # 入力をスコア降順にソート\n        points = centroid_df[[\"x\", \"y\", \"z\", \"score\", \"model_id\"]].values\n        points = points[np.argsort(points[:, 3])[::-1]]\n\n        # モデル数\n        num_models = len(centroid_df[\"model_id\"].unique())\n\n        # クラスター = [cx, cy, cz, score_sum, set_of_model_ids]\n        fused_clusters = []\n\n        for i in range(len(points)):\n            x, y, z, s, mid = points[i]\n\n            if len(fused_clusters) == 0:\n                fused_clusters.append([x, y, z, s, [mid]])\n                continue\n\n            matched_idx = None\n            for idx, (cx, cy, cz, c_score_sum, c_model_ids) in enumerate(fused_clusters):\n                dist = np.sqrt((cx - x)**2 + (cy - y)**2 + (cz - z)**2)\n                if dist <= distance_threshold:\n                    matched_idx = idx\n                    break\n\n            if matched_idx is not None:\n                # 既存クラスタにマージ → 加重平均で座標更新\n                cx, cy, cz, c_score_sum, c_model_ids = fused_clusters[matched_idx]\n\n                new_score_sum = c_score_sum + s\n                new_cx = (cx * c_score_sum + x * s) / new_score_sum\n                new_cy = (cy * c_score_sum + y * s) / new_score_sum\n                new_cz = (cz * c_score_sum + z * s) / new_score_sum\n\n                c_model_ids.append(mid)\n                fused_clusters[matched_idx] = [new_cx, new_cy, new_cz, new_score_sum, c_model_ids]\n\n            else:\n                # 新規クラスタを追加\n                fused_clusters.append([x, y, z, s, [mid]])\n\n        # クラスタ内のモデル数でスコア調整\n        results = []\n        for (cx, cy, cz, c_score_sum, c_model_ids) in fused_clusters:\n            c_num_points = len(c_model_ids)\n            c_num_models = len(set(c_model_ids))\n            avg_score = c_score_sum / c_num_points\n            penalty = (min(num_models, c_num_models) / num_models)\n            final_score = avg_score * penalty\n            results.append([cx, cy, cz, final_score])\n\n        final = np.array(results)\n        result_df = pd.DataFrame(final, columns=[\"x\", \"y\", \"z\", \"score\"])\n        return result_df\n\n    def __call__(self, centroid_df):\n        if len(centroid_df[\"experiment\"].unique()) > 1:\n            raise ValueError(\"Multiple experiments is not supported.\")\n        experiment = centroid_df[\"experiment\"].unique()[0]\n\n        wpf_df_list = []\n        for class_info, threshold in zip(self.class_infos,  self.thresholds):\n            class_name = class_info[\"name\"]\n            if class_name == \"beta-amylase\":\n                continue\n\n            distance_thresholds = class_info[\"radius\"] * 0.5\n            centroid_class_df = centroid_df[centroid_df[\"particle_type\"] == class_name]\n            if len(centroid_class_df) == 0:\n                continue\n\n            wpf_df = self._weighted_points_fusion(centroid_class_df, distance_thresholds)\n            wpf_df = wpf_df[wpf_df[\"score\"] >= threshold]\n            if len(wpf_df) == 0:\n                continue\n\n            wpf_df.insert(0, \"experiment\", experiment)\n            wpf_df.insert(1, \"particle_type\", class_name)\n            wpf_df_list.append(wpf_df)\n        \n        if len(wpf_df_list) > 0:\n            wpf_df = pd.concat(wpf_df_list, ignore_index=True).reset_index(drop=True)\n            return wpf_df\n        else:\n            return pd.DataFrame(columns=[\"experiment\", \"particle_type\", \"x\", \"y\", \"z\", \"score\"])","metadata":{"execution":{"iopub.status.busy":"2025-02-06T03:56:33.486978Z","iopub.execute_input":"2025-02-06T03:56:33.487285Z","iopub.status.idle":"2025-02-06T03:56:33.506007Z","shell.execute_reply.started":"2025-02-06T03:56:33.487263Z","shell.execute_reply":"2025-02-06T03:56:33.505355Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2D Remover Model","metadata":{}},{"cell_type":"code","source":"class GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6, p_trainable=False, flatten=True):\n        super(GeM, self).__init__()\n        if p_trainable:\n            self.p = nn.Parameter(torch.ones(1) * p)\n        else:\n            self.p = p\n        self.eps = eps\n        self.flatten = flatten\n\n    def _gem(self, x, p=3, eps=1e-6):\n        return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1.0 / p)\n\n    def forward(self, x):\n        ret = self._gem(x, p=self.p, eps=self.eps)\n        if self.flatten:\n            return ret[:, :, 0, 0]\n        else:\n            return ret\n\n    def __repr__(self):\n        if not isinstance(self.p, int):\n            return (self.__class__.__name__ + f\"(p={self.p.data.tolist()[0]:.4f},eps={self.eps})\")\n        else:\n            return (self.__class__.__name__ + f\"(p={self.p:.4f},eps={self.eps})\")\n\n\nclass CZIIRemoverNet(nn.Module):\n    def __init__(\n        self,\n        timm_model: dict,\n        loss: dict,\n    ):\n        super().__init__()\n        self._in_channels = timm_model.get(\"in_chans\", 3)\n\n        self.backbone = timm.create_model(**timm_model)\n\n        if \"efficientnet\" in timm_model[\"model_name\"]:\n            backbone_out_channels = self.backbone.num_features\n            self.backbone.global_pool = nn.Identity()\n            self.backbone.classifier = nn.Identity()\n        elif timm_model[\"model_name\"].startswith(\"convnext_\"):\n            backbone_out_channels = self.backbone.head.in_features\n            self.backbone.head = nn.Identity()\n        elif timm_model[\"model_name\"].startswith(\"resnet\"):\n            backbone_out_channels = self.backbone.fc.in_features\n            self.backbone.global_pool = nn.Identity()\n            self.backbone.fc = nn.Identity()\n        elif timm_model[\"model_name\"].startswith(\"tinynet\"):\n            backbone_out_channels = self.backbone.num_features\n            self.backbone.global_pool = nn.Identity()\n            self.backbone.classifier = nn.Identity()\n        elif timm_model[\"model_name\"].startswith(\"mobilevit\"):\n            backbone_out_channels = self.backbone.head.in_features\n            self.backbone.head = nn.Identity()\n        elif hasattr(timm_model[\"model_name\"], \"fc\"):\n            backbone_out_channels = self.backbone.fc.in_features\n        else:\n            raise ValueError(f'{timm_model[\"model_name\"]} is not supported.')\n\n        self.global_pool = GeM(flatten=True)\n        self.dropouts = nn.ModuleList([nn.Dropout(p) for p in np.linspace(0.1, 0.5, 5)])\n        self.fc = nn.Linear(backbone_out_channels, 1)\n\n    def forward(\n        self,\n        images: torch.Tensor,\n    ) -> dict:\n        outputs = dict()\n        images = images - 0.5 / 0.5\n        feats = self.backbone(images)\n        feats = self.global_pool(feats)\n        logits = self.fc(feats)\n\n        outputs[\"logits\"] = logits\n        outputs[\"scores\"] = torch.sigmoid(logits)\n\n        return outputs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-06T03:56:33.506729Z","iopub.execute_input":"2025-02-06T03:56:33.506914Z","iopub.status.idle":"2025-02-06T03:56:33.518417Z","shell.execute_reply.started":"2025-02-06T03:56:33.506897Z","shell.execute_reply":"2025-02-06T03:56:33.517748Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Remover","metadata":{}},{"cell_type":"code","source":"class Remover:\n    \"\"\"誤検出を消してくれる心強い相棒\n    \"\"\"\n    def __init__(self, particle_type: str, path_list: list[dict], image_size: int, volume_size: int, threshold: float, device: torch.device):\n        self.particle_type = particle_type\n        self.image_size = image_size\n        self.volume_size = volume_size\n        self.threshold = threshold\n        self.device = device\n\n        self.models = []\n        for path_dict in path_list:\n            config_path = path_dict[\"config\"]\n\n            for ckpt_path in path_dict[\"ckpts\"]:\n                config = get_config(config_path, [])\n                config['model']['timm_model']['pretrained'] = False\n        \n                model_name = config['model'].pop('name')\n                model_class = globals()[model_name]\n                model = model_class(**config['model'])\n        \n                checkpoint = torch.load(ckpt_path, map_location=\"cpu\")\n                if \"ema_model_state_dict\" in checkpoint:\n                    missing_keys, unexpected_keys = model.load_state_dict(checkpoint[\"ema_model_state_dict\"], strict=False)\n                else:\n                    missing_keys, unexpected_keys = model.load_state_dict(checkpoint[\"state_dict\"], strict=False)\n                for k in missing_keys:\n                    if k not in [\"loss.bce_loss.pos_weight\"]:\n                        raise ValueError(f\"Missing key: {k}\")\n                for k in unexpected_keys:\n                    if k not in [\"loss.bce_loss.pos_weight\"]:\n                        raise ValueError(f\"Unexpected key: {k}\")\n        \n                model.eval()\n                model.to(device)\n                self.models.append(model)\n\n    def _preprocess(self, volume: np.ndarray, xyz: np.ndarray) -> torch.Tensor:\n        images = []\n        half_volume_size = self.volume_size // 2\n        xyz_int = np.round(xyz / 10).astype(int)\n        for x, y, z in xyz_int:\n            x_start = max(0, x - half_volume_size)\n            x_end = min(x + half_volume_size, volume.shape[2])\n    \n            y_start = max(0, y - half_volume_size)\n            y_end = min(y + half_volume_size, volume.shape[1])\n    \n            z_start = max(0, z - half_volume_size)\n            z_end = min(z + half_volume_size, volume.shape[0])\n\n            image = volume[z_start:z_end, y_start:y_end, x_start:x_end]\n    \n            pad_x_before = max(0, -(x - half_volume_size))\n            pad_x_after = max(0, (x + half_volume_size) - volume.shape[2])\n    \n            pad_y_before = max(0, -(y - half_volume_size))\n            pad_y_after = max(0, (y + half_volume_size) - volume.shape[1])\n    \n            pad_z_before = max(0, -(z - half_volume_size))\n            pad_z_after = max(0, (z + half_volume_size) - volume.shape[0])\n    \n            if pad_x_before > 0 or pad_x_after > 0 or pad_y_before > 0 or pad_y_after > 0 or pad_z_before > 0 or pad_z_after > 0:\n                image = np.pad(image, ((pad_z_before, pad_z_after), (pad_y_before, pad_y_after), (pad_x_before, pad_x_after)), mode=\"constant\")\n    \n            # [volume_size, volume_size, volume_size]のvolumeをdepth方向に圧縮する (Minislab的な処理)\n            image = image.mean(axis=0)\n            image = cv2.resize(image, (self.image_size, self.image_size), interpolation=cv2.INTER_LINEAR)\n            images.append(image[None, ...])\n        images = np.stack(images)\n        images = torch.from_numpy(images)\n        return images\n\n    @ torch.no_grad()\n    def _forward(self, image: torch.Tensor) -> torch.Tensor:\n        score = self.model(image)[\"scores\"]\n        return score\n\n    @ torch.no_grad()\n    def __call__(self, volume_norm: np.ndarray, wpf_df: pd.DataFrame, chunk_size: int = 256) -> pd.DataFrame:\n        partical_df = wpf_df[wpf_df[\"particle_type\"] == self.particle_type].copy().reset_index(drop=True)\n        images = self._preprocess(volume_norm, partical_df[[\"x\", \"y\", \"z\"]].values)\n        images = images.to(self.device)\n\n        scores = []\n        num_images = len(images)\n\n        # chunk_size個の画像ごとにforwardする\n        for start_index in range(0, num_images, chunk_size):\n            image_chunk = images[start_index:start_index+chunk_size]\n            current_chunk_size = image_chunk.shape[0]\n            chunk_scores = np.zeros(current_chunk_size)\n            for model in self.models:\n                chunk_scores += model(image_chunk)[\"scores\"].detach().cpu().numpy().flatten()\n            chunk_scores /= len(self.models)\n            scores.append(chunk_scores)\n        scores = np.concatenate(scores)\n\n        # 閾値以上の予測結果だけ残す\n        mask = scores >= self.threshold\n        masked_partical_df = partical_df[mask]\n\n        new_wpf_df = wpf_df[wpf_df[\"particle_type\"] != self.particle_type]\n        new_wpf_df = pd.concat([new_wpf_df, masked_partical_df], ignore_index=True)\n        new_wpf_df.sort_values([\"experiment\", \"particle_type\", \"score\"], ascending=[True, True, False], inplace=True)\n        return new_wpf_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-06T03:56:33.519262Z","iopub.execute_input":"2025-02-06T03:56:33.519486Z","iopub.status.idle":"2025-02-06T03:56:33.535265Z","shell.execute_reply.started":"2025-02-06T03:56:33.519467Z","shell.execute_reply":"2025-02-06T03:56:33.534584Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"def inference(zarr_paths: list[Path], model_3d_infos: list, centroid_post_process_settings: dict, wpf_post_process_settings: dict, remover_infos: list, device: torch.device) -> pd.DataFrame:\n    detector_3d_list = [Detector3D(**model_info, device=device) for model_info in model_3d_infos]\n    centroid_post_processor = ClusteringPostProcessor(**centroid_post_process_settings)\n    wpf_post_processor = WPFPostProcessor(**wpf_post_process_settings)\n    removers = [Remover(**remover_info, device=device) for remover_info in remover_infos]\n\n    result_df_list = []\n    for zarr_path in zarr_paths:\n        run_name = zarr_path.parts[-3]\n        volume = load_tomogram(zarr_path)\n\n        # 3Dセマセグと粒子座標推定\n        centroid_df_list = []\n        for detector in detector_3d_list:\n            start_time = time.time()\n            probability = detector.detect(volume)\n            centroid_df = centroid_post_processor(probability)\n            centroid_df.insert(0, \"experiment\", run_name)\n            centroid_df.insert(1, \"model_id\", detector.name)\n            centroid_df_list.append(centroid_df)\n            end_time = time.time()\n            print(f\"{run_name=} {device=} detector={detector.name} time={end_time - start_time} sec\")  \n\n        # WPFで粒子座標の統合\n        centroid_df = pd.concat(centroid_df_list)\n        wpf_df = wpf_post_processor(centroid_df)\n\n        # Removerで低信頼度粒子の削除\n        start_time = time.time()\n        for remover in removers:\n            wpf_df = remover(volume, wpf_df)\n        end_time = time.time()\n        print(f\"{run_name=} {device=} remover time={end_time - start_time} sec\")  \n\n        result_df = wpf_df[[\"experiment\", \"particle_type\", \"x\", \"y\", \"z\"]]\n        result_df_list.append(result_df)\n\n    result_df = pd.concat(result_df_list)\n    result_df = result_df.reset_index(drop=True)\n    return result_df","metadata":{"execution":{"iopub.status.busy":"2025-02-06T03:56:33.535990Z","iopub.execute_input":"2025-02-06T03:56:33.536222Z","iopub.status.idle":"2025-02-06T03:56:33.550671Z","shell.execute_reply.started":"2025-02-06T03:56:33.536187Z","shell.execute_reply":"2025-02-06T03:56:33.549817Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"start_time = time.time()\nif CONFIG.MAX_WORKERS == 1:\n    submit_df = inference(CONFIG.ZARR_FILES,\n                          CONFIG.MODEL_3D_INFOS,\n                          CONFIG.CENTROID_POST_PROCESS_SETTINGS,\n                          CONFIG.WPF_POST_PROCESS_SETTINGS,\n                          CONFIG.REMOVER_INFOS,\n                          torch.device(0))\nelif CONFIG.MAX_WORKERS == 2:\n    center_idx = len(CONFIG.ZARR_FILES) // 2\n    zarr_paths1 = CONFIG.ZARR_FILES[:center_idx]\n    zarr_paths2 = CONFIG.ZARR_FILES[center_idx:]\n    with ProcessPoolExecutor(max_workers=CONFIG.MAX_WORKERS) as executor:\n        submit_df_list = list(\n            executor.map(\n                inference,\n                (zarr_paths1, zarr_paths2),\n                (CONFIG.MODEL_3D_INFOS, CONFIG.MODEL_3D_INFOS),\n                (CONFIG.CENTROID_POST_PROCESS_SETTINGS, CONFIG.CENTROID_POST_PROCESS_SETTINGS),\n                (CONFIG.WPF_POST_PROCESS_SETTINGS, CONFIG.WPF_POST_PROCESS_SETTINGS),\n                (CONFIG.REMOVER_INFOS, CONFIG.REMOVER_INFOS),\n                (torch.device(0), torch.device(1)),\n            )\n        )\n        submit_df = pd.concat(submit_df_list)\nelse:\n    raise ValueError(f\"Invalid {CONFIG.MAX_WORKERS=}\")\nend_time = time.time()\n\nelapsed_time = end_time - start_time\nprint(f\"Elapsed time: {elapsed_time} sec\")\n\nprocessed_count = math.ceil(len(CONFIG.ZARR_FILES) / CONFIG.MAX_WORKERS) * CONFIG.MAX_WORKERS\nprint(f\"500 files (h): {((elapsed_time / processed_count) * 500 / 3600):.2f}\")","metadata":{"execution":{"iopub.status.busy":"2025-02-06T03:56:33.551493Z","iopub.execute_input":"2025-02-06T03:56:33.551722Z","iopub.status.idle":"2025-02-06T04:01:20.256492Z","shell.execute_reply.started":"2025-02-06T03:56:33.551702Z","shell.execute_reply":"2025-02-06T04:01:20.255420Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Submit","metadata":{}},{"cell_type":"code","source":"submit_df = submit_df.reset_index(drop=True)\nsubmit_df.insert(0, 'id', range(len(submit_df)))\nsubmit_df.to_csv(\"submission.csv\", index=False)\nsubmit_df.head()","metadata":{"execution":{"iopub.status.busy":"2025-02-06T04:01:20.257655Z","iopub.execute_input":"2025-02-06T04:01:20.257926Z","iopub.status.idle":"2025-02-06T04:01:20.286201Z","shell.execute_reply.started":"2025-02-06T04:01:20.257902Z","shell.execute_reply":"2025-02-06T04:01:20.285576Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if submit_df.columns.to_list() != [\"id\", \"experiment\", \"particle_type\", \"x\", \"y\", \"z\"]:\n    raise ValueError(\"submit csv error.\")","metadata":{"execution":{"iopub.status.busy":"2025-02-06T04:01:20.287006Z","iopub.execute_input":"2025-02-06T04:01:20.287331Z","iopub.status.idle":"2025-02-06T04:01:20.291173Z","shell.execute_reply.started":"2025-02-06T04:01:20.287298Z","shell.execute_reply":"2025-02-06T04:01:20.290231Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Debug","metadata":{}},{"cell_type":"code","source":"if CONFIG.ENV == \"Notebook\":\n    from scipy.spatial import KDTree\n    from scipy.optimize import linear_sum_assignment\n    from tabulate import tabulate\n\n\n    class ParticipantVisibleError(Exception):\n        pass\n    \n    \n    def compute_metrics(\n        reference_points: np.ndarray,\n        reference_radius: float,\n        candidate_points: np.ndarray,\n    ) -> tuple:\n        # reference: https://www.kaggle.com/code/metric/czi-cryoet-84969\n        num_reference_particles = len(reference_points)\n        num_candidate_particles = len(candidate_points)\n    \n        if len(reference_points) == 0:\n            return 0, num_candidate_particles, 0\n    \n        if len(candidate_points) == 0:\n            return 0, 0, num_reference_particles\n    \n        ref_tree = KDTree(reference_points)\n        candidate_tree = KDTree(candidate_points)\n        raw_matches = candidate_tree.query_ball_tree(ref_tree, r=reference_radius)\n        matches_within_threshold = []\n        for match in raw_matches:\n            matches_within_threshold.extend(match)\n        # Prevent submitting multiple matches per particle.\n        # This won't be be strictly correct in the (extremely rare) case where true particles\n        # are very close to each other.\n        matches_within_threshold = set(matches_within_threshold)\n        tp = int(len(matches_within_threshold))\n        fp = int(num_candidate_particles - tp)\n        fn = int(num_reference_particles - tp)\n        return tp, fp, fn\n    \n    \n    def score(\n        solution: pd.DataFrame,\n        submission: pd.DataFrame,\n        row_id_column_name: str,\n        distance_multiplier: float,\n        beta: int,\n    ) -> float:\n        \"\"\"\n        F_beta (reference: https://www.kaggle.com/code/metric/czi-cryoet-84969)\n          - a true positive occurs when\n             - (a) the predicted location is within a threshold of the particle radius, and\n             - (b) the correct `particle_type` is specified\n          - raw results (TP, FP, FN) are aggregated across all experiments for each particle type\n          - f_beta is calculated for each particle type\n          - individual f_beta scores are weighted by particle type for final score\n        \"\"\"\n    \n        particle_radius = {\n            \"apo-ferritin\": 60,\n            \"beta-amylase\": 65,\n            \"beta-galactosidase\": 90,\n            \"ribosome\": 150,\n            \"thyroglobulin\": 130,\n            \"virus-like-particle\": 135,\n        }\n    \n        weights = {\n            \"apo-ferritin\": 1,\n            \"beta-amylase\": 0,\n            \"beta-galactosidase\": 2,\n            \"ribosome\": 1,\n            \"thyroglobulin\": 2,\n            \"virus-like-particle\": 1,\n        }\n    \n        particle_radius = {k: v * distance_multiplier for k, v in particle_radius.items()}\n    \n        # Filter submission to only contain experiments found in the solution split\n        # original\n        # split_experiments = set(solution[\"experiment\"].unique())\n        # submission = submission.loc[submission[\"experiment\"].isin(split_experiments)]\n        # modified (submission based)\n        split_experiments = set(submission[\"experiment\"].unique())\n        solution = solution.loc[solution[\"experiment\"].isin(split_experiments)]\n    \n        # Only allow known particle types\n        if not set(submission[\"particle_type\"].unique()).issubset(set(weights.keys())):\n            raise ParticipantVisibleError(\"Unrecognized `particle_type`.\")\n    \n        assert solution.duplicated(subset=[\"experiment\", \"x\", \"y\", \"z\"]).sum() == 0\n        assert particle_radius.keys() == weights.keys()\n    \n        results = {}\n        for particle_type in solution[\"particle_type\"].unique():\n            results[particle_type] = {\n                \"total_tp\": 0,\n                \"total_fp\": 0,\n                \"total_fn\": 0,\n            }\n    \n        for experiment in split_experiments:\n            for particle_type in solution[\"particle_type\"].unique():\n                reference_radius = particle_radius[particle_type]\n                select = (solution[\"experiment\"] == experiment) & (solution[\"particle_type\"] == particle_type)\n                reference_points = solution.loc[select, [\"x\", \"y\", \"z\"]].values\n    \n                select = (submission[\"experiment\"] == experiment) & (submission[\"particle_type\"] == particle_type)\n                candidate_points = submission.loc[select, [\"x\", \"y\", \"z\"]].values\n    \n                if len(reference_points) == 0:\n                    reference_points = np.array([])\n                    reference_radius = 1\n    \n                if len(candidate_points) == 0:\n                    candidate_points = np.array([])\n    \n                tp, fp, fn = compute_metrics(reference_points, reference_radius, candidate_points)\n    \n                results[particle_type][\"total_tp\"] += tp\n                results[particle_type][\"total_fp\"] += fp\n                results[particle_type][\"total_fn\"] += fn\n    \n        aggregate_fbeta = 0.0\n        for particle_type, totals in results.items():\n            tp = totals[\"total_tp\"]\n            fp = totals[\"total_fp\"]\n            fn = totals[\"total_fn\"]\n    \n            precision = tp / (tp + fp) if tp + fp > 0 else 0\n            recall = tp / (tp + fn) if tp + fn > 0 else 0\n            fbeta = (1 + beta**2) * (precision * recall) / (beta**2 * precision + recall) if (precision + recall) > 0 else 0.0\n            aggregate_fbeta += fbeta * weights.get(particle_type, 1.0)\n    \n        if weights:\n            aggregate_fbeta = aggregate_fbeta / sum(weights.values())\n        else:\n            aggregate_fbeta = aggregate_fbeta / len(results)\n        return aggregate_fbeta\n    \n    class Metrics:\n        def __init__(\n            self,\n            mode: bool = \"multi_label\",\n            drop_beta_amylase: bool = True,\n            gt_csv_path: str = \"/kaggle/input/czii2024-csv/train.csv\",\n        ):\n            super().__init__()\n            self.class_infos = [\n                # dict(id=0, name=\"background\", radius=None, color=(0, 0, 0)),\n                dict(id=1, name=\"apo-ferritin\", radius=60, color=(255, 0, 0)),\n                dict(id=2, name=\"beta-amylase\", radius=65, color=(0, 255, 0)),\n                dict(id=3, name=\"beta-galactosidase\", radius=90, color=(0, 0, 255)),\n                dict(id=4, name=\"ribosome\", radius=150, color=(255, 255, 0)),\n                dict(id=5, name=\"thyroglobulin\", radius=130, color=(255, 0, 255)),\n                dict(id=6, name=\"virus-like-particle\", radius=135, color=(0, 255, 255)),\n            ]\n            self.num_classes = len(self.class_infos)\n            self.mode = mode\n            if self.mode != \"multi_label\":\n                raise NotImplementedError(\"Not implemented yet\")\n            self.drop_beta_amylase = drop_beta_amylase\n            self.gt_df = pd.read_csv(gt_csv_path)\n    \n        def _do_one_eval(self, truth: np.ndarray, predict: np.ndarray, threshold: float) -> tuple:\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        def _compute_lb(self, submit_df: pd.DataFrame, gt_df: pd.DataFrame, class_infos: list[dict]) -> tuple:\n            if len(submit_df) == 0:\n                return (\n                    pd.DataFrame(columns=[\"particle_type\", \"P\", \"T\", \"hit\", \"miss\", \"fp\", \"precision\", \"recall\", \"f-beta4\"]),\n                    pd.DataFrame(columns=[\"experiment\", \"particle_type\", \"type\", \"x\", \"y\", \"z\"]),\n                    pd.DataFrame(columns=[\"experiment\", \"particle_type\", \"type\", \"x\", \"y\", \"z\"]),\n                    0.0,\n                )\n    \n            valid_id = list(submit_df[\"experiment\"].unique())\n    \n            eval_df = []\n            confu_truth_df = []\n            confu_pred_df = []\n            for experiment in valid_id:\n                gt_exp_df = gt_df[gt_df[\"experiment\"] == experiment]\n                pred_exp_df = submit_df[submit_df[\"experiment\"] == experiment]\n                for class_info in class_infos:\n                    label_name = class_info[\"name\"]\n                    radius = class_info[\"radius\"]\n                    xyz_truth = gt_exp_df[gt_exp_df[\"particle_type\"] == label_name][[\"x\", \"y\", \"z\"]].values\n                    xyz_predict = pred_exp_df[pred_exp_df[\"particle_type\"] == label_name][[\"x\", \"y\", \"z\"]].values\n                    hit, fp, miss, metric = self._do_one_eval(xyz_truth, xyz_predict, radius * 0.5)\n                    eval_df.append(\n                        dict(\n                            experiment=experiment,\n                            particle_type=label_name,\n                            P=metric[0],\n                            T=metric[1],\n                            hit=metric[2],\n                            miss=metric[3],\n                            fp=metric[4],\n                        ),\n                    )\n    \n                    for tp_p in hit[0]:\n                        tp_pred_x, tp_pred_y, tp_pred_z = (xyz_predict[tp_p] / 10).round().astype(int)\n                        confu_pred_df.append(dict(experiment=experiment, particle_type=label_name, type=\"tp\", x=tp_pred_x, y=tp_pred_y, z=tp_pred_z))\n                    for fp_p in fp:\n                        fp_pred_x, fp_pred_y, fp_pred_z = (xyz_predict[fp_p] / 10).round().astype(int)\n                        confu_pred_df.append(dict(experiment=experiment, particle_type=label_name, type=\"fp\", x=fp_pred_x, y=fp_pred_y, z=fp_pred_z))\n                    for tp_t in hit[1]:\n                        tp_truth_x, tp_truth_y, tp_truth_z = (xyz_truth[tp_t] / 10).round().astype(int)\n                        confu_truth_df.append(dict(experiment=experiment, particle_type=label_name, type=\"tp\", x=tp_truth_x, y=tp_truth_y, z=tp_truth_z))\n                    for fn_t in miss:\n                        fn_truth_x, fn_truth_y, fn_truth_z = (xyz_truth[fn_t] / 10).round().astype(int)\n                        confu_truth_df.append(dict(experiment=experiment, particle_type=label_name, type=\"fn\", x=fn_truth_x, y=fn_truth_y, z=fn_truth_z))\n    \n            eval_df = pd.DataFrame(eval_df)\n            eval_df = eval_df.groupby(\"particle_type\").agg(\"sum\").drop(columns=[\"experiment\"])\n            eval_df.loc[:, \"precision\"] = eval_df[\"hit\"] / eval_df[\"P\"]\n            eval_df.loc[:, \"precision\"] = eval_df[\"precision\"].fillna(0)\n            eval_df.loc[:, \"recall\"] = eval_df[\"hit\"] / eval_df[\"T\"]\n            eval_df.loc[:, \"recall\"] = eval_df[\"recall\"].fillna(0)\n            eval_df.loc[:, \"f-beta4\"] = 17 * eval_df[\"precision\"] * eval_df[\"recall\"] / (16 * eval_df[\"precision\"] + eval_df[\"recall\"])\n            eval_df.loc[:, \"f-beta4\"] = eval_df[\"f-beta4\"].fillna(0)\n    \n            eval_df = eval_df.sort_values(\"particle_type\").reset_index(drop=False)\n            # https://www.kaggle.com/competitions/czii-cryo-et-object-identification/discussion/544895\n            eval_df.loc[:, \"weight\"] = [1, 0, 2, 1, 2, 1]\n            lb_score = (eval_df[\"f-beta4\"] * eval_df[\"weight\"]).sum() / eval_df[\"weight\"].sum()\n    \n            confu_pred_df = pd.DataFrame(confu_pred_df)\n            confu_truth_df = pd.DataFrame(confu_truth_df)\n            return eval_df, confu_pred_df, confu_truth_df, lb_score\n    \n        def update(self, scores: torch.Tensor, masks: torch.Tensor):\n            pass\n    \n        def compute(self, submit_df: pd.DataFrame) -> tuple:\n            lb_df, confu_pred_df, confu_truth_df, lb_score = self._compute_lb(submit_df, self.gt_df.copy(), self.class_infos)\n            return lb_score, lb_df, confu_pred_df, confu_truth_df\n    \n        def reset(self):\n            pass\n\n    gt_df = pd.read_csv(\"/kaggle/input/czii2024-csv/train.csv\")\n    val_run_names = [\"TS_6_4\"]\n    val_submit_df = submit_df[submit_df[\"experiment\"].isin(val_run_names)].copy().reset_index(drop=True)\n    val_submit_df['id'] = range(len(val_submit_df))\n    custom_metrics = Metrics()\n    lb_score, lb_df, _, _ = custom_metrics.compute(val_submit_df.copy())\n    print(f\"{val_run_names=}\")\n    print(f\"cv={lb_score:.4f}\")\n    print(tabulate(lb_df, headers=\"keys\", tablefmt=\"psql\"))","metadata":{"execution":{"iopub.status.busy":"2025-02-06T04:01:20.292057Z","iopub.execute_input":"2025-02-06T04:01:20.292285Z","iopub.status.idle":"2025-02-06T04:01:20.628630Z","shell.execute_reply.started":"2025-02-06T04:01:20.292252Z","shell.execute_reply":"2025-02-06T04:01:20.627900Z"},"trusted":true},"outputs":[],"execution_count":null}]}