{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"},{"sourceId":10476214,"sourceType":"datasetVersion","datasetId":6486959},{"sourceId":10657732,"sourceType":"datasetVersion","datasetId":6599725},{"sourceId":11381245,"sourceType":"datasetVersion","datasetId":6604854},{"sourceId":11926814,"sourceType":"datasetVersion","datasetId":6989840},{"sourceId":12043892,"sourceType":"datasetVersion","datasetId":7578946},{"sourceId":12046055,"sourceType":"datasetVersion","datasetId":7579615},{"sourceId":12051796,"sourceType":"datasetVersion","datasetId":7504947}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!mkdir /kaggle/temp\n!cp -r /kaggle/input/byu-final-dataset/byu-main/byu-main /kaggle/working/byu\n!cp -r /kaggle/input/dataset-czii-pkgs/pkgs /kaggle/temp","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T00:06:13.505988Z","iopub.execute_input":"2025-06-04T00:06:13.506326Z","iopub.status.idle":"2025-06-04T00:06:44.184422Z","shell.execute_reply.started":"2025-06-04T00:06:13.506296Z","shell.execute_reply":"2025-06-04T00:06:44.183027Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pwd\n!ls -la /kaggle/working/byu","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T00:06:44.185899Z","iopub.execute_input":"2025-06-04T00:06:44.186256Z","iopub.status.idle":"2025-06-04T00:06:44.415294Z","shell.execute_reply.started":"2025-06-04T00:06:44.186210Z","shell.execute_reply":"2025-06-04T00:06:44.414312Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n!cp -r /kaggle/input/czii-additional-packages/* /kaggle/temp/pkgs\n!pip install --no-index --find-links=/kaggle/temp/pkgs/ /kaggle/temp/pkgs/pretrainedmodels-0.7.4/pretrainedmodels-0.7.4\n!pip install --no-index --find-links=/kaggle/temp/pkgs/ /kaggle/temp/pkgs/fvcore-0.1.5.post20221221/fvcore-0.1.5.post20221221\n!pip install --no-index --find-links=/kaggle/temp/pkgs/ /kaggle/temp/pkgs/pytorchvideo-0.1.5/pytorchvideo\n!pip install --no-index --find-links=/kaggle/temp/pkgs/ /kaggle/temp/pkgs/detectron2-0.6/detectron2\n\n!pip install /kaggle/temp/pkgs/asciitree-0.3.3/asciitree-0.3.3\n!pip install --no-index --find-links=/kaggle/temp/pkgs/ hydra-core copick lightning monai simplejson iopath fvcore connected-components-3d\n!pip install --no-index /kaggle/input/tensorrt-10-1-0/tensorrt*.whl\n!pip install --no-index --find-links=/kaggle/input/onnx-wheels/ onnx onnxoptimizer onnxsim onnxruntime-gpu","metadata":{"trusted":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-06-04T00:06:44.417131Z","iopub.execute_input":"2025-06-04T00:06:44.417464Z","iopub.status.idle":"2025-06-04T00:09:57.377168Z","shell.execute_reply.started":"2025-06-04T00:06:44.417434Z","shell.execute_reply":"2025-06-04T00:09:57.376077Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%cd /kaggle/working/byu","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T00:09:57.378634Z","iopub.execute_input":"2025-06-04T00:09:57.379024Z","iopub.status.idle":"2025-06-04T00:09:57.386500Z","shell.execute_reply.started":"2025-06-04T00:09:57.378983Z","shell.execute_reply":"2025-06-04T00:09:57.385572Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nsys.path.extend([\n    '/kaggle/working/byu/yagm/src',\n    '/kaggle/working/byu/src',\n    '/kaggle/working/byu/',\n    '/kaggle/working/byu/third_party/segmentation_models_pytorch_3d',\n    '/kaggle/working/byu/third_party/slowfast',\n    '/kaggle/working/byu/third_party/timm_3d'\n])\nprint('SYS PATH:', sys.path, sep = '\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T00:09:57.387476Z","iopub.execute_input":"2025-06-04T00:09:57.387853Z","iopub.status.idle":"2025-06-04T00:09:57.733567Z","shell.execute_reply.started":"2025-06-04T00:09:57.387828Z","shell.execute_reply":"2025-06-04T00:09:57.732420Z"}},"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":"# TensorRT","metadata":{}},{"cell_type":"code","source":"!rm -rf /kaggle/working/byu/assets\n!ln -s /kaggle/input/byu-final-dataset /kaggle/working/byu/assets","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T23:17:02.412201Z","iopub.execute_input":"2025-06-03T23:17:02.412486Z","iopub.status.idle":"2025-06-03T23:17:02.642011Z","shell.execute_reply.started":"2025-06-03T23:17:02.412463Z","shell.execute_reply":"2025-06-03T23:17:02.641003Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import gc\n# import logging\n# import os\n# import random\n\n# import torch\n# from omegaconf import OmegaConf\n# from torch import nn\n# from torch.nn import functional as F\n# from yagm.utils import lightning as l_utils\n# from yagm.utils.torch2onnx import onnx2trt, torch2onnx\n\n# from byu.data.datasets.heatmap_3d_dataset import Heatmap3dDataset\n# from byu.data.io import OpencvTomogramLoader\n\n# logging.basicConfig(level=logging.INFO)\n\n# # DATA_DIR = \"/home/dangnh36/datasets/.comp/byu/\"\n# # DATASET_3D_CONFIG = f\"\"\"\n# # env:\n# #     data_dir: {DATA_DIR}\n# # cv:\n# #     strategy: skf4_rd42\n# #     num_folds: 4\n# #     fold_idx: 0\n# #     train_on_all: False\n# # loader:\n# #     num_workers: 1\n# # data:\n# #     label_fname: gt_v2  # gt | gt_v2 | gt_v3 | all_gt | all_gt_v3\n# #     patch_size: [224, 448, 448]\n# #     start: [0, 0, 0]\n# #     overlap: [0, 0, 0]\n# #     border: [0, 0, 0]\n\n# #     sigma: 0.2\n# #     fast_val_workers: 1\n# #     fast_val_prefetch: 1\n# #     io_backend: cv2  # cv2, cv2_seq, npy, cache\n# #     crop_outside: True\n# #     ensure_fg: False\n# #     label_smooth: [0.0,1.0]\n# #     filter_rule: null  # null | eq1 | le1\n\n# #     sampling:\n# #         method: pre_patch  # pre_patch | rand_crop\n# #         pre_patch:\n# #             fg_max_dup: 1\n# #             bg_ratio: 0.0\n# #             bg_from_pos_ratio: 0.01\n# #             overlap: [0, 0, 0]\n# #         rand_crop:\n# #             random_center: True\n# #             margin: 0.25\n# #             auto_correct_center: True\n# #             pos_weight: 1\n# #             neg_weight: 1\n# #     transform:\n# #         resample_mode: trilinear # F.grid_sample() mode\n# #         target_spacing: [16,16,16]\n\n# #         heatmap_mode: gaussian\n# #         heatmap_stride: [1,1,1]\n# #         heatmap_same_sigma: False\n# #         heatmap_same_std: False\n# #         lazy: True\n# #         device: null\n# #     aug:\n# #         enable: True\n# #         zoom_prob: 0.4\n# #         zoom_range: [[0.6, 1.2], [0.6, 1.2], [0.6, 1.2]]  # (X, Y, Z) or (H, W, D)\n# #         # affine1\n# #         affine1_prob: 0.5\n# #         affine1_scale: 0.3  # max_skew_xy = 1.3 / 0.7 = 1.86\n# #         # affine2\n# #         affine2_prob: 0.25\n# #         affine2_rotate_xy: 15 # degrees\n# #         affine2_scale: 0.3  # max_skew_xy = 1.3 / 0.7 = 1.86\n# #         affine2_shear: 0.2\n\n# #         rand_shift: False  # only used in rand_crop, very useful if random_center=auto_correct_center=False\n\n# #         # no lazy, can't properly transform points\n# #         grid_distort_prob: 0.0\n# #         smooth_deform_prob: 0.0\n\n# #         intensity_prob: 0.5\n# #         smooth_prob: 0.0\n# #         hist_equalize: False\n# #         downsample_prob: 0.2\n# #         coarse_dropout_prob: 0.1\n\n# #         # MIXER\n# #         mixup_prob: 0.0\n# #         cutmix_prob: 0.0\n# #         mixer_alpha: 1.0\n# #         mixup_target_mode: max\n# #     tta:\n# #         # enable: [zyx]\n# #         enable: [zyx, zxy, zyx_x, zyx_y]\n# # \"\"\"\n\n\n# # global_cfg = OmegaConf.create(DATASET_3D_CONFIG)\n# # dataset = Heatmap3dDataset(global_cfg, stage=\"train\")\n# # if dataset.stage != \"train\":\n# #     dataset.fast_val_tomo_loader.start()\n\n# # imgs = []\n# # random.seed(42)\n# # # select_idxs = random.choices(list(range(len(dataset))), k = 50)\n# # select_idxs = list(range(1))\n# # print(\"SELECT RANDOM INDICES:\", select_idxs)\n# # for idx in select_idxs:\n# #     sample = dataset[idx]\n# #     img = sample[\"image\"]\n# #     assert img.dtype == torch.uint8\n# #     print(idx, img.shape, img.dtype)\n# #     imgs.append(img[None].float().contiguous())\n# # del dataset\n# # gc.collect()\n\n\n# tomo_loader = OpencvTomogramLoader()\n# # TOMO_ROOT_DIR = \"/home/dangnh36/datasets/.comp/byu/raw/train/\"\n# TOMO_ROOT_DIR = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train/'\n# tomo = tomo_loader.load(os.path.join(TOMO_ROOT_DIR, \"tomo_bfd5ea\"))\n# tomo = (\n#     torch.from_numpy(tomo[None, 10:13, 0:896, 0:896]).float().contiguous()\n# )  # (1, 1, Z, Y, X)\n# print(tomo.shape)\n# imgs = [tomo]\n\n\n# ############### LOAD MODEL ###############\n# DEVICE = \"cuda:0\"\n# CONVERT_ONNX = True\n# CONVERT_TRT = True\n\n\n\n\n\n\n# CONFIG_PATH = \"/kaggle/input/byu-final-dataset/EXP25_COATLITEMEDIUM_ALLGTV3_config.yaml\"\n# WEIGHT_PATH = \"/kaggle/input/byu-final-dataset/EXP25_COATLITEMEDIUM_ALLGTV3_ep4_step20000.ckpt\"\n# TRT_SAVE_PATH = '/kaggle/working/EXP25_COATLITEMEDIUM_ALLGTV3_ep4_step20000.engine'\n# WORKSPACE_SIZE = int(5 * (2**30))\n# PRECISION = 'fp16'\n\n\n\n# ONNX_SAVE_PATH = None\n# MIN_BS, OPT_BS, MAX_BS = 8, 8, 8\n\n# if TRT_SAVE_PATH is None:\n#     TRT_SAVE_PATH = WEIGHT_PATH.replace(\".ckpt\", \".engine\")\n    \n# if ONNX_SAVE_PATH is None:\n#     ONNX_SAVE_PATH = TRT_SAVE_PATH.replace(\".engine\", \".onnx\")\n\n# os.makedirs(os.path.dirname(ONNX_SAVE_PATH), exist_ok=True)\n# os.makedirs(os.path.dirname(TRT_SAVE_PATH), exist_ok=True)\n\n\n# if CONVERT_ONNX:\n#     device = torch.device(DEVICE)\n#     cfg = OmegaConf.load(CONFIG_PATH)\n\n#     ######### OVERWRITE SOME CONFIGS ##########\n#     cfg.misc.log_model = False\n#     cfg.ema.val_decays = [0.99]\n#     cfg.ckpt.strict = False\n\n#     ######### OVERWRITE SOME CONFIGS ##########\n#     # OVERWRITE SOME CONFIG\n#     try:\n#         cfg.model.encoder.pretrained = False\n#     except:\n#         pass\n\n#     # for backward-compatibility\n#     OmegaConf.set_readonly(cfg, False)\n#     # just keep outputing heatmap\n#     cfg.model.reg_head.enable = False\n#     cfg.model.dsnt = OmegaConf.create()\n#     cfg.model.dsnt.enable = False\n#     OmegaConf.set_readonly(cfg, True)\n#     ##########\n\n#     task = l_utils.build_task(cfg)\n#     print(task.device)\n#     l_utils.load_lightning_state_dict(\n#         model=task,\n#         ckpt_path=WEIGHT_PATH,\n#         cfg=cfg,\n#     )\n#     print(f\"Loaded state dict\")\n\n#     class ModelWrapper(nn.Module):\n#         def __init__(self, model):\n#             super().__init__()\n#             self._model = model\n\n#         def forward(self, x):\n#             ret = self._model(x)\n#             assert isinstance(ret, (list, tuple)) and len(ret) == 4\n#             return ret[2]\n\n#     torch_model = ModelWrapper(task.model)\n#     print(torch_model)\n#     torch_model.eval().to(device)\n#     print('DEVICE:', device)\n\n#     torch2onnx(\n#         torch_model=torch_model,\n#         sample_inputs=[imgs[0]],\n#         input_names=[\"image\"],\n#         output_names=[\"heatmap\"],\n#         save_path=ONNX_SAVE_PATH,\n#         precision=\"fp32\",\n#         device=DEVICE,\n#         dynamic_batching=True,\n#         batch_axis=0,\n#         validate_fn=None,\n#         rtol=1e-3,\n#         atol=1e-4,\n#         verbose=True,\n#     )\n#     del torch_model\n#     gc.collect()\n#     torch.cuda.empty_cache()\n\n\n# if CONVERT_TRT:\n#     try:\n#         engine = onnx2trt(\n#             onnx_file_path=ONNX_SAVE_PATH,\n#             engine_file_path=TRT_SAVE_PATH,\n#             input_shape=imgs[0].shape[1:],\n#             precision=PRECISION,\n#             min_batch=MIN_BS,\n#             opt_batch=OPT_BS,\n#             max_batch=MAX_BS,\n#             enable_dynamic_batching=True,\n#             workspace_size=WORKSPACE_SIZE,\n#         )\n\n#         if engine is not None:\n#             print(\"Conversion successful!\")\n#         else:\n#             print(\"Conversion failed!\")\n#     except Exception as e:\n#         print(f\"Conversion error: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T00:09:57.734892Z","iopub.execute_input":"2025-06-04T00:09:57.735211Z","iopub.status.idle":"2025-06-04T00:09:57.746346Z","shell.execute_reply.started":"2025-06-04T00:09:57.735174Z","shell.execute_reply":"2025-06-04T00:09:57.745529Z"},"jupyter":{"source_hidden":true}},"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":"","metadata":{"trusted":true},"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":"","metadata":{"trusted":true},"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":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3D CODE","metadata":{}},{"cell_type":"code","source":"%%writefile /kaggle/working/byu/submit_3d.py\n\nimport sys\nsys.path.extend([\n    '/kaggle/working/byu/yagm/src',\n    '/kaggle/working/byu/src',\n    '/kaggle/working/byu/',\n    '/kaggle/working/byu/third_party/segmentation_models_pytorch_3d',\n    '/kaggle/working/byu/third_party/slowfast',\n    '/kaggle/working/byu/third_party/timm_3d'\n])\nprint('SYS PATH:', sys.path, sep = '\\n')\n\n\nimport warnings\n\nwarnings.simplefilter(\"ignore\")\nimport gc\nimport json\nimport logging\nimport multiprocessing as mp\nimport os\nimport queue\nimport shutil\nimport sys\nimport threading\nimport time\nfrom functools import partial\n\nimport cc3d\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom omegaconf import OmegaConf\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.utils.data import default_collate\nfrom tqdm import tqdm\nfrom yagm.tasks.base_task import BaseTask\nfrom yagm.transforms.keypoints.decode import decode_heatmap_3d, decode_segment_mask_3d\nfrom yagm.transforms.sliding_window import get_sliding_patch_positions\nfrom yagm.transforms.tta_3d import build_tta\nfrom yagm.utils import hydra as hydra_utils\nfrom yagm.utils import lightning as l_utils\nfrom yagm.utils.concurrent import ShmTensor\n\n# Setting up custom logging\nfrom yagm.utils.logging import init_logging, setup_logging\n\nfrom byu.data.io import MultithreadOpencvTomogramLoader\nfrom byu.inference import zoo\nfrom byu.inference.tensorrt_engine import ThreadsafeTRTEngine\nfrom byu.utils.data import SubmissionDataFrame\nfrom byu.utils.metrics import compute_metrics, kaggle_score\nfrom byu.utils.viz import viz_byu\nfrom byu.utils.wbf import weighted_boxes_fusion_3d\n\ntry:\n    hydra_utils.init_hydra()\nexcept:\n    print(\"SKIP RE-INIT HYDRA\")\n\n### SETUP LOGGER FOR IPYTHON ###\n# ref: https://github.com/ipython/ipykernel/issues/111\n# Create logger\nlogger = logging.getLogger()\nlogger.setLevel(logging.INFO)\n# Create STDERR handler\nhandler = logging.StreamHandler(sys.stderr)\n# Create formatter and add it to the handler\nformatter = logging.Formatter(\"[%(levelname)s] %(message)s\")\nhandler.setFormatter(formatter)\n# Set STDERR handler as the only handler\nlogger.handlers = [handler]\n\n\n##################################################\n###### MODEL SIGNATURES OVERRIDES\n##################################################\nALL_SUPPORTED_TTAS = [\n    \"zyx\",\n    \"zyx_x\",\n    \"zyx_y\",\n    \"zyx_z\",\n    \"zyx_xy\",\n    \"zyx_xz\",\n    \"zyx_yz\",\n    \"zyx_xyz\",\n    \"zxy\",\n    \"zxy_x\",\n    \"zxy_y\",\n    \"zxy_z\",\n    \"zxy_xy\",\n    \"zxy_xz\",\n    \"zxy_yz\",\n    \"zxy_xyz\",\n]\n\n_SIG_TEMPLATE = {\n    \"backend\": \"torch\",  # torch | trt\n    \"trt_path\": \"\",\n    \"config_path\": \"\",\n    \"torch_path\": \"\",\n    \"ema\": 0.99,\n}\n\n# DENSENET121 (LB 86.0)\n# https://www.kaggle.com/code/dangnh0611/submit-heatmap-final?scriptVersionId=243314913\nEXP28_DENSNET121_ALLGTV3 = {\n    \"backend\": \"torch\",  # torch | trt\n    \"trt_path\": \"\",\n    \"config_path\": \"/kaggle/input/byu-final-dataset/EXP28_DENSNET121_ALLGTV3_config.yaml\",\n    \"torch_path\": \"/kaggle/input/byu-final-dataset/EXP28_DENSNET121_ALLGTV3_ep1_step32000.ckpt\",\n    \"ema\": 0.99,\n}\n\n# X3D (LB 85.8)\n# https://www.kaggle.com/code/dangnh0611/submit-heatmap-final?scriptVersionId=242747515\nEXP16_X3DM_ALL_TRAIN_EXT = {\n    \"backend\": \"torch\",  # torch | trt\n    \"trt_path\": \"\",\n    \"config_path\": \"/kaggle/input/byu-final-dataset/EXP16_X3DM_ALL_TRAIN_EXT_config.yaml\",\n    \"torch_path\": \"/kaggle/input/byu-final-dataset/EXP16_X3DM_ALL_TRAIN_EXT_ep1_step30000.ckpt\",\n    \"ema\": 0.99,\n}\n\n# RESNEXT50 (LB 85.8)\n# https://www.kaggle.com/code/dangnh0611/submit-heatmap-final?scriptVersionId=243403282\nEXP31_RESNEXT50_ALLGTV3 = {\n    \"backend\": \"torch\",  # torch | trt\n    \"trt_path\": \"\",\n    \"config_path\": \"/kaggle/input/byu-final-dataset/EXP31_RESNEXT50_ALLGTV3_config.yaml\",\n    \"torch_path\": \"/kaggle/input/byu-final-dataset/EXP31_RESNEXT50_ALLGTV3_ep1_step32000.ckpt\",\n    \"ema\": 0.99,\n}\n\n# R50 (LB 85.4)\n# https://www.kaggle.com/code/dangnh0611/submit-heatmap-final?scriptVersionId=243127441\nEXP26_R50_ALLGTV3 = {\n    \"backend\": \"torch\",  # torch | trt\n    \"trt_path\": \"\",\n    \"config_path\": \"/kaggle/input/byu-final-dataset/EXP26_R50_ALLGTV3_config.yaml\",\n    \"torch_path\": \"/kaggle/input/byu-final-dataset/EXP26_R50_ALLGTV3_ep1_step32000.ckpt\",\n    \"ema\": 0.99,\n}\n\n# 2.5D CONVNEXT_TINY_AVG (LB 83.7), for diversity\n# https://www.kaggle.com/code/dangnh0611/submit-heatmap-final?scriptVersionId=243123609\nEXP24_CONVNEXT_TINY_AVG_ALLGTV3 = {\n    \"backend\": \"torch\",  # torch | trt\n    \"trt_path\": \"\",\n    \"config_path\": \"/kaggle/input/byu-final-dataset/EXP24_CONVNEXT_TINY_AVG_ALLGTV3_config.yaml\",\n    \"torch_path\": \"/kaggle/input/byu-final-dataset/EXP24_CONVNEXT_TINY_AVG_ALLGTV3_ep2_step36000.ckpt\",\n    \"ema\": 0.99,\n}\n\n\n\nJOBS = [\n    [\n        {\n            \"sig\": EXP28_DENSNET121_ALLGTV3,\n            \"tta\": [\"zyx\", \"zxy\"],\n            \"weight\": [1.0, 1.0],\n            \"agg_name\": [\"DENSENET-zyx\", \"DENSENET-zxy\"],\n        },\n        {\n            \"sig\": EXP16_X3DM_ALL_TRAIN_EXT,\n            \"tta\": [\"zyx_x\", \"zyx_y\"],\n            \"weight\": [1.0, 1.0],\n            \"agg_name\": [\"X3D-zyx_x\", \"X3D-zyx_y\"],\n        },\n        {\n            \"sig\": EXP31_RESNEXT50_ALLGTV3,\n            \"tta\": [\"zyx_xy\"],\n            \"weight\": [1.0],\n            \"agg_name\": [\"RESNEXT50-zyx_xy\"],\n        },\n        # {\n        #     \"sig\": EXP26_R50_ALLGTV3,\n        #     \"tta\": [\"zxy_z\"],\n        #     \"weight\": [1.0],\n        #     \"agg_name\": [\"R50-zyx_xy\"],\n        # },\n        {\n            \"sig\": EXP24_CONVNEXT_TINY_AVG_ALLGTV3,\n            \"tta\": [\"zxy_xy\"],\n            \"weight\": [1.0],\n            \"agg_name\": [\"CONVNEXT-zxy_xy\"],\n        },\n    ]\n]\n\n\n\n# JOBS = [\n#     [\n#         {\n#             \"sig\": EXP28_DENSNET121_ALLGTV3,\n#             \"tta\": [\"zyx\", \"zxy\"],\n#             \"weight\": [1.0, 1.0],\n#             \"agg_name\": [\"DENSENET-zyx\", \"DENSENET-zxy\"],\n#         },\n#         {\n#             \"sig\": EXP26_R50_ALLGTV3,\n#             \"tta\": [\"zyx_x\", \"zyx_y\"],\n#             \"weight\": [1.0, 1.0],\n#             \"agg_name\": [\"R50-zyx_x\", \"R50-zyx_y\"],\n#         },\n#         {\n#             \"sig\": EXP16_X3DM_ALL_TRAIN_EXT,\n#             \"tta\": [\"zyx_xy\", \"zxy_xy\"],\n#             \"weight\": [1.0, 1.0],\n#             \"agg_name\": [\"X3D-zyx_xy\", \"X3D-zxy_xy\"],\n#         },\n#     ]\n# ]\n\n\n##################################################\n###### GLOBAL VARS\n##################################################\n_CV2_TOMO_LOADER_NUM_THREADS = 8\n\nPATCH_SIZE = [224, 448, 448]\nPATCH_BORDER = [0, 0, 0]\nPATCH_OVERLAP = [0, 0, 0]\n\nIO_BACKEND = \"cv2\"\nBASE_SPACING = 13.1  # avg 15.6\n# TARGET_SPACINGS = [13.1, 16.0, 19.7]\nTARGET_SPACINGS = [[16.0, 16.0, 16.0]]\nTOMO_SPACING_DEVICE = \"cpu\"\nTOMO_SPACING_MODE = \"trilinear\"\n\nHEATMAP_AGG_MODE = \"avg\"\nHEATMAP_AGG_LOGITS = True\nHEATMAP_INTERPOLATION_MODE = \"trilinear\"\nFORWARD_SAVE_MEMORY = False\nBATCH_SIZE = 1\nHEATMAP_STRIDE = 8\n\n# MASK CC3D DECODING\nDECODE_CC3D_CONF_THRES = 0.05\nDECODE_CC3D_RADIUS_FACTOR = 0.1\nDECODE_CC3D_CONF_MODE = \"prob\"  # volume | prob | fuse\nDECODE_CC3D_PROB_MODE = \"center\"  # center | mean | max\n# HEATMAP NMS DECODING\nDECODE_NMS_BLUR_SIGMA = None\n###### MAIN DECODE PARAMS ######\nDECODE_METHOD = \"nms\"  # nms, cc3d\nDECODE_CONF_THRES = 0.05\nQUANTILE_THRES = 0.55\n\nSAVE_HEATMAP_NPY = False\nDECODE_HEATMAP_TO_CSV = True\nCSV_ENSEMBLE_MODE = \"wbf\"  # max | wbf\nWBF_CONF_TYPE = \"avg\"  # avg | max\nWBF_CONF_THRES = 0.2\n\nVIZ_ENABLE = False\n\n\n# automatic determine the MODE\nif os.path.isdir(\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025\"):\n    MODE = \"KAGGLE\"\nelse:\n    MODE = \"LOCAL\"\n# manual override\n# MODE = \"KAGGLE_VAL\"\n\n# LOCAL DEVELOPMENT ENV\nif MODE == \"LOCAL\":\n    DEVICES = [0]\n    ASSETS_DIR = \"assets/\"\n    WORKING_DIR = \"outputs/submit/working/\"\n\n    DATA_DIR = \"/home/dangnh36/datasets/.comp/byu/raw/train\"\n    VAL_GT_CSV_PATH = \"/home/dangnh36/datasets/.comp/byu/processed/gt_v3.csv\"\n    with open(\n        \"/home/dangnh36/datasets/.comp/byu/processed/cv/v3/skf4_rd42.json\", \"r\"\n    ) as f:\n        cv_meta = json.load(f)\n        ALL_TOMO_IDS = cv_meta[\"folds\"][0][\"val\"][:]  # fold 0\n\n    # DATA_DIR = \"/home/dangnh36/datasets/.comp/byu/processed/pseudo_test/tomograms\"\n    # VAL_GT_CSV_PATH = \"/home/dangnh36/datasets/.comp/byu/processed/pseudo_test/gt.csv\"\n    # ALL_TOMO_IDS = sorted(os.listdir(DATA_DIR)) * 10\n\n    GT_DF = pd.read_csv(VAL_GT_CSV_PATH)\n    GT_DF = GT_DF[GT_DF[\"tomo_id\"].isin(ALL_TOMO_IDS)].reset_index(drop=True)\n    _tomo2spacing = {}\n    for i, row in GT_DF.iterrows():\n        _tomo2spacing[row[\"tomo_id\"]] = row[\"voxel_spacing\"]\n    ALL_TOMO_SPACINGS = [_tomo2spacing[tomo_id] for tomo_id in ALL_TOMO_IDS]\n\n    # TOMO_PRODUCER_NUM_WORKERS = 8\n    # TOMO_PRODUCER_PREFETCH = 8\n    # PATCHES_PRODUCER_NUM_WORKERS = 2\n    # PATCHES_PRODUCER_PREFETCH = 8\n    # DATALOADER_PREFETCH = 8\n\n    TOMO_PRODUCER_NUM_WORKERS = 1\n    TOMO_PRODUCER_PREFETCH = 1\n    PATCHES_PRODUCER_NUM_WORKERS = 1\n    PATCHES_PRODUCER_PREFETCH = 2\n    DATALOADER_PREFETCH = 8\n\n# VALIDATION ON KAGGLE\nelif MODE == \"KAGGLE_VAL\":\n    DEVICES = [0, 1]\n    ASSETS_DIR = \"/kaggle/input/byu-final-dataset/\"\n    DATA_DIR = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train\"\n    WORKING_DIR = \"/kaggle/working/\"\n    VAL_GT_CSV_PATH = \"/kaggle/input/byu-checkpoints/gt_v2.csv\"\n    with open(\"/kaggle/input/byu-checkpoints/skf4_rd42.json\", \"r\") as f:\n        cv_meta = json.load(f)\n        ALL_TOMO_IDS = cv_meta[\"folds\"][0][\"val\"][:50]  # fold 0, top 10 first tomo\n    GT_DF = pd.read_csv(VAL_GT_CSV_PATH)\n    GT_DF = GT_DF[GT_DF[\"tomo_id\"].isin(ALL_TOMO_IDS)].reset_index(drop=True)\n    _tomo2spacing = {}\n    for i, row in GT_DF.iterrows():\n        _tomo2spacing[row[\"tomo_id\"]] = row[\"voxel_spacing\"]\n    ALL_TOMO_SPACINGS = [_tomo2spacing[tomo_id] for tomo_id in ALL_TOMO_IDS]\n\n    TOMO_PRODUCER_NUM_WORKERS = 1\n    TOMO_PRODUCER_PREFETCH = 1\n    PATCHES_PRODUCER_NUM_WORKERS = 1\n    PATCHES_PRODUCER_PREFETCH = 2\n    DATALOADER_PREFETCH = 8\n\n# PRIVATE TEST SUBMISSION\nelif MODE == \"KAGGLE\":\n    DEVICES = [0, 1]\n    ASSETS_DIR = \"/kaggle/input/byu-final-dataset/\"\n    DATA_DIR = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test\"\n    WORKING_DIR = \"/kaggle/working/\"\n    ALL_TOMO_IDS = sorted(os.listdir(DATA_DIR))\n    ALL_TOMO_SPACINGS = [BASE_SPACING] * len(ALL_TOMO_IDS)\n\n    TOMO_PRODUCER_NUM_WORKERS = 1\n    TOMO_PRODUCER_PREFETCH = 1\n    PATCHES_PRODUCER_NUM_WORKERS = 1\n    PATCHES_PRODUCER_PREFETCH = 2\n    DATALOADER_PREFETCH = 8\n\nelse:\n    raise ValueError\n\nTMP_CSV_DIR = os.path.join(WORKING_DIR, \"tmp_csv\")\nTMP_VIZ_DIR = os.path.join(WORKING_DIR, \"tmp_viz\")\nSAVE_HEATMAP_NPY_DIR = os.path.join(\"/kaggle/temp/tmp_heatmap\")\n\n# RECHECK SOME CONDITIONS\nassert len(ALL_TOMO_IDS) == len(ALL_TOMO_SPACINGS)\n# assert not (SAVE_HEATMAP_NPY and DECODE_HEATMAP_TO_CSV) and (SAVE_HEATMAP_NPY or DECODE_HEATMAP_TO_CSV)\n\n\n###################################################\n###### FUNCTIONS/CLASSES DEFINITION\n###################################################\ndef exception_printer(func):\n    def wrapper(*args, **kwargs):\n        try:\n            return func(*args, **kwargs)\n        except Exception as e:\n            logger.exception(\"Exception %s in %s: %s\", type(e), func.__name__, e)\n            import traceback\n\n            traceback.print_exc()\n            raise e\n\n    return wrapper\n\n\ndef clear_dir_content(folder_path):\n    if os.path.isdir(folder_path):\n        for item in os.listdir(folder_path):\n            item_path = os.path.join(folder_path, item)\n            if os.path.isdir(item_path):\n                shutil.rmtree(item_path)\n            else:\n                os.remove(item_path)\n\n\nclass Model3dWrapper(nn.Module):\n    def __init__(self, sig, gpu_id, act=\"sigmoid\"):\n        super().__init__()\n        if sig[\"backend\"] == \"trt\":\n            logger.info(\"Loading TRT model with config:\\n%s\", sig)\n            self.model = ThreadsafeTRTEngine(sig[\"trt_path\"], gpu_id)\n        elif sig[\"backend\"] == \"torch\":\n            cfg = OmegaConf.load(sig[\"config_path\"])\n            cfg.misc.log_model = False\n            cfg.ema.val_decays = [sig[\"ema\"] if sig[\"ema\"] is not None else 0]\n            cfg.model.head.ms = [0]\n            cfg.ckpt.strict = True\n\n            ######### OVERWRITE SOME CONFIGS ##########\n            if (\n                \"x3d\" in cfg.model.encoder._target_\n                or \"i3d\" in cfg.model.encoder._target_\n            ):\n                cfg.model.encoder.pretrained = None\n            elif \"smp\" in cfg.model.encoder._target_:\n                cfg.model.encoder.weights = None\n                if \"resnet101\" in cfg.model.encoder.model_name:\n                    cfg.model.encoder.model_name = \"resnet101\"\n            else:\n                pass\n            # for 2.5D encoder\n            try:\n                cfg.model.encoder.encoder_2d.pretrained = False\n            except:\n                pass\n            try:\n                if (\n                    cfg.model.encoder.encoder_2d.model_name\n                    == \"convnext_nano.r384_ad_in12k\"\n                ):\n                    cfg.model.encoder.encoder_2d.model_name = \"convnext_nano\"\n            except:\n                pass\n            ##########\n\n            task: BaseTask = l_utils.build_task(cfg)\n            l_utils.load_lightning_state_dict(\n                model=task,\n                ckpt_path=sig[\"torch_path\"],\n                cfg=cfg,\n            )\n            logger.info(\"Loaded Pytorch state dict from %s\", sig[\"torch_path\"])\n            # self.model = TorchModelWrapper(task.model)\n            self.model = task.model\n            del task\n            gc.collect()\n        else:\n            raise ValueError\n\n        self.act_name = act\n        if act == \"sigmoid\":\n            self.act = F.sigmoid\n            self.nan_to_num = partial(torch.nan_to_num, nan=0.0)\n        elif act == \"identity\":\n            self.act = lambda x: x\n            self.nan_to_num = partial(\n                torch.nan_to_num, nan=torch.finfo(torch.float32).min\n            )\n        else:\n            raise ValueError\n\n    def forward(self, x):\n        heatmap = self.act(self.model(x)[0])\n        heatmap = self.nan_to_num(heatmap)\n        # if torch.isnan(heatmap).any():\n        #     print('\\n\\nNAN!!!')\n        #     raise Exception\n        return heatmap\n\n\ndef decode_heatmap(pred_heatmap, target_spacing, stride, blur_operator=None):\n    radius_voxel = [1000.0 / e / stride for e in target_spacing]\n    radius_thres = max(radius_voxel)\n    # @TODO - currently, using fixed pool_ksize=3 due to performance issue\n    # when using large kernel size with Pytorch (>10 sec on 224x448x448)\n\n    # # maximum ood number which <= radius_thres\n    # # 0.7071067811865475 = 1 / sqrt(2)\n    # pool_ksize = int((2 * radius_thres * 0.7071067811865475 - 1) // 2 * 2 + 1)\n    # # at least 3, [1,1,1] == no pooling, which increase maxRecall but significantly reduce other metrics\n    # pool_ksize = max(3, pool_ksize)\n    # pool_ksize = [pool_ksize, pool_ksize, pool_ksize]\n\n    pool_ksize = [3, 3, 3]\n    logger.debug(\n        \"Heatmap decode with radius=%s, pool_ksize=%s, radius_thres=%s\",\n        radius_voxel,\n        pool_ksize,\n        radius_thres,\n    )\n    del target_spacing\n    assert pred_heatmap.shape[0] == 1\n    ret = []\n    for channel_idx, heatmap in enumerate(pred_heatmap):\n        outputs = decode_heatmap_3d(\n            heatmap=heatmap,\n            pool_ksize=pool_ksize,\n            nms_radius_thres=radius_thres,\n            blur_operator=blur_operator,\n            conf_thres=DECODE_CONF_THRES,\n            # max_dets=5 if (VIZ_ENABLE and MODE != \"KAGGLE\") else 1,\n            max_dets=100,\n            timeout=None,\n        )\n        outputs = outputs.cpu()\n        ret.append(outputs)\n    return ret\n\n\ndef decode_segment_mask(\n    pred_heatmap,\n    target_spacing,\n    stride,\n    prob_thres,\n    radius_factor_thres,\n    conf_mode=\"fuse\",\n    prob_mode=\"avg\",\n):\n    radius_voxel_thres = [\n        1000 / e / stride * radius_factor_thres for e in target_spacing\n    ]\n    volume_thres = (\n        4\n        / 3\n        * np.pi\n        * radius_voxel_thres[0]\n        * radius_voxel_thres[1]\n        * radius_voxel_thres[2]\n    )\n    assert len(pred_heatmap.shape) == 4 and pred_heatmap.shape[0] == 1\n    ret = []\n    for channel_idx, heatmap in enumerate(pred_heatmap):\n        keypoints = decode_segment_mask_3d(\n            heatmap,\n            prob_thres,\n            volume_thres,\n            max_dets=100,\n            conf_mode=conf_mode,\n            prob_mode=prob_mode,\n        )\n        ret.append(keypoints)\n    return ret\n\n\ndef crop_tomo_patch(tomo, patch_position):\n    # start = time.time()\n    tomo_shape = tomo.shape\n    _roi_start, _roi_end, patch_start, patch_end = patch_position\n    top_pad_z, top_pad_y, top_pad_x = [max(-start, 0) for start in patch_start]\n    bot_pad_z, bot_pad_y, bot_pad_x = [\n        max(0, end - size) for end, size in zip(patch_end, tomo_shape)\n    ]\n    actual_crop_start = [max(0, start) for start in patch_start]\n    actual_crop_end = [min(end, size) for end, size in zip(patch_end, tomo_shape)]\n    crop_slices = tuple(\n        slice(start, end) for start, end in zip(actual_crop_start, actual_crop_end)\n    )\n    crop = tomo[crop_slices]\n    # pad if needed\n    pad = (top_pad_x, bot_pad_x, top_pad_y, bot_pad_y, top_pad_z, bot_pad_z)\n    if any(pad):\n        crop = F.pad(crop, pad, mode=\"constant\", value=0)\n    # end = time.time()\n    # logger.debug(\"Crop take %f sec: %s\", end - start, patch_position)\n    return crop\n\n\n@torch.inference_mode()\ndef spacing_torch(\n    ori_tomo, ori_spacing, target_spacing, device=\"cuda\", mode=\"trilinear\"\n):\n    if tuple(ori_spacing) == tuple(target_spacing):\n        return ori_tomo\n    # NOTE: https://docs.opencv.org/3.4/da/d54/group__imgproc__transform.html#ga47a974309e9102f5f08231edc7e7529d\n    # shrink -> AREA, enlarge -> CUBIC or LINEAR\n    assert mode in [\"trilinear\", \"nearest\", \"area\", \"nearest-exact\"]\n    assert len(ori_tomo.shape) == 3\n    ori_tomo = ori_tomo[None, None].to(device)  # 11ZYX\n    if mode != \"nearest\":\n        ori_tomo = ori_tomo.float()\n\n    scale_factor = tuple(ori / tgt for ori, tgt in zip(ori_spacing, target_spacing))\n    ##### ISSUES #####\n    ### with interpolation mode `nearest`\n    # RuntimeError: upsample_nearest3d only supports output tensors with less than INT_MAX elements, but got [1, 1, 1178, 1367, 1367]\n    # ref: https://github.com/pytorch/pytorch/issues/144855\n    ### with interpolation mode `trilinear` on GPU\n    # RuntimeError: CUDA error: invalid configuration argument\n    ##### CURRENT SIMPLE FIX: if CUDA-interpolation-kernel fails, use CPU instead\n    try:\n        start = time.time()\n        spaced_tomo = F.interpolate(\n            ori_tomo,\n            size=None,\n            scale_factor=scale_factor,\n            mode=mode,\n            align_corners=False if mode == \"trilinear\" else None,\n            recompute_scale_factor=False,\n        )[0, 0]\n        end = time.time()\n        logger.debug(\"INTERPOLATE WITH MODE=%s TAKE %.2f sec\", mode, end - start)\n    except RuntimeError as e:\n        logger.warning(\n            \"EXCEPTION with torch.interpolate():\\n%s\\nAttempt CPU interpolation..\", e\n        )\n        spaced_tomo = F.interpolate(\n            ori_tomo.cpu(),\n            size=None,\n            scale_factor=scale_factor,\n            mode=mode,\n            align_corners=False if mode == \"trilinear\" else None,\n            recompute_scale_factor=False,\n        )[0, 0]\n    if mode != \"nearest\":\n        spaced_tomo = spaced_tomo.to(torch.uint8)  # ZYX\n    spaced_tomo = spaced_tomo.cpu()\n    assert spaced_tomo.dtype == torch.uint8\n    return spaced_tomo\n\n\ndef scale_range(start, end, scale, minv, maxv):\n    n = round((end - start) * scale)\n    if round(end * scale) - round(start * scale) == n:\n        return round(start * scale), round(end * scale)\n    start1 = int(start * scale)\n    start2 = round(start * scale)\n    candidate_starts = []\n    if minv <= start1 < start1 + n <= maxv:\n        candidate_starts.append(start1)\n    if minv <= start2 < start2 + n <= maxv:\n        candidate_starts.append(start2)\n    if len(candidate_starts) == 1:\n        return candidate_starts[0], candidate_starts[0] + n\n    elif len(candidate_starts) == 2:\n        if candidate_starts[0] == candidate_starts[1]:\n            return candidate_starts[0], candidate_starts[0] + n\n        diffs = [\n            abs(s - start * scale) + abs(s + n - end * scale) for s in candidate_starts\n        ]\n        if diffs[0] < diffs[1]:\n            return candidate_starts[0], candidate_starts[0] + n\n        else:\n            return candidate_starts[1], candidate_starts[1] + n\n    else:\n        raise ValueError\n\n\n@exception_printer\ndef tomo_producer_worker_func(\n    in_queue: mp.Queue,\n    out_queue: queue.Queue,\n    num_remain_workers: mp.Value,\n    lock: mp.Lock,\n    device: str = \"cpu\",\n    mode: str = \"trilinear\",\n    worker_id: str = \"\",\n):\n    \"\"\"Read tomo, spacing, then save to shared memory.\n    Should be run at process level (multiprocessing)\n    \"\"\"\n    logger.info(\"[TOMO PRODUCER %s] STARTING..\", worker_id)\n    _old_num_threads = torch.get_num_threads()\n    if _old_num_threads < 16:\n        torch.set_num_threads(16)\n    logger.info(\n        \"[TOMO PRODUCER %s] Number of threads change from %d to %d\",\n        worker_id,\n        _old_num_threads,\n        torch.get_num_threads(),\n    )\n    tomo_loader = MultithreadOpencvTomogramLoader(\n        num_workers=_CV2_TOMO_LOADER_NUM_THREADS\n    )\n    while True:\n        item = in_queue.get()\n        if item is None:\n            with lock:\n                num_remain_workers.value -= 1\n                if num_remain_workers.value == 0:\n                    # all other workers are stopped, this worker is the last one alive\n                    # put None to indicate stopping to out_queue's consumer\n                    out_queue.put(None)\n            # other workers should stop too..\n            in_queue.put(None)\n            logger.debug(\n                \"[TOMO PRODUCER %s] STOPPING..\",\n                worker_id,\n            )\n            return\n        start = time.time()\n        tomo_idx, tomo_id, ori_spacing = item\n        # fetch new item\n        tomo = tomo_loader.load(os.path.join(DATA_DIR, tomo_id))\n        ori_shape = tomo.shape\n        target_spacings = TARGET_SPACINGS\n        tomo = torch.from_numpy(tomo)\n        assert tomo.dtype == torch.uint8\n\n        # FIX BUG WRONG QUANTILE\n        hist = torch.bincount(tomo.view(-1), minlength=256).cpu().numpy()\n        assert hist.shape[0] == 256\n        CUTOFF = 25\n        is_wrong_quantile_dark = hist[:CUTOFF].sum() / hist[CUTOFF:].sum() > 0.08\n        is_wrong_quantile_light = hist[-CUTOFF:].sum() / hist[:-CUTOFF].sum() > 0.08\n        if is_wrong_quantile_dark or is_wrong_quantile_light:\n            tomo = tomo + 127\n\n        # Multiscale TTA\n        logger.debug(\n            \"[TOMO PRODUCER %s] Read tomo_idx=%d tomo_id=%s tomo_shape=%s\",\n            worker_id,\n            tomo_idx,\n            tomo_id,\n            tomo.shape,\n        )\n        proxy_tomos = []\n        for spacing_idx, target_spacing in enumerate(target_spacings):\n            spaced_tomo = spacing_torch(\n                tomo, ori_spacing, target_spacing, device=device, mode=mode\n            )\n            gc.collect()\n            logger.debug(\n                \"[TOMO PRODUCER %s] Spacing tomo_idx=%d tomo_id=%s shape %s -> %s\",\n                worker_id,\n                tomo_idx,\n                tomo_id,\n                tomo.shape,\n                spaced_tomo.shape,\n            )\n            proxy_tomos.append(\n                ShmTensor.from_tensor(\n                    spaced_tomo, f\"{tomo_idx}_{tomo_id}_{spacing_idx}\"\n                )\n            )\n        del tomo\n        gc.collect()\n        end = time.time()\n        logger.info(\n            \"[TOMO PRODUCER %s] Done tomo_idx=%d tomo_id=%s spacings %s -> %s with shape %s -> %s take %.2f sec\",\n            worker_id,\n            tomo_idx,\n            tomo_id,\n            ori_spacing,\n            target_spacings,\n            ori_shape,\n            tuple(e.shape for e in proxy_tomos),\n            end - start,\n        )\n        out_queue.put((tomo_idx, tomo_id, ori_spacing, target_spacings, proxy_tomos))\n\n\n@exception_printer\ndef patches_producer_worker_func(\n    in_queue: mp.Queue,\n    out_queue: queue.Queue,\n    num_remain_workers: mp.Value,\n    lock: threading.Lock,\n    patch_size,\n    border,\n    overlap,\n    worker_id=\"\",\n):\n    \"\"\"\n    Should be lightweight, and expected to be run in thread level.\n\n    Args:\n        in_queue: Input queue contain item of (tomo_idx, tomo_id, ori_spacing, target_spacing, proxy_tomo)\n        out_queue: Each item in output queue contains multiple patches belong to a same tomo\n    \"\"\"\n    logger.info(\"[PATCHES PRODUCER %s] STARTING..\", worker_id)\n    while True:\n        item = in_queue.get()\n        if item is None:\n            with lock:\n                num_remain_workers.value -= 1\n                if num_remain_workers.value == 0:\n                    # put final None so out_queue's consumer should stop too\n                    out_queue.put(None)\n            # put None back to in_queue so other in_queue's consumer workers will stop too\n            in_queue.put(None)\n            logger.info(\n                \"\\n\\n[PATCHES PRODUCER %s] STOPPING..\\n\\n\",\n                worker_id,\n            )\n            return\n        tomo_idx, tomo_id, ori_spacing, target_spacings, proxy_tomos = item\n        start = time.time()\n        assert len(proxy_tomos) == len(target_spacings) == len(TARGET_SPACINGS)\n\n        items = []\n        total_patches = 0\n        cur_patch_idx = -1\n        shm_tomos = []\n        for target_spacing, proxy_tomo in zip(target_spacings, proxy_tomos):\n            shm_tomo, tomo = proxy_tomo.to_tensor()\n            tomo_shape = tomo.shape\n            shm_tomos.append(shm_tomo)\n            patch_positions = get_sliding_patch_positions(\n                img_size=tomo.shape,\n                patch_size=patch_size,\n                border=border,\n                overlap=overlap,\n                validate=False,\n            )\n            # patch_positions = torch.from_numpy(patch_positions)\n            patch_positions = patch_positions.tolist()\n            total_patches += len(patch_positions)\n\n            for patch_pos in patch_positions:\n                # crop tomo patch\n                crop = crop_tomo_patch(tomo, patch_pos)\n                crop = crop[None]  # 1ZYX, uint8\n                cur_patch_idx += 1\n                item = {\n                    \"tomo_idx\": torch.tensor(tomo_idx, dtype=torch.int16),\n                    \"patch_pos\": torch.tensor(patch_pos, dtype=torch.float32),\n                    \"tomo_shape\": torch.tensor(tomo_shape, dtype=torch.int16),\n                    \"ori_spacing\": torch.tensor(ori_spacing, dtype=torch.float32),\n                    \"target_spacing\": torch.tensor(target_spacing, dtype=torch.float32),\n                    \"cur_patch_idx\": torch.tensor(cur_patch_idx, dtype=torch.int16),\n                    \"image\": crop,\n                }\n                items.append(item)\n        assert cur_patch_idx == total_patches - 1\n        for item in items:\n            item[\"total_patches\"] = torch.tensor(total_patches, dtype=torch.int16)\n\n        end = time.time()\n        logger.debug(\n            \"[PATCHES PRODUCER %s] Done tomo_idx=%d tomo_id=%s take %.2f sec\",\n            worker_id,\n            tomo_idx,\n            tomo_id,\n            end - start,\n        )\n        out_queue.put((tomo_idx, shm_tomos, items))\n\n\n@exception_printer\ndef batch_producer_worker_func(\n    in_queue: queue.Queue,\n    out_queue: queue.Queue,\n    batch_size,\n    device=\"cpu\",\n):\n    \"\"\"Should be lightweight, and expected to be run in thread level.\"\"\"\n    count = 0\n    items = []\n    should_continue = True\n    shm_dict = {}\n    last_tomo_idx = -1\n    while should_continue:\n        item = in_queue.get()\n        if item is None:\n            # put None so in_queue's consumer should notice and stop too\n            in_queue.put(None)\n            should_continue = False\n        else:\n            tomo_idx, shm_tomos, new_items = item\n            # @TODO - remove this line\n            assert all([tomo_idx == e[\"tomo_idx\"].item() for e in new_items])\n            items.extend(new_items)\n            shm_dict[tomo_idx] = shm_tomos\n\n        while True:\n            if len(items) == 0 or (len(items) < batch_size and should_continue):\n                break\n            start = time.time()\n            cur_bs = min(batch_size, len(items))\n            batch = items[:cur_bs]\n            batch = default_collate(batch)\n            if device != \"cpu\":\n                batch = {k: v.to(device) for k, v in batch.items()}\n            assert batch[\"image\"].dtype == torch.uint8\n            end = time.time()\n            logger.debug(\n                \"[BATCH PRODUCER] new batch %d (tomo %d) take %.4f\",\n                count,\n                tomo_idx,\n                end - start,\n            )\n            out_queue.put(batch)\n            count += 1\n            last_tomo_idx = batch[\"tomo_idx\"][-1].item()\n            items = items[cur_bs:]\n        for shm_tomo_idx in list(shm_dict.keys()):\n            # un-collated batch should start with tomo_idx >= last_tomo_idx\n            # so we're safe to unlink previous one\n            # since default_collate() copy to a newly allocated tensor\n            if shm_tomo_idx < last_tomo_idx:\n                for shm in shm_dict[shm_tomo_idx]:\n                    shm.unlink()\n                del shm_dict[shm_tomo_idx]\n\n    for shms in shm_dict.values():\n        for shm in shms:\n            shm.unlink()\n    del shm_dict\n    gc.collect()\n    # put a final None to indicate `all done`, stop inference worker too..\n    out_queue.put(None)\n    logger.info(\"[BATCH PRODUCER] Stopped since tomo_idx=%d\", tomo_idx)\n    assert len(items) == 0\n\n\n@exception_printer\ndef inference_worker(worker_id, input_queue, chunk_idx, gpu_id, model_cfgs):\n    logger.info(\n        \"STARTING INFERENCE WORKER %d with device cuda:%d, current input queue size %d, %d models with configs:\\n%s\",\n        worker_id,\n        gpu_id,\n        input_queue.qsize(),\n        len(model_cfgs),\n        model_cfgs,\n    )\n    device = torch.device(f\"cuda:{gpu_id}\")\n    patch_size = PATCH_SIZE\n    border = PATCH_BORDER\n    overlap = PATCH_OVERLAP\n\n    ############ START ALL WORKERS ###########\n    ##########################################\n    tomo_queue = mp.Queue(maxsize=TOMO_PRODUCER_PREFETCH)\n    tomo_producer_num_remain_workers = mp.Value(\"i\", TOMO_PRODUCER_NUM_WORKERS)\n    tomo_producer_lock = mp.Lock()\n    patches_queue = queue.Queue(maxsize=PATCHES_PRODUCER_PREFETCH)\n    patches_producer_num_remain_workers = mp.Value(\"i\", PATCHES_PRODUCER_NUM_WORKERS)\n    patches_producer_lock = threading.Lock()\n    batch_queue = queue.Queue(maxsize=DATALOADER_PREFETCH)\n\n    # PRODUCE SPACED TOMO\n    tomo_producer_workers = []\n    for _worker_id in range(TOMO_PRODUCER_NUM_WORKERS):\n        worker = mp.Process(\n            group=None,\n            target=tomo_producer_worker_func,\n            name=f\"tomo_producer_{worker_id}.{_worker_id}\",\n            args=(\n                input_queue,\n                tomo_queue,\n                tomo_producer_num_remain_workers,\n                tomo_producer_lock,\n            ),\n            kwargs={\n                \"device\": TOMO_SPACING_DEVICE,\n                \"mode\": TOMO_SPACING_MODE,\n                \"worker_id\": f\"{worker_id}.{_worker_id}\",\n            },\n        )\n        worker.start()\n        tomo_producer_workers.append(worker)\n\n    # PRODUCER TOMO PATCHES\n    patches_producer_workers = []\n    for _worker_id in range(PATCHES_PRODUCER_NUM_WORKERS):\n        worker = threading.Thread(\n            group=None,\n            target=patches_producer_worker_func,\n            name=f\"patches_producer_{worker_id}.{_worker_id}\",\n            args=(\n                tomo_queue,\n                patches_queue,\n                patches_producer_num_remain_workers,\n                patches_producer_lock,\n                patch_size,\n                border,\n                overlap,\n            ),\n            kwargs={\"worker_id\": f\"{worker_id}.{_worker_id}\"},\n        )\n        worker.start()\n        patches_producer_workers.append(worker)\n\n    # BATCH PRODUCER\n    batch_producer_worker = threading.Thread(\n        group=None,\n        target=batch_producer_worker_func,\n        name=\"batch_producer\",\n        args=(patches_queue, batch_queue, BATCH_SIZE),\n        kwargs={\"device\": \"cuda:0\"},\n    )\n    batch_producer_worker.start()\n\n    ############## LOAD MODEL ##################\n    ############################################\n    logger.info(\"LOADING MODEL..\")\n    models = []\n    unique_tta_names = set()\n    unique_agg_names = set()\n    sessions = []\n    for session_idx, model_cfg in enumerate(model_cfgs):\n        model_tta_names = model_cfg[\"tta\"]\n        model_agg_weights = model_cfg[\"weight\"]\n        model_agg_names = model_cfg[\"agg_name\"]\n        assert len(model_tta_names) == len(model_agg_weights) == len(model_agg_names)\n        assert len(model_tta_names) == len(set(model_tta_names))\n        model = Model3dWrapper(\n            model_cfg[\"sig\"],\n            gpu_id,\n            act=\"identity\" if HEATMAP_AGG_LOGITS else \"sigmoid\",\n        )\n        model.eval().to(device)\n        models.append(model)\n        unique_tta_names.update(model_tta_names)\n        unique_agg_names.update(model_agg_names)\n        logger.info(\n            \"Loaded model %d with TTA=%s, WEIGHTS=%s\",\n            session_idx,\n            model_tta_names,\n            model_agg_weights,\n        )\n        for tta_name, agg_weight, agg_name in zip(\n            model_tta_names, model_agg_weights, model_agg_names\n        ):\n            assert tta_name in ALL_SUPPORTED_TTAS\n            sessions.append((model, tta_name, agg_weight, agg_name))\n    unique_tta_names = list(unique_tta_names)\n    unique_agg_names = list(unique_agg_names)\n    num_sessions = len(sessions)\n    _session_agg_names = [e[-1] for e in sessions]\n    agg_name_to_first_session_idx = {\n        agg_name: _session_agg_names.index(agg_name) for agg_name in unique_agg_names\n    }\n    agg_name_to_last_session_idx = {\n        agg_name: num_sessions - 1 - _session_agg_names[::-1].index(agg_name)\n        for agg_name in unique_agg_names\n    }\n    logger.info(\n        \"FIRST/LAST SESSION IDX:\\n%s\\n%s\",\n        agg_name_to_first_session_idx,\n        agg_name_to_last_session_idx,\n    )\n\n    logger.info(\"-----------\\nLOADED %d MODELS:\\n%s------------\", len(models), models)\n    logger.info(\n        \"TOTALLY ENABLE %d MODELS, %d TTA %s, %d AGG_NAMES %s => %d SESSIONS\",\n        len(models),\n        len(unique_tta_names),\n        unique_tta_names,\n        len(unique_agg_names),\n        unique_agg_names,\n        len(sessions),\n    )\n    unique_tta_tfs = {\n        tta_name: build_tta(tta_name, 1.0, ori_dims=\"zyx\")\n        for tta_name in unique_tta_names\n    }\n\n    # Postprocess: heatmap blurring\n    if DECODE_NMS_BLUR_SIGMA is not None:\n        from yagm.transforms import monai_custom as CT\n\n        sigma_voxel = [\n            1000 / e / HEATMAP_STRIDE * DECODE_NMS_BLUR_SIGMA\n            for e in TARGET_SPACINGS[0]\n        ]\n        blur_operator = (\n            CT.CustomGaussianFilter(\n                spatial_dims=3,\n                sigma=sigma_voxel,\n                truncated=4,\n                approx=\"erf\",\n                requires_grad=False,\n            )\n            .eval()\n            .to(device)\n        )\n    else:\n        blur_operator = None\n    logger.info(\n        \"POST-PROCESSING BLUR OPERATOR:\\n%s on device %s\",\n        blur_operator,\n        getattr(blur_operator, \"device\", None),\n    )\n\n    cache_results = {\n        agg_name: {\n            \"cur_pred_heatmap_sum\": None,\n            \"cur_pred_heatmap_count\": None,\n            \"cur_tomo_idx\": None,\n            \"cur_tomo_id\": None,\n            \"main_target_spacing\": None,\n            \"submission\": SubmissionDataFrame(),\n        }\n        for agg_name in unique_agg_names\n    }\n    logger.info(\"\\n\\n\\nWORKER %d STARTING INFERENCE..\\n\", worker_id)\n    infer_start = time.time()\n    cur_batch_idx = 0\n    pbar = tqdm(desc=f\"===== INFERENCE {worker_id} =====\")\n    with torch.inference_mode(), torch.autocast(\n        device_type=\"cuda\", dtype=torch.float16\n    ):\n        while True:\n            batch = batch_queue.get()\n            if batch is None:\n                break\n            # print(\"GOT BATCH:\", {(k, v.shape, v.dtype) for k, v in batch.items()})\n            batch_image = batch[\"image\"].to(device)\n            assert batch_image.dtype == torch.uint8\n            B = batch_image.shape[0]\n\n            unique_tta_batch_images = {\n                tta_name: tta_tf.transform(batch_image)[0].float().contiguous()\n                for tta_name, tta_tf in unique_tta_tfs.items()\n            }\n\n            batch_outputs = []\n            for model, tta_name, agg_weight, agg_name in sessions:\n                tta_tf = unique_tta_tfs[tta_name]\n                tta_batch_image = unique_tta_batch_images[tta_name]\n                tta_batch_heatmap = model(tta_batch_image)\n                batch_heatmap = tta_tf.invert(tta_batch_heatmap)[0]\n                batch_outputs.append((batch_heatmap, agg_weight, agg_name))\n\n            assert len(batch_outputs) == num_sessions\n\n            for i in range(B):\n                for session_idx, (batch_heatmap, agg_weight, agg_name) in enumerate(\n                    batch_outputs\n                ):\n                    pred_patch_heatmap = batch_heatmap[i]\n                    C = pred_patch_heatmap.shape[0]\n                    assert C == 1\n\n                    tomo_idx = batch[\"tomo_idx\"][i]\n                    tomo_id = ALL_TOMO_IDS[tomo_idx]\n                    tomo_shape = batch[\"tomo_shape\"][i].tolist()\n                    ori_spacing = batch[\"ori_spacing\"][i].tolist()\n                    target_spacing = batch[\"target_spacing\"][i].tolist()\n                    cur_patch_idx = int(batch[\"cur_patch_idx\"][i].item())\n                    total_patches = int(batch[\"total_patches\"][i].item())\n                    patch_position = batch[\"patch_pos\"][i].cpu()\n\n                    is_first = (cur_patch_idx == 0) and (\n                        session_idx == agg_name_to_first_session_idx[agg_name]\n                    )\n                    is_last = (cur_patch_idx == total_patches - 1) and (\n                        session_idx == agg_name_to_last_session_idx[agg_name]\n                    )\n\n                    # load cache for this tomo+agg combination\n                    agg_result = cache_results[agg_name]\n                    ### start IS_FIRST\n                    if is_first:\n                        logger.info(\n                            \"[INFER %d] Model %s receive first patch of tomo %s (index %d)\",\n                            worker_id,\n                            session_idx,\n                            tomo_id,\n                            tomo_idx,\n                        )\n                        assert all(\n                            [\n                                agg_result[k] is None\n                                for k in [\n                                    \"cur_tomo_idx\",\n                                    \"cur_tomo_id\",\n                                    \"cur_pred_heatmap_sum\",\n                                    \"cur_pred_heatmap_count\",\n                                    \"main_target_spacing\",\n                                ]\n                            ]\n                        )\n\n                        # allocate new heatmap_sum and heatmap_count\n                        agg_result[\"cur_tomo_idx\"] = tomo_idx\n                        agg_result[\"cur_tomo_id\"] = tomo_id\n                        shared_agg_heatmap_shape = [\n                            round(e / HEATMAP_STRIDE) for e in tomo_shape\n                        ]\n                        # float16 to save memory\n                        agg_result[\"cur_pred_heatmap_sum\"] = torch.zeros(\n                            (C, *shared_agg_heatmap_shape),\n                            dtype=torch.float16,\n                            device=pred_patch_heatmap.device,\n                        )\n                        agg_result[\"cur_pred_heatmap_count\"] = torch.zeros(\n                            (C, *shared_agg_heatmap_shape),\n                            dtype=torch.float16,\n                            device=pred_patch_heatmap.device,\n                        )\n                        agg_result[\"main_target_spacing\"] = target_spacing\n                        assert all(\n                            [\n                                abs(a - b) < 1e-4\n                                for a, b in zip(target_spacing, TARGET_SPACINGS[0])\n                            ]\n                        )\n                    ### end IS_FIRST\n\n                    assert tomo_idx == agg_result[\"cur_tomo_idx\"]\n                    assert all(\n                        [\n                            a <= b\n                            for a, b in zip(\n                                agg_result[\"main_target_spacing\"], target_spacing\n                            )\n                        ]\n                    )\n                    cur_pred_heatmap_sum = agg_result[\"cur_pred_heatmap_sum\"]\n                    cur_pred_heatmap_count = agg_result[\"cur_pred_heatmap_count\"]\n                    main_target_spacing = agg_result[\"main_target_spacing\"]\n                    cur_submission: SubmissionDataFrame = agg_result[\"submission\"]\n                    shared_agg_heatmap_shape = cur_pred_heatmap_sum.shape[1:]\n\n                    scale_zyx = [\n                        a / b for a, b in zip(shared_agg_heatmap_shape, tomo_shape)\n                    ]\n                    scale_z, scale_y, scale_x = scale_zyx\n\n                    # interpolate to shared coordinate space\n                    C, Z2, Y2, X2 = pred_patch_heatmap.shape\n                    assert C == 1\n                    patch_heatmap_shape = [\n                        round(e * s) for e, s in zip(patch_size, scale_zyx)\n                    ]\n                    if [Z2, Y2, X2] != patch_heatmap_shape:\n                        logger.debug(\n                            \"PATCH HEATMAP INTERPOLATE: %s --> %s\",\n                            (Z2, Y2, X2),\n                            patch_heatmap_shape,\n                        )\n                        pred_patch_heatmap = F.interpolate(\n                            pred_patch_heatmap[None],  # (1, C, Z2, Y2, X2)\n                            size=patch_heatmap_shape,\n                            mode=HEATMAP_INTERPOLATION_MODE,\n                            align_corners=False,\n                        )[\n                            0\n                        ]  # (C, Z3, Y3, X3)\n                    else:\n                        logger.debug(\"PATCH HEATMAP SAME: %s\", pred_patch_heatmap.shape)\n                        pass\n\n                    roi_start, roi_end, patch_start, patch_end = patch_position.tolist()\n                    top_pad_z, top_pad_y, top_pad_x = [\n                        max(-start, 0) for start in roi_start\n                    ]\n                    bot_pad_z, bot_pad_y, bot_pad_x = [\n                        max(0, end - size) for end, size in zip(roi_end, tomo_shape)\n                    ]\n                    assert (\n                        tuple(\n                            round((pe - ps) * s)\n                            for pe, ps, s in zip(patch_end, patch_start, scale_zyx)\n                        )\n                        == pred_patch_heatmap.shape[1:]\n                    )\n                    sz, sy, sx = tuple(\n                        rs - ps for rs, ps in zip(roi_start, patch_start)\n                    )\n                    ez, ey, ex = tuple(\n                        ps - ce + pe\n                        for ps, ce, pe in zip(patch_size, patch_end, roi_end)\n                    )\n\n                    assert (roi_end[0] - bot_pad_z) - (roi_start[0] + top_pad_z) == (\n                        ez - bot_pad_z\n                    ) - (sz + top_pad_z)\n                    assert (roi_end[1] - bot_pad_y) - (roi_start[1] + top_pad_y) == (\n                        ey - bot_pad_y\n                    ) - (sy + top_pad_y)\n                    assert (roi_end[2] - bot_pad_x) - (roi_start[2] + top_pad_x) == (\n                        ex - bot_pad_x\n                    ) - (sx + top_pad_x)\n\n                    dst_shape = shared_agg_heatmap_shape\n                    dst_slices = [\n                        slice(None),\n                        slice(\n                            *scale_range(\n                                roi_start[0] + top_pad_z,\n                                roi_end[0] - bot_pad_z,\n                                scale_z,\n                                0,\n                                dst_shape[0],\n                            )\n                        ),\n                        slice(\n                            *scale_range(\n                                roi_start[1] + top_pad_y,\n                                roi_end[1] - bot_pad_y,\n                                scale_y,\n                                0,\n                                dst_shape[1],\n                            )\n                        ),\n                        slice(\n                            *scale_range(\n                                roi_start[2] + top_pad_x,\n                                roi_end[2] - bot_pad_x,\n                                scale_x,\n                                0,\n                                dst_shape[2],\n                            )\n                        ),\n                    ]\n                    src_shape = pred_patch_heatmap.shape[1:]\n                    src_slices = [\n                        slice(None),\n                        slice(\n                            *scale_range(\n                                sz + top_pad_z, ez - bot_pad_z, scale_z, 0, src_shape[0]\n                            )\n                        ),\n                        slice(\n                            *scale_range(\n                                sy + top_pad_y, ey - bot_pad_y, scale_y, 0, src_shape[1]\n                            )\n                        ),\n                        slice(\n                            *scale_range(\n                                sx + top_pad_x, ex - bot_pad_x, scale_x, 0, src_shape[2]\n                            )\n                        ),\n                    ]\n\n                    if HEATMAP_AGG_MODE == \"avg\":\n                        cur_pred_heatmap_sum[dst_slices] += (\n                            pred_patch_heatmap[src_slices] * agg_weight\n                        )\n                        cur_pred_heatmap_count[dst_slices] += agg_weight\n                    elif HEATMAP_AGG_MODE == \"max\":\n                        torch.maximum(\n                            cur_pred_heatmap_sum[dst_slices],\n                            pred_patch_heatmap[src_slices],\n                            out=cur_pred_heatmap_sum[dst_slices],\n                        )\n                    else:\n                        raise ValueError\n\n                    if is_last:\n                        # DECODE HEATMAP TO COORDINATE\n                        tomo_id = agg_result[\"cur_tomo_id\"]\n                        logger.info(\n                            \"[INFER %d] Receive last patch of tomo %s (index %d), decoding..\",\n                            worker_id,\n                            tomo_id,\n                            tomo_idx,\n                        )\n\n                        if HEATMAP_AGG_MODE == \"avg\":\n                            # ensure prediction cover all the tomogram\n                            # assert torch.all(cur_pred_heatmap_count > 0)\n                            heatmap = torch.div(\n                                cur_pred_heatmap_sum,\n                                cur_pred_heatmap_count,\n                                out=cur_pred_heatmap_sum,\n                            )\n                        elif HEATMAP_AGG_MODE == \"max\":\n                            heatmap = cur_pred_heatmap_sum\n                        else:\n                            raise ValueError\n                        if HEATMAP_AGG_LOGITS:\n                            heatmap = F.sigmoid(heatmap)\n\n                        if SAVE_HEATMAP_NPY:\n                            heatmap_npy = heatmap.half().cpu().numpy()\n                            assert heatmap_npy.dtype == np.float16\n                            save_heatmap_path = os.path.join(\n                                SAVE_HEATMAP_NPY_DIR, agg_name, f\"{tomo_id}.npy\"\n                            )\n                            os.makedirs(\n                                os.path.dirname(save_heatmap_path), exist_ok=True\n                            )\n                            np.save(save_heatmap_path, heatmap_npy)\n\n                        if DECODE_HEATMAP_TO_CSV:\n                            # Decode heatmap\n                            _start = time.time()\n                            if DECODE_METHOD == \"nms\":\n                                outputs = decode_heatmap(\n                                    heatmap,\n                                    main_target_spacing,\n                                    HEATMAP_STRIDE,\n                                    blur_operator=blur_operator,\n                                )\n                            elif DECODE_METHOD == \"cc3d\":\n                                outputs = decode_segment_mask(\n                                    heatmap,\n                                    main_target_spacing,\n                                    HEATMAP_STRIDE,\n                                    prob_thres=DECODE_CC3D_CONF_THRES,\n                                    radius_factor_thres=DECODE_CC3D_RADIUS_FACTOR,\n                                    conf_mode=DECODE_CC3D_CONF_MODE,\n                                    prob_mode=DECODE_CC3D_PROB_MODE,\n                                )\n                            else:\n                                raise ValueError\n                            _end = time.time()\n                            logger.debug(\n                                \"Decode heatmap of shape %s take %.4f sec\",\n                                heatmap.shape,\n                                _end - _start,\n                            )\n                            # Convert to submission coordinates\n                            assert len(outputs) == 1\n                            main_target_spacing_tensor = torch.tensor(\n                                [main_target_spacing], dtype=torch.float32\n                            )  # (1, 3)\n                            ori_spacing_tensor = torch.tensor(\n                                [ori_spacing], dtype=torch.float32\n                            )  # (1, 3)\n                            for _channel_idx, keypoints in enumerate(outputs):\n                                # back to original coordinate space\n                                keypoints[:, :3] = (keypoints[:, :3] + 0.5) * (\n                                    HEATMAP_STRIDE\n                                    * main_target_spacing_tensor\n                                    / ori_spacing_tensor\n                                ) - 0.5\n                                keypoints = keypoints.tolist()\n\n                                ###### VISUALIZATION ######\n                                if VIZ_ENABLE and MODE != \"KAGGLE\":\n                                    tomo_dir = os.path.join(DATA_DIR, tomo_id)\n                                    row = GT_DF[GT_DF[\"tomo_id\"] == tomo_id].iloc[0]\n                                    ori_shape = (row[\"Z\"], row[\"Y\"], row[\"X\"])\n                                    gt_zyx = [\n                                        row[\"motor_z\"],\n                                        row[\"motor_y\"],\n                                        row[\"motor_x\"],\n                                    ]\n                                    viz = viz_byu(\n                                        tomo_dir,\n                                        ori_spacing,\n                                        ori_shape,\n                                        gt_zyx,\n                                        heatmap[0],\n                                        keypoints,\n                                    )\n                                    # get the kind/judge of current prediction, one of TP, TN, FP, FN\n                                    if (\n                                        len(keypoints) > 0\n                                        and keypoints[0][3] >= DECODE_CONF_THRES\n                                    ):\n                                        # predicted as positive\n                                        pred_zyx = keypoints[0][:3]\n                                        if gt_zyx[0] == -1:\n                                            # GT is negative\n                                            judge = \"FP\"\n                                        else:\n                                            dist = (\n                                                (gt_zyx[0] - pred_zyx[0]) ** 2\n                                                + (gt_zyx[1] - pred_zyx[1]) ** 2\n                                                + (gt_zyx[2] - pred_zyx[2]) ** 2\n                                            ) ** 0.5\n                                            if dist <= 1000 / ori_spacing:\n                                                judge = \"TP\"\n                                            else:\n                                                judge = \"FN\"\n                                    else:\n                                        # predicted as negative\n                                        if gt_zyx[0] == -1:\n                                            # GT is negative\n                                            judge = \"TN\"\n                                        else:\n                                            judge = \"FN\"\n\n                                    save_name = f\"{judge}-{tomo_idx}-{tomo_id}-{'x'.join([str(int(e)) for e in ori_shape])}-spacing{ori_spacing}.jpg\"\n                                    save_path = os.path.join(\n                                        TMP_VIZ_DIR, judge, agg_name, save_name\n                                    )\n                                    os.makedirs(\n                                        os.path.dirname(save_path), exist_ok=True\n                                    )\n                                    cv2.imwrite(\n                                        save_path, cv2.cvtColor(viz, cv2.COLOR_RGB2BGR)\n                                    )\n\n                                if len(keypoints) > 0:\n                                    for kpt in keypoints:\n                                        z, y, x, conf = kpt\n                                        cur_submission.add_row(tomo_id, x, y, z, conf)\n                                else:\n                                    # no motor detected\n                                    cur_submission.add_row(tomo_id, -1, -1, -1, 0.0)\n\n                        # clear the cache\n                        for k in [\n                            \"cur_tomo_idx\",\n                            \"cur_tomo_id\",\n                            \"cur_pred_heatmap_sum\",\n                            \"cur_pred_heatmap_count\",\n                            \"main_target_spacing\",\n                        ]:\n                            agg_result[k] = None\n                        gc.collect()\n                        # torch.cuda.empty_cache()\n            cur_batch_idx += 1\n            pbar.update(1)\n    pbar.close()\n    infer_end = time.time()\n\n    logger.info(\n        \"\\n\\n\\n INFER WORKER %d INFER ON %d BATCHES TAKE %.4f sec\",\n        worker_id,\n        cur_batch_idx,\n        infer_end - infer_start,\n    )\n\n    assert len(cache_results) == len(unique_agg_names)\n    if DECODE_HEATMAP_TO_CSV:\n        for agg_name, agg_result in cache_results.items():\n            submision = agg_result[\"submission\"]\n            df = submision.to_pandas(submit=False)\n            csv_path = os.path.join(TMP_CSV_DIR, f\"{agg_name}_chunk{chunk_idx}.csv\")\n            df.to_csv(csv_path, index=False)\n\n    # JOIN\n    logger.info(\"Worker %d waiting for all worker to be finished..\", worker_id)\n    all_workers = [\n        batch_producer_worker,\n        *tomo_producer_workers,\n        *patches_producer_workers,\n    ]\n    for worker in all_workers:\n        worker.join()\n    return None\n\n\n########################### MAIN CODE #############################\n###################################################################\nlogger.warning(\"CLEARING ALL CONTENTS IN: %s\", [TMP_CSV_DIR, TMP_CSV_DIR, TMP_VIZ_DIR])\nclear_dir_content(\"/dev/shm/\")\n# clear_dir_content(TMP_CSV_DIR)\nclear_dir_content(TMP_VIZ_DIR)\nif MODE == \"LOCAL\":\n    clear_dir_content(WORKING_DIR)\nos.makedirs(TMP_CSV_DIR, exist_ok=True)\nos.makedirs(TMP_VIZ_DIR, exist_ok=True)\n\n\nglobal_start = time.time()\n\nall_unique_agg_names = set()\nfor job_idx, model_cfgs in enumerate(JOBS):\n    # fill the input queue\n    input_queue = mp.Queue()\n    for tomo_idx, (tomo_id, ori_spacing) in enumerate(\n        zip(ALL_TOMO_IDS, ALL_TOMO_SPACINGS)\n    ):\n        input_queue.put((tomo_idx, tomo_id, [ori_spacing] * 3))\n    # None to indicate ending and workers should stop when receive this None\n    input_queue.put(None)\n\n    job_unique_agg_names = set()\n    for model_cfg in model_cfgs:\n        job_unique_agg_names.update(model_cfg[\"agg_name\"])\n    assert len(all_unique_agg_names.intersection(job_unique_agg_names)) == 0\n    all_unique_agg_names.update(job_unique_agg_names)\n\n    job_workers = []\n    for chunk_idx, device_id in enumerate(DEVICES):\n        # worker_id, input_queue, chunk_idx, gpu_id, model_cfgs\n        worker = mp.Process(\n            group=None,\n            target=inference_worker,\n            name=f\"INFERENCE_WORKER_job={job_idx}_chunk={chunk_idx}\",\n            kwargs={\n                \"worker_id\": chunk_idx,\n                \"input_queue\": input_queue,\n                \"chunk_idx\": chunk_idx,\n                \"gpu_id\": device_id,\n                \"model_cfgs\": model_cfgs,\n            },\n        )\n        worker.start()\n        job_workers.append(worker)\n\n    for worker in job_workers:\n        worker.join()\n\nglobal_end = time.time()\nlogger.info(\n    \"\\n\\n\\n>>>>>> FINISH ALL INFERENCE TASK WITHIN %.4f sec (%.4f min)\",\n    global_end - global_start,\n    (global_end - global_start) / 60.0,\n)\nall_unique_agg_names = list(all_unique_agg_names)\nprint('ALL UNIQUE AGG NAMES:', all_unique_agg_names)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T00:27:36.535463Z","iopub.execute_input":"2025-06-04T00:27:36.535898Z","iopub.status.idle":"2025-06-04T00:27:36.551555Z","shell.execute_reply.started":"2025-06-04T00:27:36.535858Z","shell.execute_reply":"2025-06-04T00:27:36.550629Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python3 /kaggle/working/byu/submit_3d.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T00:27:37.783477Z","iopub.execute_input":"2025-06-04T00:27:37.783754Z","iopub.status.idle":"2025-06-04T00:30:51.234661Z","shell.execute_reply.started":"2025-06-04T00:27:37.783733Z","shell.execute_reply":"2025-06-04T00:30:51.233827Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2D CODE","metadata":{}},{"cell_type":"code","source":"%%writefile /kaggle/working/byu/submit_2d.py\n\nimport sys\nsys.path.extend([\n    '/kaggle/working/byu/yagm/src',\n    '/kaggle/working/byu/src',\n    '/kaggle/working/byu/',\n    '/kaggle/working/byu/third_party/segmentation_models_pytorch_3d',\n    '/kaggle/working/byu/third_party/slowfast',\n    '/kaggle/working/byu/third_party/timm_3d'\n])\nprint('SYS PATH:', sys.path, sep = '\\n')\n\n\nimport warnings\n\nwarnings.simplefilter(\"ignore\")\nimport gc\nimport json\nimport logging\nimport multiprocessing as mp\nimport os\nimport queue\nimport shutil\nimport sys\nimport threading\nimport time\nfrom functools import partial\n\nimport cc3d\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom omegaconf import OmegaConf\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.utils.data import default_collate\nfrom tqdm import tqdm\nfrom yagm.tasks.base_task import BaseTask\nfrom yagm.transforms.keypoints.decode import decode_heatmap_3d, decode_segment_mask_3d\nfrom yagm.transforms.sliding_window import get_sliding_patch_positions\nfrom yagm.transforms.tta_3d import build_tta\nfrom yagm.utils import hydra as hydra_utils\nfrom yagm.utils import lightning as l_utils\nfrom yagm.utils.concurrent import ShmTensor\n\n# Setting up custom logging\nfrom yagm.utils.logging import init_logging, setup_logging\n\nfrom byu.data.io import MultithreadOpencvTomogramLoader\nfrom byu.inference import zoo\nfrom byu.inference.tensorrt_engine import ThreadsafeTRTEngine\nfrom byu.utils.data import SubmissionDataFrame\nfrom byu.utils.metrics import compute_metrics, kaggle_score\nfrom byu.utils.viz import viz_byu\nfrom byu.utils.wbf import weighted_boxes_fusion_3d\n\ntry:\n    hydra_utils.init_hydra()\nexcept:\n    print(\"SKIP RE-INIT HYDRA\")\n\n### SETUP LOGGER FOR IPYTHON ###\n# ref: https://github.com/ipython/ipykernel/issues/111\n# Create logger\nlogger = logging.getLogger()\nlogger.setLevel(logging.INFO)\n# Create STDERR handler\nhandler = logging.StreamHandler(sys.stderr)\n# Create formatter and add it to the handler\nformatter = logging.Formatter(\"[%(levelname)s] %(message)s\")\nhandler.setFormatter(formatter)\n# Set STDERR handler as the only handler\nlogger.handlers = [handler]\n\n\n##################################################\n###### MODEL SIGNATURES OVERRIDES\n##################################################\nALL_SUPPORTED_TTAS = [\"yx\", \"yx_x\", \"yx_y\", \"yx_xy\", \"xy\", \"xy_x\", \"xy_y\", \"xy_xy\"]\n\n_SIG_TEMPLATE = {\n    \"backend\": \"torch\",  # torch | trt\n    \"trt_path\": \"\",\n    \"config_path\": \"\",\n    \"torch_path\": \"\",\n    \"ema\": 0.99,\n}\n\n\n# MAXVIT, LB 85.4\n# https://www.kaggle.com/code/dangnh0611/submit-heatmap-final?scriptVersionId=243128311\nEXP27_2D_MAXVIT_TINY_ALLGTV3 = {\n    \"backend\": \"torch\",  # torch | trt\n    \"trt_path\": \"\",\n    \"config_path\": \"/kaggle/input/byu-final-dataset/EXP27_2D_MAXVIT_TINY_ALLGTV3_config.yaml\",\n    \"torch_path\": \"/kaggle/input/byu-final-dataset/EXP27_2D_MAXVIT_TINY_ALLGTV3_ep5_step20000.ckpt\",\n    \"ema\": 0.99,\n}\n\n\n# COAT, LB 84.8\n# https://www.kaggle.com/code/dangnh0611/submit-heatmap-final?scriptVersionId=242969781\nEXP19_2D_COATLITEMEDIUM_ALLGTV2 = {\n    \"backend\": \"torch\",  # torch | trt\n    \"trt_path\": \"\",\n    \"config_path\": \"/kaggle/input/byu-final-dataset/EXP19_2D_COATLITEMEDIUM_ALLGTV2_config.yaml\",\n    \"torch_path\": \"/kaggle/input/byu-final-dataset/EXP19_2D_COATLITEMEDIUM_ALLGTV2_ep2_step10000_val_Fbeta0.949087_val_PAP0.935705.ckpt\",\n    \"ema\": 0.99,\n}\n\n\nJOBS = [\n    [\n        {\n            \"sig\": EXP27_2D_MAXVIT_TINY_ALLGTV3,\n            \"tta\": [\"yx\", \"xy\"],\n            \"weight\": [1.0, 1.0],\n            \"agg_name\": [\"MAXVIT-yx\", \"MAXVIT-xy\"],\n        }\n    ]\n]\n\n\n##################################################\n###### GLOBAL VARS\n##################################################\n_CV2_TOMO_LOADER_NUM_THREADS = 8\n\nPATCH_SIZE = [3, 896, 896]\nPATCH_BORDER = [0, 0, 0]\nPATCH_OVERLAP = [0, 0, 0]\n\nIO_BACKEND = \"cv2\"\nBASE_SPACING = 13.1  # avg 15.6\n# TARGET_SPACINGS = [13.1, 16.0, 19.7]\nTARGET_SPACINGS = [[32.0, 16.0, 16.0]]\nTOMO_SPACING_DEVICE = \"cpu\"\nTOMO_SPACING_MODE = \"trilinear\"\n\nHEATMAP_AGG_MODE = \"avg\"\nHEATMAP_AGG_LOGITS = True\nHEATMAP_INTERPOLATION_MODE = \"bilinear\"\nFORWARD_SAVE_MEMORY = False\nBATCH_SIZE = 8\nHEATMAP_STRIDE = [4, 8, 8]\n\n# MASK CC3D DECODING\nDECODE_CC3D_CONF_THRES = 0.05\nDECODE_CC3D_RADIUS_FACTOR = 0.1\nDECODE_CC3D_CONF_MODE = \"prob\"  # volume | prob | fuse\nDECODE_CC3D_PROB_MODE = \"center\"  # center | mean | max\n# HEATMAP NMS DECODING\nDECODE_NMS_BLUR_SIGMA = None\n###### MAIN DECODE PARAMS ######\nDECODE_METHOD = \"nms\"  # nms, cc3d\nDECODE_CONF_THRES = 0.01\nQUANTILE_THRES = 0.55\n\nSAVE_HEATMAP_NPY = False\nDECODE_HEATMAP_TO_CSV = True\nCSV_ENSEMBLE_MODE = \"wbf\"  # max | wbf\nWBF_CONF_TYPE = \"avg\"  # avg | max\nWBF_CONF_THRES = 0.2\n\nVIZ_ENABLE = False\n\n\n# automatic determine the MODE\nif os.path.isdir(\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025\"):\n    MODE = \"KAGGLE\"\nelse:\n    MODE = \"LOCAL\"\n# manual override\n# MODE = \"KAGGLE_VAL\"\n\n# LOCAL DEVELOPMENT ENV\nif MODE == \"LOCAL\":\n    DEVICES = [0]\n    ASSETS_DIR = \"assets/\"\n    WORKING_DIR = \"outputs/submit/working_2d/\"\n\n    DATA_DIR = \"/home/dangnh36/datasets/.comp/byu/raw/train\"\n    VAL_GT_CSV_PATH = \"/home/dangnh36/datasets/.comp/byu/processed/gt_v3.csv\"\n    with open(\n        \"/home/dangnh36/datasets/.comp/byu/processed/cv/v3/skf4_rd42.json\", \"r\"\n    ) as f:\n        cv_meta = json.load(f)\n        ALL_TOMO_IDS = cv_meta[\"folds\"][0][\"val\"][:]  # fold 0\n\n    # DATA_DIR = \"/home/dangnh36/datasets/.comp/byu/processed/pseudo_test/tomograms\"\n    # VAL_GT_CSV_PATH = \"/home/dangnh36/datasets/.comp/byu/processed/pseudo_test/gt.csv\"\n    # ALL_TOMO_IDS = sorted(os.listdir(DATA_DIR)) * 10\n\n    GT_DF = pd.read_csv(VAL_GT_CSV_PATH)\n    GT_DF = GT_DF[GT_DF[\"tomo_id\"].isin(ALL_TOMO_IDS)].reset_index(drop=True)\n    _tomo2spacing = {}\n    for i, row in GT_DF.iterrows():\n        _tomo2spacing[row[\"tomo_id\"]] = row[\"voxel_spacing\"]\n    ALL_TOMO_SPACINGS = [_tomo2spacing[tomo_id] for tomo_id in ALL_TOMO_IDS]\n\n    # TOMO_PRODUCER_NUM_WORKERS = 8\n    # TOMO_PRODUCER_PREFETCH = 8\n    # PATCHES_PRODUCER_NUM_WORKERS = 2\n    # PATCHES_PRODUCER_PREFETCH = 8\n    # DATALOADER_PREFETCH = 8\n\n    TOMO_PRODUCER_NUM_WORKERS = 1\n    TOMO_PRODUCER_PREFETCH = 1\n    PATCHES_PRODUCER_NUM_WORKERS = 1\n    PATCHES_PRODUCER_PREFETCH = 2\n    DATALOADER_PREFETCH = 8\n\n# VALIDATION ON KAGGLE\nelif MODE == \"KAGGLE_VAL\":\n    DEVICES = [0, 1]\n    ASSETS_DIR = \"/kaggle/input/byu-final-dataset/\"\n    DATA_DIR = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train\"\n    WORKING_DIR = \"/kaggle/working/\"\n    VAL_GT_CSV_PATH = \"/kaggle/input/byu-checkpoints/gt_v2.csv\"\n    with open(\"/kaggle/input/byu-checkpoints/skf4_rd42.json\", \"r\") as f:\n        cv_meta = json.load(f)\n        ALL_TOMO_IDS = cv_meta[\"folds\"][0][\"val\"][:50]  # fold 0, top 10 first tomo\n    GT_DF = pd.read_csv(VAL_GT_CSV_PATH)\n    GT_DF = GT_DF[GT_DF[\"tomo_id\"].isin(ALL_TOMO_IDS)].reset_index(drop=True)\n    _tomo2spacing = {}\n    for i, row in GT_DF.iterrows():\n        _tomo2spacing[row[\"tomo_id\"]] = row[\"voxel_spacing\"]\n    ALL_TOMO_SPACINGS = [_tomo2spacing[tomo_id] for tomo_id in ALL_TOMO_IDS]\n\n    TOMO_PRODUCER_NUM_WORKERS = 1\n    TOMO_PRODUCER_PREFETCH = 1\n    PATCHES_PRODUCER_NUM_WORKERS = 1\n    PATCHES_PRODUCER_PREFETCH = 2\n    DATALOADER_PREFETCH = 8\n\n# PRIVATE TEST SUBMISSION\nelif MODE == \"KAGGLE\":\n    DEVICES = [0, 1]\n    ASSETS_DIR = \"/kaggle/input/byu-final-dataset/\"\n    DATA_DIR = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test\"\n    WORKING_DIR = \"/kaggle/working/\"\n    ALL_TOMO_IDS = sorted(os.listdir(DATA_DIR))\n    ALL_TOMO_SPACINGS = [BASE_SPACING] * len(ALL_TOMO_IDS)\n\n    TOMO_PRODUCER_NUM_WORKERS = 1\n    TOMO_PRODUCER_PREFETCH = 1\n    PATCHES_PRODUCER_NUM_WORKERS = 1\n    PATCHES_PRODUCER_PREFETCH = 2\n    DATALOADER_PREFETCH = 8\n\nelse:\n    raise ValueError\n\nTMP_CSV_DIR = os.path.join(WORKING_DIR, \"tmp_csv\")\nTMP_VIZ_DIR = os.path.join(WORKING_DIR, \"tmp_viz\")\nSAVE_HEATMAP_NPY_DIR = os.path.join(WORKING_DIR, \"tmp_heatmap\")\n\n# RECHECK SOME CONDITIONS\nassert len(ALL_TOMO_IDS) == len(ALL_TOMO_SPACINGS)\n# assert not (SAVE_HEATMAP_NPY and DECODE_HEATMAP_TO_CSV) and (SAVE_HEATMAP_NPY or DECODE_HEATMAP_TO_CSV)\n# @TODO - refactor to support arbitary HEATMAP_STRIDE\n# assert HEATMAP_STRIDE[0] == 1, f\"Currently, heatmap stride along Z must be 1\"\n\n\n###################################################\n###### FUNCTIONS/CLASSES DEFINITION\n###################################################\ndef exception_printer(func):\n    def wrapper(*args, **kwargs):\n        try:\n            return func(*args, **kwargs)\n        except Exception as e:\n            logger.exception(\"Exception %s in %s: %s\", type(e), func.__name__, e)\n            import traceback\n\n            traceback.print_exc()\n            raise e\n\n    return wrapper\n\n\ndef clear_dir_content(folder_path):\n    if os.path.isdir(folder_path):\n        for item in os.listdir(folder_path):\n            item_path = os.path.join(folder_path, item)\n            if os.path.isdir(item_path):\n                shutil.rmtree(item_path)\n            else:\n                os.remove(item_path)\n\n\nclass Model2dWrapper(nn.Module):\n    def __init__(self, sig, gpu_id, act=\"sigmoid\"):\n        super().__init__()\n        if sig[\"backend\"] == \"trt\":\n            logger.info(\"Loading TRT model with config:\\n%s\", sig)\n            self.model = ThreadsafeTRTEngine(sig[\"trt_path\"], gpu_id)\n        elif sig[\"backend\"] == \"torch\":\n            cfg = OmegaConf.load(sig[\"config_path\"])\n            cfg.misc.log_model = False\n            cfg.ema.val_decays = [sig[\"ema\"] if sig[\"ema\"] is not None else 0]\n            cfg.ckpt.strict = False\n\n            ######### OVERWRITE SOME CONFIGS ##########\n            # OVERWRITE SOME CONFIG\n            try:\n                cfg.model.encoder.pretrained = False\n            except:\n                pass\n\n            # for backward-compatibility\n            OmegaConf.set_readonly(cfg, False)\n            # just keep outputing heatmap\n            cfg.model.reg_head.enable = False\n            cfg.model.dsnt = OmegaConf.create()\n            cfg.model.dsnt.enable = False\n            OmegaConf.set_readonly(cfg, True)\n            ##########\n\n            task: BaseTask = l_utils.build_task(cfg)\n            l_utils.load_lightning_state_dict(\n                model=task,\n                ckpt_path=sig[\"torch_path\"],\n                cfg=cfg,\n            )\n            logger.info(\"Loaded Pytorch state dict from %s\", sig[\"torch_path\"])\n            # self.model = TorchModelWrapper(task.model)\n            self.model = task.model\n            del task\n            gc.collect()\n        else:\n            raise ValueError\n\n        self.act_name = act\n        if act == \"sigmoid\":\n            self.act = F.sigmoid\n            self.nan_to_num = partial(torch.nan_to_num, nan=0.0)\n        elif act == \"identity\":\n            self.act = lambda x: x\n            self.nan_to_num = partial(\n                torch.nan_to_num, nan=torch.finfo(torch.float32).min\n            )\n        else:\n            raise ValueError\n\n    def forward(self, x):\n        heatmap = self.act(self.model(x)[2])\n        heatmap = self.nan_to_num(heatmap)\n        # if torch.isnan(heatmap).any():\n        #     print('\\n\\nNAN!!!')\n        #     raise Exception\n        return heatmap\n\n\ndef decode_heatmap(pred_heatmap, target_spacing, stride, blur_operator=None):\n    radius_voxel = [1000.0 / e / s for e, s in zip(target_spacing, stride)]\n    radius_thres = max(radius_voxel)\n    # @TODO - currently, using fixed pool_ksize=3 due to performance issue\n    # when using large kernel size with Pytorch (>10 sec on 224x448x448)\n\n    # # maximum ood number which <= radius_thres\n    # # 0.7071067811865475 = 1 / sqrt(2)\n    # pool_ksize = int((2 * radius_thres * 0.7071067811865475 - 1) // 2 * 2 + 1)\n    # # at least 3, [1,1,1] == no pooling, which increase maxRecall but significantly reduce other metrics\n    # pool_ksize = max(3, pool_ksize)\n    # pool_ksize = [pool_ksize, pool_ksize, pool_ksize]\n\n    pool_ksize = [3, 3, 3]\n    logger.debug(\n        \"Heatmap decode with radius=%s, pool_ksize=%s, radius_thres=%s\",\n        radius_voxel,\n        pool_ksize,\n        radius_thres,\n    )\n    del target_spacing\n    assert pred_heatmap.shape[0] == 1\n    ret = []\n    for channel_idx, heatmap in enumerate(pred_heatmap):\n        outputs = decode_heatmap_3d(\n            heatmap=heatmap,\n            pool_ksize=pool_ksize,\n            nms_radius_thres=radius_thres,\n            blur_operator=blur_operator,\n            conf_thres=DECODE_CONF_THRES,\n            # max_dets=5 if (VIZ_ENABLE and MODE != \"KAGGLE\") else 1,\n            max_dets=100,\n            timeout=None,\n        )\n        outputs = outputs.cpu()\n        ret.append(outputs)\n    return ret\n\n\ndef decode_segment_mask(\n    pred_heatmap,\n    target_spacing,\n    stride,\n    prob_thres,\n    radius_factor_thres,\n    conf_mode=\"fuse\",\n    prob_mode=\"avg\",\n):\n    radius_voxel_thres = [\n        1000 / e / s * radius_factor_thres for e, s in zip(target_spacing, stride)\n    ]\n    volume_thres = (\n        4\n        / 3\n        * np.pi\n        * radius_voxel_thres[0]\n        * radius_voxel_thres[1]\n        * radius_voxel_thres[2]\n    )\n    assert len(pred_heatmap.shape) == 4 and pred_heatmap.shape[0] == 1\n    ret = []\n    for channel_idx, heatmap in enumerate(pred_heatmap):\n        keypoints = decode_segment_mask_3d(\n            heatmap,\n            prob_thres,\n            volume_thres,\n            max_dets=100,\n            conf_mode=conf_mode,\n            prob_mode=prob_mode,\n        )\n        ret.append(keypoints)\n    return ret\n\n\ndef crop_tomo_patch(tomo, patch_position):\n    # start = time.time()\n    tomo_shape = tomo.shape\n    _roi_start, _roi_end, patch_start, patch_end = patch_position\n    top_pad_z, top_pad_y, top_pad_x = [max(-start, 0) for start in patch_start]\n    bot_pad_z, bot_pad_y, bot_pad_x = [\n        max(0, end - size) for end, size in zip(patch_end, tomo_shape)\n    ]\n    actual_crop_start = [max(0, start) for start in patch_start]\n    actual_crop_end = [min(end, size) for end, size in zip(patch_end, tomo_shape)]\n    crop_slices = tuple(\n        slice(start, end) for start, end in zip(actual_crop_start, actual_crop_end)\n    )\n    crop = tomo[crop_slices]\n    # pad if needed\n    pad = (top_pad_x, bot_pad_x, top_pad_y, bot_pad_y, top_pad_z, bot_pad_z)\n    if any(pad):\n        crop = F.pad(crop, pad, mode=\"constant\", value=0)\n    # end = time.time()\n    # logger.debug(\"Crop take %f sec: %s\", end - start, patch_position)\n    return crop\n\n\n@torch.inference_mode()\ndef spacing_torch(\n    ori_tomo, ori_spacing, target_spacing, device=\"cuda\", mode=\"trilinear\"\n):\n    if tuple(ori_spacing) == tuple(target_spacing):\n        return ori_tomo\n    # NOTE: https://docs.opencv.org/3.4/da/d54/group__imgproc__transform.html#ga47a974309e9102f5f08231edc7e7529d\n    # shrink -> AREA, enlarge -> CUBIC or LINEAR\n    assert mode in [\"trilinear\", \"nearest\", \"area\", \"nearest-exact\"]\n    assert len(ori_tomo.shape) == 3\n    ori_tomo = ori_tomo[None, None].to(device)  # 11ZYX\n    if mode != \"nearest\":\n        ori_tomo = ori_tomo.float()\n\n    scale_factor = tuple(ori / tgt for ori, tgt in zip(ori_spacing, target_spacing))\n    ##### ISSUES #####\n    ### with interpolation mode `nearest`\n    # RuntimeError: upsample_nearest3d only supports output tensors with less than INT_MAX elements, but got [1, 1, 1178, 1367, 1367]\n    # ref: https://github.com/pytorch/pytorch/issues/144855\n    ### with interpolation mode `trilinear` on GPU\n    # RuntimeError: CUDA error: invalid configuration argument\n    ##### CURRENT SIMPLE FIX: if CUDA-interpolation-kernel fails, use CPU instead\n    try:\n        start = time.time()\n        spaced_tomo = F.interpolate(\n            ori_tomo,\n            size=None,\n            scale_factor=scale_factor,\n            mode=mode,\n            align_corners=False if mode == \"trilinear\" else None,\n            recompute_scale_factor=False,\n        )[0, 0]\n        end = time.time()\n        logger.debug(\"INTERPOLATE WITH MODE=%s TAKE %.2f sec\", mode, end - start)\n    except RuntimeError as e:\n        logger.warning(\n            \"EXCEPTION with torch.interpolate():\\n%s\\nAttempt CPU interpolation..\", e\n        )\n        spaced_tomo = F.interpolate(\n            ori_tomo.cpu(),\n            size=None,\n            scale_factor=scale_factor,\n            mode=mode,\n            align_corners=False if mode == \"trilinear\" else None,\n            recompute_scale_factor=False,\n        )[0, 0]\n    if mode != \"nearest\":\n        spaced_tomo = spaced_tomo.to(torch.uint8)  # ZYX\n    spaced_tomo = spaced_tomo.cpu()\n    assert spaced_tomo.dtype == torch.uint8\n    return spaced_tomo\n\n\ndef scale_range(start, end, scale, minv, maxv):\n    n = max(1, round((end - start) * scale))\n    if round(end * scale) - round(start * scale) == n:\n        return round(start * scale), round(end * scale)\n    start1 = int(start * scale)\n    start2 = round(start * scale)\n    candidate_starts = []\n    if minv <= start1 < start1 + n <= maxv:\n        candidate_starts.append(start1)\n    if minv <= start2 < start2 + n <= maxv:\n        candidate_starts.append(start2)\n    if len(candidate_starts) == 1:\n        return candidate_starts[0], candidate_starts[0] + n\n    elif len(candidate_starts) == 2:\n        if candidate_starts[0] == candidate_starts[1]:\n            return candidate_starts[0], candidate_starts[0] + n\n        diffs = [\n            abs(s - start * scale) + abs(s + n - end * scale) for s in candidate_starts\n        ]\n        if diffs[0] < diffs[1]:\n            return candidate_starts[0], candidate_starts[0] + n\n        else:\n            return candidate_starts[1], candidate_starts[1] + n\n    else:\n        raise ValueError\n\n\n@exception_printer\ndef tomo_producer_worker_func(\n    in_queue: mp.Queue,\n    out_queue: queue.Queue,\n    num_remain_workers: mp.Value,\n    lock: mp.Lock,\n    device: str = \"cpu\",\n    mode: str = \"trilinear\",\n    worker_id: str = \"\",\n):\n    \"\"\"Read tomo, spacing, then save to shared memory.\n    Should be run at process level (multiprocessing)\n    \"\"\"\n    logger.info(\"[TOMO PRODUCER %s] STARTING..\", worker_id)\n    _old_num_threads = torch.get_num_threads()\n    if _old_num_threads < 16:\n        torch.set_num_threads(16)\n    logger.info(\n        \"[TOMO PRODUCER %s] Number of threads change from %d to %d\",\n        worker_id,\n        _old_num_threads,\n        torch.get_num_threads(),\n    )\n    tomo_loader = MultithreadOpencvTomogramLoader(\n        num_workers=_CV2_TOMO_LOADER_NUM_THREADS\n    )\n    while True:\n        item = in_queue.get()\n        if item is None:\n            with lock:\n                num_remain_workers.value -= 1\n                if num_remain_workers.value == 0:\n                    # all other workers are stopped, this worker is the last one alive\n                    # put None to indicate stopping to out_queue's consumer\n                    out_queue.put(None)\n            # other workers should stop too..\n            in_queue.put(None)\n            logger.debug(\n                \"[TOMO PRODUCER %s] STOPPING..\",\n                worker_id,\n            )\n            return\n        start = time.time()\n        tomo_idx, tomo_id, ori_spacing = item\n        # fetch new item\n        tomo = tomo_loader.load(os.path.join(DATA_DIR, tomo_id))\n        ori_shape = tomo.shape\n        target_spacings = TARGET_SPACINGS\n        tomo = torch.from_numpy(tomo)\n        assert tomo.dtype == torch.uint8\n\n        # FIX BUG WRONG QUANTILE\n        hist = torch.bincount(tomo.view(-1), minlength=256).cpu().numpy()\n        assert hist.shape[0] == 256\n        CUTOFF = 25\n        is_wrong_quantile_dark = hist[:CUTOFF].sum() / hist[CUTOFF:].sum() > 0.08\n        is_wrong_quantile_light = hist[-CUTOFF:].sum() / hist[:-CUTOFF].sum() > 0.08\n        if is_wrong_quantile_dark or is_wrong_quantile_light:\n            tomo = tomo + 127\n\n        # Multiscale TTA\n        logger.debug(\n            \"[TOMO PRODUCER %s] Read tomo_idx=%d tomo_id=%s tomo_shape=%s\",\n            worker_id,\n            tomo_idx,\n            tomo_id,\n            tomo.shape,\n        )\n        proxy_tomos = []\n        for spacing_idx, target_spacing in enumerate(target_spacings):\n            spaced_tomo = spacing_torch(\n                tomo, ori_spacing, target_spacing, device=device, mode=mode\n            )\n            gc.collect()\n            logger.debug(\n                \"[TOMO PRODUCER %s] Spacing tomo_idx=%d tomo_id=%s shape %s -> %s\",\n                worker_id,\n                tomo_idx,\n                tomo_id,\n                tomo.shape,\n                spaced_tomo.shape,\n            )\n            proxy_tomos.append(\n                ShmTensor.from_tensor(\n                    spaced_tomo, f\"{tomo_idx}_{tomo_id}_{spacing_idx}\"\n                )\n            )\n        del tomo\n        gc.collect()\n        end = time.time()\n        logger.info(\n            \"[TOMO PRODUCER %s] Done tomo_idx=%d tomo_id=%s spacings %s -> %s with shape %s -> %s take %.2f sec\",\n            worker_id,\n            tomo_idx,\n            tomo_id,\n            ori_spacing,\n            target_spacings,\n            ori_shape,\n            tuple(e.shape for e in proxy_tomos),\n            end - start,\n        )\n        out_queue.put((tomo_idx, tomo_id, ori_spacing, target_spacings, proxy_tomos))\n\n\n@exception_printer\ndef patches_producer_worker_func(\n    in_queue: mp.Queue,\n    out_queue: queue.Queue,\n    num_remain_workers: mp.Value,\n    lock: threading.Lock,\n    patch_size,\n    border,\n    overlap,\n    worker_id=\"\",\n):\n    \"\"\"\n    Should be lightweight, and expected to be run in thread level.\n\n    Args:\n        in_queue: Input queue contain item of (tomo_idx, tomo_id, ori_spacing, target_spacing, proxy_tomo)\n        out_queue: Each item in output queue contains multiple patches belong to a same tomo\n    \"\"\"\n    logger.info(\"[PATCHES PRODUCER %s] STARTING..\", worker_id)\n    while True:\n        item = in_queue.get()\n        if item is None:\n            with lock:\n                num_remain_workers.value -= 1\n                if num_remain_workers.value == 0:\n                    # put final None so out_queue's consumer should stop too\n                    out_queue.put(None)\n            # put None back to in_queue so other in_queue's consumer workers will stop too\n            in_queue.put(None)\n            logger.info(\n                \"\\n\\n[PATCHES PRODUCER %s] STOPPING..\\n\\n\",\n                worker_id,\n            )\n            return\n        tomo_idx, tomo_id, ori_spacing, target_spacings, proxy_tomos = item\n        start = time.time()\n        assert len(proxy_tomos) == len(target_spacings) == len(TARGET_SPACINGS)\n\n        items = []\n        total_patches = 0\n        cur_patch_idx = -1\n        shm_tomos = []\n        for target_spacing, proxy_tomo in zip(target_spacings, proxy_tomos):\n            shm_tomo, tomo = proxy_tomo.to_tensor()\n            tomo_shape = tomo.shape\n            shm_tomos.append(shm_tomo)\n            patch_positions = get_sliding_patch_positions(\n                img_size=tomo.shape,\n                patch_size=patch_size,\n                border=border,\n                overlap=overlap,\n                validate=False,\n            )\n            # patch_positions = torch.from_numpy(patch_positions)\n            patch_positions = patch_positions.tolist()\n            total_patches += len(patch_positions)\n\n            for patch_pos in patch_positions:\n                # crop tomo patch\n                crop = crop_tomo_patch(tomo, patch_pos)  # ZYX, uint8\n                cur_patch_idx += 1\n                item = {\n                    \"tomo_idx\": torch.tensor(tomo_idx, dtype=torch.int16),\n                    \"patch_pos\": torch.tensor(patch_pos, dtype=torch.float32),\n                    \"tomo_shape\": torch.tensor(tomo_shape, dtype=torch.int16),\n                    \"ori_spacing\": torch.tensor(ori_spacing, dtype=torch.float32),\n                    \"target_spacing\": torch.tensor(target_spacing, dtype=torch.float32),\n                    \"cur_patch_idx\": torch.tensor(cur_patch_idx, dtype=torch.int16),\n                    \"image\": crop,\n                }\n                items.append(item)\n        assert cur_patch_idx == total_patches - 1\n        for item in items:\n            item[\"total_patches\"] = torch.tensor(total_patches, dtype=torch.int16)\n\n        end = time.time()\n        logger.debug(\n            \"[PATCHES PRODUCER %s] Done tomo_idx=%d tomo_id=%s take %.2f sec\",\n            worker_id,\n            tomo_idx,\n            tomo_id,\n            end - start,\n        )\n        out_queue.put((tomo_idx, shm_tomos, items))\n\n\n@exception_printer\ndef batch_producer_worker_func(\n    in_queue: queue.Queue,\n    out_queue: queue.Queue,\n    batch_size,\n    device=\"cpu\",\n):\n    \"\"\"Should be lightweight, and expected to be run in thread level.\"\"\"\n    count = 0\n    items = []\n    should_continue = True\n    shm_dict = {}\n    last_tomo_idx = -1\n    while should_continue:\n        item = in_queue.get()\n        if item is None:\n            # put None so in_queue's consumer should notice and stop too\n            in_queue.put(None)\n            should_continue = False\n        else:\n            tomo_idx, shm_tomos, new_items = item\n            # @TODO - remove this line\n            assert all([tomo_idx == e[\"tomo_idx\"].item() for e in new_items])\n            items.extend(new_items)\n            shm_dict[tomo_idx] = shm_tomos\n\n        while True:\n            if len(items) == 0 or (len(items) < batch_size and should_continue):\n                break\n            start = time.time()\n            cur_bs = min(batch_size, len(items))\n            batch = items[:cur_bs]\n            batch = default_collate(batch)\n            if device != \"cpu\":\n                batch = {k: v.to(device) for k, v in batch.items()}\n            assert batch[\"image\"].dtype == torch.uint8\n            end = time.time()\n            logger.debug(\n                \"[BATCH PRODUCER] new batch %d (tomo %d) take %.4f\",\n                count,\n                tomo_idx,\n                end - start,\n            )\n            out_queue.put(batch)\n            count += 1\n            last_tomo_idx = batch[\"tomo_idx\"][-1].item()\n            items = items[cur_bs:]\n        for shm_tomo_idx in list(shm_dict.keys()):\n            # un-collated batch should start with tomo_idx >= last_tomo_idx\n            # so we're safe to unlink previous one\n            # since default_collate() copy to a newly allocated tensor\n            if shm_tomo_idx < last_tomo_idx:\n                for shm in shm_dict[shm_tomo_idx]:\n                    shm.unlink()\n                del shm_dict[shm_tomo_idx]\n\n    for shms in shm_dict.values():\n        for shm in shms:\n            shm.unlink()\n    del shm_dict\n    gc.collect()\n    # put a final None to indicate `all done`, stop inference worker too..\n    out_queue.put(None)\n    logger.info(\"[BATCH PRODUCER] Stopped since tomo_idx=%d\", tomo_idx)\n    assert len(items) == 0\n\n\n@exception_printer\ndef inference_worker(worker_id, input_queue, chunk_idx, gpu_id, model_cfgs):\n    logger.info(\n        \"STARTING INFERENCE WORKER %d with device cuda:%d, current input queue size %d, %d models with configs:\\n%s\",\n        worker_id,\n        gpu_id,\n        input_queue.qsize(),\n        len(model_cfgs),\n        model_cfgs,\n    )\n    device = torch.device(f\"cuda:{gpu_id}\")\n    patch_size = PATCH_SIZE\n    border = PATCH_BORDER\n    overlap = PATCH_OVERLAP\n\n    ############ START ALL WORKERS ###########\n    ##########################################\n    tomo_queue = mp.Queue(maxsize=TOMO_PRODUCER_PREFETCH)\n    tomo_producer_num_remain_workers = mp.Value(\"i\", TOMO_PRODUCER_NUM_WORKERS)\n    tomo_producer_lock = mp.Lock()\n    patches_queue = queue.Queue(maxsize=PATCHES_PRODUCER_PREFETCH)\n    patches_producer_num_remain_workers = mp.Value(\"i\", PATCHES_PRODUCER_NUM_WORKERS)\n    patches_producer_lock = threading.Lock()\n    batch_queue = queue.Queue(maxsize=DATALOADER_PREFETCH)\n\n    # PRODUCE SPACED TOMO\n    tomo_producer_workers = []\n    for _worker_id in range(TOMO_PRODUCER_NUM_WORKERS):\n        worker = mp.Process(\n            group=None,\n            target=tomo_producer_worker_func,\n            name=f\"tomo_producer_{worker_id}.{_worker_id}\",\n            args=(\n                input_queue,\n                tomo_queue,\n                tomo_producer_num_remain_workers,\n                tomo_producer_lock,\n            ),\n            kwargs={\n                \"device\": TOMO_SPACING_DEVICE,\n                \"mode\": TOMO_SPACING_MODE,\n                \"worker_id\": f\"{worker_id}.{_worker_id}\",\n            },\n        )\n        worker.start()\n        tomo_producer_workers.append(worker)\n\n    # PRODUCER TOMO PATCHES\n    patches_producer_workers = []\n    for _worker_id in range(PATCHES_PRODUCER_NUM_WORKERS):\n        worker = threading.Thread(\n            group=None,\n            target=patches_producer_worker_func,\n            name=f\"patches_producer_{worker_id}.{_worker_id}\",\n            args=(\n                tomo_queue,\n                patches_queue,\n                patches_producer_num_remain_workers,\n                patches_producer_lock,\n                patch_size,\n                border,\n                overlap,\n            ),\n            kwargs={\"worker_id\": f\"{worker_id}.{_worker_id}\"},\n        )\n        worker.start()\n        patches_producer_workers.append(worker)\n\n    # BATCH PRODUCER\n    batch_producer_worker = threading.Thread(\n        group=None,\n        target=batch_producer_worker_func,\n        name=\"batch_producer\",\n        args=(patches_queue, batch_queue, BATCH_SIZE),\n        kwargs={\"device\": \"cuda:0\"},\n    )\n    batch_producer_worker.start()\n\n    ############## LOAD MODEL ##################\n    ############################################\n    logger.info(\"LOADING MODEL..\")\n    models = []\n    unique_tta_names = set()\n    unique_agg_names = set()\n    sessions = []\n    for session_idx, model_cfg in enumerate(model_cfgs):\n        model_tta_names = model_cfg[\"tta\"]\n        model_agg_weights = model_cfg[\"weight\"]\n        model_agg_names = model_cfg[\"agg_name\"]\n        assert len(model_tta_names) == len(model_agg_weights) == len(model_agg_names)\n        assert len(model_tta_names) == len(set(model_tta_names))\n        model = Model2dWrapper(\n            model_cfg[\"sig\"],\n            gpu_id,\n            act=\"identity\" if HEATMAP_AGG_LOGITS else \"sigmoid\",\n        )\n        model.eval().to(device)\n        models.append(model)\n        unique_tta_names.update(model_tta_names)\n        unique_agg_names.update(model_agg_names)\n        logger.info(\n            \"Loaded model %d with TTA=%s, WEIGHTS=%s\",\n            session_idx,\n            model_tta_names,\n            model_agg_weights,\n        )\n        for tta_name, agg_weight, agg_name in zip(\n            model_tta_names, model_agg_weights, model_agg_names\n        ):\n            assert tta_name in ALL_SUPPORTED_TTAS\n            sessions.append((model, tta_name, agg_weight, agg_name))\n    unique_tta_names = list(unique_tta_names)\n    unique_agg_names = list(unique_agg_names)\n    num_sessions = len(sessions)\n    _session_agg_names = [e[-1] for e in sessions]\n    agg_name_to_first_session_idx = {\n        agg_name: _session_agg_names.index(agg_name) for agg_name in unique_agg_names\n    }\n    agg_name_to_last_session_idx = {\n        agg_name: num_sessions - 1 - _session_agg_names[::-1].index(agg_name)\n        for agg_name in unique_agg_names\n    }\n    logger.info(\n        \"FIRST/LAST SESSION IDX:\\n%s\\n%s\",\n        agg_name_to_first_session_idx,\n        agg_name_to_last_session_idx,\n    )\n\n    logger.info(\"-----------\\nLOADED %d MODELS:\\n%s------------\", len(models), models)\n    logger.info(\n        \"TOTALLY ENABLE %d MODELS, %d TTA %s, %d AGG_NAMES %s => %d SESSIONS\",\n        len(models),\n        len(unique_tta_names),\n        unique_tta_names,\n        len(unique_agg_names),\n        unique_agg_names,\n        len(sessions),\n    )\n    unique_tta_tfs = {\n        tta_name: build_tta(tta_name, 1.0, ori_dims=\"yx\")\n        for tta_name in unique_tta_names\n    }\n\n    # Postprocess: heatmap blurring\n    if DECODE_NMS_BLUR_SIGMA is not None:\n        from yagm.transforms import monai_custom as CT\n\n        sigma_voxel = [\n            1000 / e / s * DECODE_NMS_BLUR_SIGMA\n            for e, s in zip(TARGET_SPACINGS[0], HEATMAP_STRIDE)\n        ]\n        blur_operator = (\n            CT.CustomGaussianFilter(\n                spatial_dims=3,\n                sigma=sigma_voxel,\n                truncated=4,\n                approx=\"erf\",\n                requires_grad=False,\n            )\n            .eval()\n            .to(device)\n        )\n    else:\n        blur_operator = None\n    logger.info(\n        \"POST-PROCESSING BLUR OPERATOR:\\n%s on device %s\",\n        blur_operator,\n        getattr(blur_operator, \"device\", None),\n    )\n\n    cache_results = {\n        agg_name: {\n            \"cur_pred_heatmap_sum\": None,\n            \"cur_pred_heatmap_count\": None,\n            \"cur_tomo_idx\": None,\n            \"cur_tomo_id\": None,\n            \"main_target_spacing\": None,\n            \"submission\": SubmissionDataFrame(),\n        }\n        for agg_name in unique_agg_names\n    }\n    logger.info(\"\\n\\n\\nWORKER %d STARTING INFERENCE..\\n\", worker_id)\n    infer_start = time.time()\n    cur_batch_idx = 0\n    pbar = tqdm(desc=f\"===== INFERENCE {worker_id} =====\")\n    with torch.inference_mode(), torch.autocast(\n        device_type=\"cuda\", dtype=torch.float16\n    ):\n        while True:\n            batch = batch_queue.get()\n            if batch is None:\n                break\n            # print(\"GOT BATCH:\", {(k, v.shape, v.dtype) for k, v in batch.items()})\n            batch_image = batch[\"image\"].to(device)\n            assert batch_image.dtype == torch.uint8\n            B = batch_image.shape[0]\n\n            unique_tta_batch_images = {\n                tta_name: tta_tf.transform(batch_image)[0].float().contiguous()\n                for tta_name, tta_tf in unique_tta_tfs.items()\n            }\n\n            batch_outputs = []\n            for model, tta_name, agg_weight, agg_name in sessions:\n                tta_tf = unique_tta_tfs[tta_name]\n                tta_batch_image = unique_tta_batch_images[tta_name]\n                tta_batch_heatmap = model(tta_batch_image)\n                batch_heatmap = tta_tf.invert(tta_batch_heatmap)[0]\n                batch_outputs.append((batch_heatmap, agg_weight, agg_name))\n\n            assert len(batch_outputs) == num_sessions\n\n            for i in range(B):\n                for session_idx, (batch_heatmap, agg_weight, agg_name) in enumerate(\n                    batch_outputs\n                ):\n                    pred_patch_heatmap = batch_heatmap[i]\n                    C = pred_patch_heatmap.shape[0]\n                    assert C == 1\n\n                    tomo_idx = batch[\"tomo_idx\"][i]\n                    tomo_id = ALL_TOMO_IDS[tomo_idx]\n                    tomo_shape = batch[\"tomo_shape\"][i].tolist()\n                    ori_spacing = batch[\"ori_spacing\"][i].tolist()\n                    target_spacing = batch[\"target_spacing\"][i].tolist()\n                    cur_patch_idx = int(batch[\"cur_patch_idx\"][i].item())\n                    total_patches = int(batch[\"total_patches\"][i].item())\n                    patch_position = batch[\"patch_pos\"][i].cpu()\n\n                    is_first = (cur_patch_idx == 0) and (\n                        session_idx == agg_name_to_first_session_idx[agg_name]\n                    )\n                    is_last = (cur_patch_idx == total_patches - 1) and (\n                        session_idx == agg_name_to_last_session_idx[agg_name]\n                    )\n\n                    # load cache for this tomo+agg combination\n                    agg_result = cache_results[agg_name]\n                    ### start IS_FIRST\n                    if is_first:\n                        logger.info(\n                            \"[INFER %d] Model %s receive first patch of tomo %s (index %d)\",\n                            worker_id,\n                            session_idx,\n                            tomo_id,\n                            tomo_idx,\n                        )\n                        assert all(\n                            [\n                                agg_result[k] is None\n                                for k in [\n                                    \"cur_tomo_idx\",\n                                    \"cur_tomo_id\",\n                                    \"cur_pred_heatmap_sum\",\n                                    \"cur_pred_heatmap_count\",\n                                    \"main_target_spacing\",\n                                ]\n                            ]\n                        )\n\n                        # allocate new heatmap_sum and heatmap_count\n                        agg_result[\"cur_tomo_idx\"] = tomo_idx\n                        agg_result[\"cur_tomo_id\"] = tomo_id\n                        shared_agg_heatmap_shape = [\n                            round(e / stride)\n                            for e, stride in zip(tomo_shape, HEATMAP_STRIDE)\n                        ]\n                        print(tomo_shape, '-->', shared_agg_heatmap_shape)\n                        # float16 to save memory\n                        agg_result[\"cur_pred_heatmap_sum\"] = torch.zeros(\n                            (C, *shared_agg_heatmap_shape),\n                            dtype=torch.float16,\n                            device=pred_patch_heatmap.device,\n                        )\n                        agg_result[\"cur_pred_heatmap_count\"] = torch.zeros(\n                            (C, *shared_agg_heatmap_shape),\n                            dtype=torch.float16,\n                            device=pred_patch_heatmap.device,\n                        )\n                        agg_result[\"main_target_spacing\"] = target_spacing\n                        assert all(\n                            [\n                                abs(a - b) < 1e-4\n                                for a, b in zip(target_spacing, TARGET_SPACINGS[0])\n                            ]\n                        )\n                    ### end IS_FIRST\n\n                    assert tomo_idx == agg_result[\"cur_tomo_idx\"]\n                    assert all(\n                        [\n                            a <= b\n                            for a, b in zip(\n                                agg_result[\"main_target_spacing\"], target_spacing\n                            )\n                        ]\n                    )\n                    cur_pred_heatmap_sum = agg_result[\"cur_pred_heatmap_sum\"]\n                    cur_pred_heatmap_count = agg_result[\"cur_pred_heatmap_count\"]\n                    main_target_spacing = agg_result[\"main_target_spacing\"]\n                    cur_submission: SubmissionDataFrame = agg_result[\"submission\"]\n                    shared_agg_heatmap_shape = cur_pred_heatmap_sum.shape[1:]\n\n                    scale_zyx = [\n                        a / b for a, b in zip(shared_agg_heatmap_shape, tomo_shape)\n                    ]\n                    scale_z, scale_y, scale_x = scale_zyx\n\n                    # interpolate to shared coordinate space\n                    C, Y2, X2 = pred_patch_heatmap.shape\n                    assert C == 1\n                    patch_heatmap_shape = [\n                        round(e * s) for e, s in zip(patch_size, scale_zyx)\n                    ]\n                    if [Y2, X2] != patch_heatmap_shape[1:]:\n                        logger.debug(\n                            \"PATCH HEATMAP INTERPOLATE: %s --> %s\",\n                            (Y2, X2),\n                            patch_heatmap_shape[1:],\n                        )\n                        pred_patch_heatmap = F.interpolate(\n                            pred_patch_heatmap[None],  # (1, C, Y2, X2)\n                            size=patch_heatmap_shape[1:],\n                            mode=HEATMAP_INTERPOLATION_MODE,\n                            align_corners=False,\n                        )[\n                            0\n                        ]  # (C, Y3, X3)\n                    else:\n                        logger.debug(\"PATCH HEATMAP SAME: %s\", pred_patch_heatmap.shape)\n                        pass\n\n                    roi_start, roi_end, patch_start, patch_end = patch_position.tolist()\n                    top_pad_z, top_pad_y, top_pad_x = [\n                        max(-start, 0) for start in roi_start\n                    ]\n                    bot_pad_z, bot_pad_y, bot_pad_x = [\n                        max(0, end - size) for end, size in zip(roi_end, tomo_shape)\n                    ]\n                    assert (\n                        top_pad_z == bot_pad_z == 0\n                        and roi_end[0] - roi_start[0] == patch_size[0]\n                    )\n                    assert (\n                        tuple(\n                            round((pe - ps) * s)\n                            for pe, ps, s in zip(\n                                patch_end[1:], patch_start[1:], scale_zyx[1:]\n                            )\n                        )\n                        == pred_patch_heatmap.shape[1:]\n                    )\n                    sz, sy, sx = tuple(\n                        rs - ps for rs, ps in zip(roi_start, patch_start)\n                    )\n                    ez, ey, ex = tuple(\n                        ps - pe + re\n                        for ps, pe, re in zip(\n                            (1, patch_size[1], patch_size[2]), patch_end, roi_end\n                        )\n                    )\n                    assert sz == 0 and ez == 1 and top_pad_z == bot_pad_z == 0\n                    assert (roi_end[1] - bot_pad_y) - (roi_start[1] + top_pad_y) == (\n                        ey - bot_pad_y\n                    ) - (sy + top_pad_y)\n                    assert (roi_end[2] - bot_pad_x) - (roi_start[2] + top_pad_x) == (\n                        ex - bot_pad_x\n                    ) - (sx + top_pad_x)\n\n                    dst_shape = shared_agg_heatmap_shape\n                    dst_slices = [\n                        slice(None),\n                        slice(\n                            *scale_range(\n                                roi_start[0] + top_pad_z,\n                                roi_end[0] - bot_pad_z,\n                                scale_z,\n                                0,\n                                dst_shape[0],\n                            )\n                        ),\n                        slice(\n                            *scale_range(\n                                roi_start[1] + top_pad_y,\n                                roi_end[1] - bot_pad_y,\n                                scale_y,\n                                0,\n                                dst_shape[1],\n                            )\n                        ),\n                        slice(\n                            *scale_range(\n                                roi_start[2] + top_pad_x,\n                                roi_end[2] - bot_pad_x,\n                                scale_x,\n                                0,\n                                dst_shape[2],\n                            )\n                        ),\n                    ]\n                    src_shape = pred_patch_heatmap.shape[1:]\n                    src_slices = [\n                        slice(None),\n                        None,\n                        slice(\n                            *scale_range(\n                                sy + top_pad_y, ey - bot_pad_y, scale_y, 0, src_shape[0]\n                            )\n                        ),\n                        slice(\n                            *scale_range(\n                                sx + top_pad_x, ex - bot_pad_x, scale_x, 0, src_shape[1]\n                            )\n                        ),\n                    ]\n\n                    if HEATMAP_AGG_MODE == \"avg\":\n                        cur_pred_heatmap_sum[dst_slices] += (\n                            pred_patch_heatmap[src_slices] * agg_weight\n                        )\n                        cur_pred_heatmap_count[dst_slices] += agg_weight\n                    elif HEATMAP_AGG_MODE == \"max\":\n                        torch.maximum(\n                            cur_pred_heatmap_sum[dst_slices],\n                            pred_patch_heatmap[src_slices],\n                            out=cur_pred_heatmap_sum[dst_slices],\n                        )\n                    else:\n                        raise ValueError\n\n                    if is_last:\n                        # DECODE HEATMAP TO COORDINATE\n                        tomo_id = agg_result[\"cur_tomo_id\"]\n                        logger.info(\n                            \"[INFER %d] Receive last patch of tomo %s (index %d), decoding..\",\n                            worker_id,\n                            tomo_id,\n                            tomo_idx,\n                        )\n\n                        if HEATMAP_AGG_MODE == \"avg\":\n                            # ensure prediction cover all the tomogram\n                            # assert torch.all(cur_pred_heatmap_count > 0)\n                            heatmap = torch.div(\n                                cur_pred_heatmap_sum,\n                                cur_pred_heatmap_count,\n                                out=cur_pred_heatmap_sum,\n                            )\n                        elif HEATMAP_AGG_MODE == \"max\":\n                            heatmap = cur_pred_heatmap_sum\n                        else:\n                            raise ValueError\n                        if HEATMAP_AGG_LOGITS:\n                            heatmap = F.sigmoid(heatmap)\n\n                        if SAVE_HEATMAP_NPY:\n                            heatmap_npy = heatmap.half().cpu().numpy()\n                            assert heatmap_npy.dtype == np.float16\n                            save_heatmap_path = os.path.join(\n                                SAVE_HEATMAP_NPY_DIR, agg_name, f\"{tomo_id}.npy\"\n                            )\n                            os.makedirs(\n                                os.path.dirname(save_heatmap_path), exist_ok=True\n                            )\n                            np.save(save_heatmap_path, heatmap_npy)\n\n                        if DECODE_HEATMAP_TO_CSV:\n                            # Decode heatmap\n                            _start = time.time()\n                            if DECODE_METHOD == \"nms\":\n                                outputs = decode_heatmap(\n                                    heatmap,\n                                    main_target_spacing,\n                                    HEATMAP_STRIDE,\n                                    blur_operator=blur_operator,\n                                )\n                            elif DECODE_METHOD == \"cc3d\":\n                                outputs = decode_segment_mask(\n                                    heatmap,\n                                    main_target_spacing,\n                                    HEATMAP_STRIDE,\n                                    prob_thres=DECODE_CC3D_CONF_THRES,\n                                    radius_factor_thres=DECODE_CC3D_RADIUS_FACTOR,\n                                    conf_mode=DECODE_CC3D_CONF_MODE,\n                                    prob_mode=DECODE_CC3D_PROB_MODE,\n                                )\n                            else:\n                                raise ValueError\n                            _end = time.time()\n                            logger.debug(\n                                \"Decode heatmap of shape %s take %.4f sec\",\n                                heatmap.shape,\n                                _end - _start,\n                            )\n                            # Convert to submission coordinates\n                            assert len(outputs) == 1\n                            heatmap_stride_tensor = torch.tensor(\n                                [HEATMAP_STRIDE], dtype=torch.float32\n                            )  # (1, 3)\n                            main_target_spacing_tensor = torch.tensor(\n                                [main_target_spacing], dtype=torch.float32\n                            )  # (1, 3)\n                            ori_spacing_tensor = torch.tensor(\n                                [ori_spacing], dtype=torch.float32\n                            )  # (1, 3)\n                            for _channel_idx, keypoints in enumerate(outputs):\n                                # back to original coordinate space\n                                keypoints[:, :3] = (keypoints[:, :3] + 0.5) * (\n                                    heatmap_stride_tensor\n                                    * main_target_spacing_tensor\n                                    / ori_spacing_tensor\n                                ) - 0.5\n                                keypoints = keypoints.tolist()\n\n                                ###### VISUALIZATION ######\n                                if VIZ_ENABLE and MODE != \"KAGGLE\":\n                                    tomo_dir = os.path.join(DATA_DIR, tomo_id)\n                                    row = GT_DF[GT_DF[\"tomo_id\"] == tomo_id].iloc[0]\n                                    ori_shape = (row[\"Z\"], row[\"Y\"], row[\"X\"])\n                                    gt_zyx = [\n                                        row[\"motor_z\"],\n                                        row[\"motor_y\"],\n                                        row[\"motor_x\"],\n                                    ]\n                                    viz = viz_byu(\n                                        tomo_dir,\n                                        ori_spacing,\n                                        ori_shape,\n                                        gt_zyx,\n                                        heatmap[0],\n                                        keypoints,\n                                    )\n                                    # get the kind/judge of current prediction, one of TP, TN, FP, FN\n                                    if (\n                                        len(keypoints) > 0\n                                        and keypoints[0][3] >= DECODE_CONF_THRES\n                                    ):\n                                        # predicted as positive\n                                        pred_zyx = keypoints[0][:3]\n                                        if gt_zyx[0] == -1:\n                                            # GT is negative\n                                            judge = \"FP\"\n                                        else:\n                                            dist = (\n                                                (gt_zyx[0] - pred_zyx[0]) ** 2\n                                                + (gt_zyx[1] - pred_zyx[1]) ** 2\n                                                + (gt_zyx[2] - pred_zyx[2]) ** 2\n                                            ) ** 0.5\n                                            if dist <= 1000 / ori_spacing:\n                                                judge = \"TP\"\n                                            else:\n                                                judge = \"FN\"\n                                    else:\n                                        # predicted as negative\n                                        if gt_zyx[0] == -1:\n                                            # GT is negative\n                                            judge = \"TN\"\n                                        else:\n                                            judge = \"FN\"\n\n                                    save_name = f\"{judge}-{tomo_idx}-{tomo_id}-{'x'.join([str(int(e)) for e in ori_shape])}-spacing{ori_spacing}.jpg\"\n                                    save_path = os.path.join(\n                                        TMP_VIZ_DIR, judge, agg_name, save_name\n                                    )\n                                    os.makedirs(\n                                        os.path.dirname(save_path), exist_ok=True\n                                    )\n                                    cv2.imwrite(\n                                        save_path, cv2.cvtColor(viz, cv2.COLOR_RGB2BGR)\n                                    )\n\n                                if len(keypoints) > 0:\n                                    for kpt in keypoints:\n                                        z, y, x, conf = kpt\n                                        cur_submission.add_row(tomo_id, x, y, z, conf)\n                                else:\n                                    # no motor detected\n                                    cur_submission.add_row(tomo_id, -1, -1, -1, 0.0)\n\n                        # clear the cache\n                        for k in [\n                            \"cur_tomo_idx\",\n                            \"cur_tomo_id\",\n                            \"cur_pred_heatmap_sum\",\n                            \"cur_pred_heatmap_count\",\n                            \"main_target_spacing\",\n                        ]:\n                            agg_result[k] = None\n                        gc.collect()\n                        # torch.cuda.empty_cache()\n            cur_batch_idx += 1\n            pbar.update(1)\n    pbar.close()\n    infer_end = time.time()\n\n    logger.info(\n        \"\\n\\n\\n INFER WORKER %d INFER ON %d BATCHES TAKE %.4f sec\",\n        worker_id,\n        cur_batch_idx,\n        infer_end - infer_start,\n    )\n\n    assert len(cache_results) == len(unique_agg_names)\n    if DECODE_HEATMAP_TO_CSV:\n        for agg_name, agg_result in cache_results.items():\n            submision = agg_result[\"submission\"]\n            df = submision.to_pandas(submit=False)\n            csv_path = os.path.join(TMP_CSV_DIR, f\"{agg_name}_chunk{chunk_idx}.csv\")\n            df.to_csv(csv_path, index=False)\n\n    # JOIN\n    logger.info(\"Worker %d waiting for all worker to be finished..\", worker_id)\n    all_workers = [\n        batch_producer_worker,\n        *tomo_producer_workers,\n        *patches_producer_workers,\n    ]\n    for worker in all_workers:\n        worker.join()\n    return None\n\n\n########################### MAIN CODE #############################\n###################################################################\nlogger.warning(\"CLEARING ALL CONTENTS IN: %s\", [TMP_CSV_DIR, TMP_CSV_DIR, TMP_VIZ_DIR])\nclear_dir_content(\"/dev/shm/\")\n# clear_dir_content(TMP_CSV_DIR)\nclear_dir_content(TMP_VIZ_DIR)\nif MODE == \"LOCAL\":\n    clear_dir_content(WORKING_DIR)\nos.makedirs(TMP_CSV_DIR, exist_ok=True)\nos.makedirs(TMP_VIZ_DIR, exist_ok=True)\n\n\nglobal_start = time.time()\n\nall_unique_agg_names = set()\nfor job_idx, model_cfgs in enumerate(JOBS):\n    # fill the input queue\n    input_queue = mp.Queue()\n    for tomo_idx, (tomo_id, ori_spacing) in enumerate(\n        zip(ALL_TOMO_IDS, ALL_TOMO_SPACINGS)\n    ):\n        input_queue.put((tomo_idx, tomo_id, [ori_spacing] * 3))\n    # None to indicate ending and workers should stop when receive this None\n    input_queue.put(None)\n\n    job_unique_agg_names = set()\n    for model_cfg in model_cfgs:\n        job_unique_agg_names.update(model_cfg[\"agg_name\"])\n    assert len(all_unique_agg_names.intersection(job_unique_agg_names)) == 0\n    all_unique_agg_names.update(job_unique_agg_names)\n\n    job_workers = []\n    for chunk_idx, device_id in enumerate(DEVICES):\n        # worker_id, input_queue, chunk_idx, gpu_id, model_cfgs\n        worker = mp.Process(\n            group=None,\n            target=inference_worker,\n            name=f\"INFERENCE_WORKER_job={job_idx}_chunk={chunk_idx}\",\n            kwargs={\n                \"worker_id\": chunk_idx,\n                \"input_queue\": input_queue,\n                \"chunk_idx\": chunk_idx,\n                \"gpu_id\": device_id,\n                \"model_cfgs\": model_cfgs,\n            },\n        )\n        worker.start()\n        job_workers.append(worker)\n\n    for worker in job_workers:\n        worker.join()\n\nglobal_end = time.time()\nlogger.info(\n    \"\\n\\n\\n>>>>>> FINISH ALL INFERENCE TASK WITHIN %.4f sec (%.4f min)\",\n    global_end - global_start,\n    (global_end - global_start) / 60.0,\n)\nall_unique_agg_names = list(all_unique_agg_names)\nprint('ALL UNIQUE AGG NAMES:', all_unique_agg_names)","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-06-04T00:30:51.236013Z","iopub.execute_input":"2025-06-04T00:30:51.236279Z","iopub.status.idle":"2025-06-04T00:30:51.249556Z","shell.execute_reply.started":"2025-06-04T00:30:51.236259Z","shell.execute_reply":"2025-06-04T00:30:51.248648Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python3 /kaggle/working/byu/submit_2d.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T00:30:51.549060Z","iopub.execute_input":"2025-06-04T00:30:51.549297Z","iopub.status.idle":"2025-06-04T00:32:16.114729Z","shell.execute_reply.started":"2025-06-04T00:30:51.549270Z","shell.execute_reply":"2025-06-04T00:32:16.113898Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# POST PROCESSING","metadata":{}},{"cell_type":"code","source":"import warnings\n\nwarnings.simplefilter(\"ignore\")\nimport gc\nimport json\nimport logging\nimport multiprocessing as mp\nimport os\nimport queue\nimport shutil\nimport sys\nimport threading\nimport time\nfrom functools import partial\n\nimport cc3d\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom omegaconf import OmegaConf\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.utils.data import default_collate\nfrom tqdm import tqdm\nfrom yagm.tasks.base_task import BaseTask\nfrom yagm.transforms.keypoints.decode import decode_heatmap_3d, decode_segment_mask_3d\nfrom yagm.transforms.sliding_window import get_sliding_patch_positions\nfrom yagm.transforms.tta_3d import build_tta\nfrom yagm.utils import hydra as hydra_utils\nfrom yagm.utils import lightning as l_utils\nfrom yagm.utils.concurrent import ShmTensor\n\n# Setting up custom logging\nfrom yagm.utils.logging import init_logging, setup_logging\n\nfrom byu.data.io import MultithreadOpencvTomogramLoader\nfrom byu.inference import zoo\nfrom byu.inference.tensorrt_engine import ThreadsafeTRTEngine\nfrom byu.utils.data import SubmissionDataFrame\nfrom byu.utils.metrics import compute_metrics, kaggle_score\nfrom byu.utils.viz import viz_byu\nfrom byu.utils.wbf import weighted_boxes_fusion_3d\n\ntry:\n    hydra_utils.init_hydra()\nexcept:\n    print(\"SKIP RE-INIT HYDRA\")\n\n### SETUP LOGGER FOR IPYTHON ###\n# ref: https://github.com/ipython/ipykernel/issues/111\n# Create logger\nlogger = logging.getLogger()\nlogger.setLevel(logging.INFO)\n# Create STDERR handler\nhandler = logging.StreamHandler(sys.stderr)\n# Create formatter and add it to the handler\nformatter = logging.Formatter(\"[%(levelname)s] %(message)s\")\nhandler.setFormatter(formatter)\n# Set STDERR handler as the only handler\nlogger.handlers = [handler]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T00:32:58.563031Z","iopub.execute_input":"2025-06-04T00:32:58.563394Z","iopub.status.idle":"2025-06-04T00:33:04.447807Z","shell.execute_reply.started":"2025-06-04T00:32:58.563366Z","shell.execute_reply":"2025-06-04T00:33:04.446899Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_unique_agg_names = [\n    'DENSENET-zyx', 'DENSENET-zxy', 'X3D-zyx_y', 'X3D-zyx_x', 'RESNEXT50-zyx_xy', 'CONVNEXT-zxy_xy', 'MAXVIT-yx', 'MAXVIT-xy'\n]\nall_unique_agg_weights = {agg_name: 1.0 for agg_name in all_unique_agg_names}\n\n\n##################################################\n###### GLOBAL VARS\n##################################################\n_CV2_TOMO_LOADER_NUM_THREADS = 8\n\nPATCH_SIZE = [224, 448, 448]\nPATCH_BORDER = [0, 0, 0]\nPATCH_OVERLAP = [0, 0, 0]\n\nIO_BACKEND = \"cv2\"\nBASE_SPACING = 13.1  # avg 15.6\n# TARGET_SPACINGS = [13.1, 16.0, 19.7]\nTARGET_SPACINGS = [[16.0, 16.0, 16.0]]\nTOMO_SPACING_DEVICE = \"cpu\"\nTOMO_SPACING_MODE = \"trilinear\"\n\nHEATMAP_AGG_MODE = \"avg\"\nHEATMAP_AGG_LOGITS = True\nHEATMAP_INTERPOLATION_MODE = \"trilinear\"\nFORWARD_SAVE_MEMORY = False\nBATCH_SIZE = 1\nHEATMAP_STRIDE = 8\n\n# MASK CC3D DECODING\nDECODE_CC3D_CONF_THRES = 0.05\nDECODE_CC3D_RADIUS_FACTOR = 0.1\nDECODE_CC3D_CONF_MODE = \"prob\"  # volume | prob | fuse\nDECODE_CC3D_PROB_MODE = \"center\"  # center | mean | max\n# HEATMAP NMS DECODING\nDECODE_NMS_BLUR_SIGMA = None\n###### MAIN DECODE PARAMS ######\nDECODE_METHOD = \"nms\"  # nms, cc3d\nDECODE_CONF_THRES = 0.05\nQUANTILE_THRES = 0.55\n\nSAVE_HEATMAP_NPY = False\nDECODE_HEATMAP_TO_CSV = True\nCSV_ENSEMBLE_MODE = \"max\"  # max | wbf\nWBF_CONF_TYPE = \"avg\"  # avg | max\nWBF_CONF_THRES = 0.2\n\nVIZ_ENABLE = False\n\n\n# automatic determine the MODE\nif os.path.isdir(\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025\"):\n    MODE = \"KAGGLE\"\nelse:\n    MODE = \"LOCAL\"\n# manual override\n# MODE = \"KAGGLE_VAL\"\n\n# LOCAL DEVELOPMENT ENV\nif MODE == \"LOCAL\":\n    DEVICES = [0]\n    ASSETS_DIR = \"assets/\"\n    WORKING_DIR = \"outputs/submit/working/\"\n\n    DATA_DIR = \"/home/dangnh36/datasets/.comp/byu/raw/train\"\n    VAL_GT_CSV_PATH = \"/home/dangnh36/datasets/.comp/byu/processed/gt_v3.csv\"\n    with open(\n        \"/home/dangnh36/datasets/.comp/byu/processed/cv/v3/skf4_rd42.json\", \"r\"\n    ) as f:\n        cv_meta = json.load(f)\n        ALL_TOMO_IDS = cv_meta[\"folds\"][0][\"val\"][:]  # fold 0\n\n    # DATA_DIR = \"/home/dangnh36/datasets/.comp/byu/processed/pseudo_test/tomograms\"\n    # VAL_GT_CSV_PATH = \"/home/dangnh36/datasets/.comp/byu/processed/pseudo_test/gt.csv\"\n    # ALL_TOMO_IDS = sorted(os.listdir(DATA_DIR)) * 10\n\n    GT_DF = pd.read_csv(VAL_GT_CSV_PATH)\n    GT_DF = GT_DF[GT_DF[\"tomo_id\"].isin(ALL_TOMO_IDS)].reset_index(drop=True)\n    _tomo2spacing = {}\n    for i, row in GT_DF.iterrows():\n        _tomo2spacing[row[\"tomo_id\"]] = row[\"voxel_spacing\"]\n    ALL_TOMO_SPACINGS = [_tomo2spacing[tomo_id] for tomo_id in ALL_TOMO_IDS]\n\n    # TOMO_PRODUCER_NUM_WORKERS = 8\n    # TOMO_PRODUCER_PREFETCH = 8\n    # PATCHES_PRODUCER_NUM_WORKERS = 2\n    # PATCHES_PRODUCER_PREFETCH = 8\n    # DATALOADER_PREFETCH = 8\n\n    TOMO_PRODUCER_NUM_WORKERS = 1\n    TOMO_PRODUCER_PREFETCH = 1\n    PATCHES_PRODUCER_NUM_WORKERS = 1\n    PATCHES_PRODUCER_PREFETCH = 2\n    DATALOADER_PREFETCH = 8\n\n# VALIDATION ON KAGGLE\nelif MODE == \"KAGGLE_VAL\":\n    DEVICES = [0, 1]\n    ASSETS_DIR = \"/kaggle/input/byu-final-dataset/\"\n    DATA_DIR = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train\"\n    WORKING_DIR = \"/kaggle/working/\"\n    VAL_GT_CSV_PATH = \"/kaggle/input/byu-checkpoints/gt_v2.csv\"\n    with open(\"/kaggle/input/byu-checkpoints/skf4_rd42.json\", \"r\") as f:\n        cv_meta = json.load(f)\n        ALL_TOMO_IDS = cv_meta[\"folds\"][0][\"val\"][:50]  # fold 0, top 10 first tomo\n    GT_DF = pd.read_csv(VAL_GT_CSV_PATH)\n    GT_DF = GT_DF[GT_DF[\"tomo_id\"].isin(ALL_TOMO_IDS)].reset_index(drop=True)\n    _tomo2spacing = {}\n    for i, row in GT_DF.iterrows():\n        _tomo2spacing[row[\"tomo_id\"]] = row[\"voxel_spacing\"]\n    ALL_TOMO_SPACINGS = [_tomo2spacing[tomo_id] for tomo_id in ALL_TOMO_IDS]\n\n    TOMO_PRODUCER_NUM_WORKERS = 1\n    TOMO_PRODUCER_PREFETCH = 1\n    PATCHES_PRODUCER_NUM_WORKERS = 1\n    PATCHES_PRODUCER_PREFETCH = 2\n    DATALOADER_PREFETCH = 8\n\n# PRIVATE TEST SUBMISSION\nelif MODE == \"KAGGLE\":\n    DEVICES = [0, 1]\n    ASSETS_DIR = \"/kaggle/input/byu-final-dataset/\"\n    DATA_DIR = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test\"\n    WORKING_DIR = \"/kaggle/working/\"\n    ALL_TOMO_IDS = sorted(os.listdir(DATA_DIR))\n    ALL_TOMO_SPACINGS = [BASE_SPACING] * len(ALL_TOMO_IDS)\n\n    TOMO_PRODUCER_NUM_WORKERS = 1\n    TOMO_PRODUCER_PREFETCH = 1\n    PATCHES_PRODUCER_NUM_WORKERS = 1\n    PATCHES_PRODUCER_PREFETCH = 2\n    DATALOADER_PREFETCH = 8\n\nelse:\n    raise ValueError\n\nTMP_CSV_DIR = os.path.join(WORKING_DIR, \"tmp_csv\")\nTMP_VIZ_DIR = os.path.join(WORKING_DIR, \"tmp_viz\")\nSAVE_HEATMAP_NPY_DIR = os.path.join(WORKING_DIR, \"tmp_heatmap\")\n\n# RECHECK SOME CONDITIONS\nassert len(ALL_TOMO_IDS) == len(ALL_TOMO_SPACINGS)\n# assert not (SAVE_HEATMAP_NPY and DECODE_HEATMAP_TO_CSV) and (SAVE_HEATMAP_NPY or DECODE_HEATMAP_TO_CSV)\n\n\n###################################################\n###### FUNCTIONS/CLASSES DEFINITION\n###################################################\ndef exception_printer(func):\n    def wrapper(*args, **kwargs):\n        try:\n            return func(*args, **kwargs)\n        except Exception as e:\n            logger.exception(\"Exception %s in %s: %s\", type(e), func.__name__, e)\n            import traceback\n\n            traceback.print_exc()\n            raise e\n\n    return wrapper\n\n\ndef clear_dir_content(folder_path):\n    if os.path.isdir(folder_path):\n        for item in os.listdir(folder_path):\n            item_path = os.path.join(folder_path, item)\n            if os.path.isdir(item_path):\n                shutil.rmtree(item_path)\n            else:\n                os.remove(item_path)\n                \n\nif DECODE_HEATMAP_TO_CSV:\n    all_agg_dfs = {}\n    for agg_name in all_unique_agg_names:\n        logger.info(\"PROCESSING AGG_NAME=%s\", agg_name)\n        chunk_agg_dfs = []\n        for chunk_idx in range(len(DEVICES)):\n            csv_name = f\"{agg_name}_chunk{chunk_idx}.csv\"\n            csv_path = os.path.join(TMP_CSV_DIR, csv_name)\n            try:\n                df = pd.read_csv(csv_path)\n                chunk_agg_dfs.append(df)\n                logger.info(\n                    \"Append new chunk Dataframe with shape %s, columns=%s\",\n                    df.shape,\n                    list(df.columns),\n                )\n            except Exception as e:\n                logger.warning(\n                    \"EXCEPTION OCCUR when processing agg_name=%s:\\n%s\\nIGNORE CSV RESULT FILE AT %s\",\n                    agg_name,\n                    e,\n                    csv_path,\n                )\n        agg_df = pd.concat(chunk_agg_dfs, axis=0, ignore_index=True).reset_index(\n            drop=True\n        )\n        logger.info(\n            \"CONCAT ALL TEMPORARY DATAFRAMES INTO A SINGLE DATAFRAME WITH SHAPE %s, COLUMNS=%s\\n%s\",\n            df.shape,\n            list(df.columns),\n            df,\n        )\n        all_agg_dfs[agg_name] = agg_df\n        print(\"-------------------------------------------\\n\\n\")\n    logger.info(\"TOTAL %d AGG DataFrame\", len(all_agg_dfs))\n\nif CSV_ENSEMBLE_MODE == \"max\":\n    df = pd.concat(list(all_agg_dfs.values()), axis=0, ignore_index=True).reset_index(\n        drop=True\n    )\n    df = (\n        df.sort_values(\"conf\", ascending=False)\n        .groupby(\"tomo_id\", as_index=False)\n        .first()\n    )\nelif CSV_ENSEMBLE_MODE == \"wbf\":\n    submission_df = SubmissionDataFrame()\n    assert len(ALL_TOMO_IDS) == len(ALL_TOMO_SPACINGS)\n    for tomo_idx, tomo_id in tqdm(enumerate(ALL_TOMO_IDS), desc=\"WBF\"):\n        tomo_zyxs = []\n        tomo_confs = []\n        tomo_labels = []\n        weights = []\n        for agg_name, agg_df in all_agg_dfs.items():\n            preds = agg_df[agg_df[\"tomo_id\"] == tomo_id][\n                [\"motor_z\", \"motor_y\", \"motor_x\", \"conf\"]\n            ].to_numpy()\n            keep_idxs = np.all(preds[:, :3] != -1, axis=1) & (preds[:, -1] > 0.0)\n            preds = preds[keep_idxs]\n            tomo_zyxs.append(preds[:, :3])\n            tomo_confs.append(preds[:, 3])\n            tomo_labels.append([0] * len(preds))\n            weights.append(all_unique_agg_weights[agg_name])\n        # print(len(tomo_zyxs), len(tomo_confs), len(tomo_labels))\n        # print(tomo_zyxs, tomo_confs, tomo_labels, sep = '\\n\\n')\n        zyxs, confs, labels = weighted_boxes_fusion_3d(\n            tomo_zyxs,\n            tomo_confs,\n            tomo_labels,\n            weights=weights,\n            dist_thr=1000.0 / ALL_TOMO_SPACINGS[tomo_idx],\n            skip_box_thr=WBF_CONF_THRES,\n            conf_type=WBF_CONF_TYPE,\n            allows_overflow=False,\n            rescale_conf=True,\n        )\n        if len(zyxs):\n            # select the highest conf\n            select_idx = np.argmax(confs)\n            z, y, x = zyxs[select_idx]\n            submission_df.add_row(tomo_id, x, y, z, confs[select_idx])\n        else:\n            submission_df.add_row(tomo_id, -1, -1, -1, 0.0)\n    df = submission_df.to_pandas(submit=False)\nelse:\n    raise ValueError\n\nlogger.info(\n    \"PREDICTION DATAFRAME WITH SHAPE %s COLUMNS=%s:\\n%s\",\n    df.shape,\n    list(df.columns),\n    df,\n)\n\nif MODE in [\"LOCAL\", \"KAGGLE_VAL\"]:\n    assert (\n        len(GT_DF) == len(df) == len(ALL_TOMO_IDS)\n    ), f\"{GT_DF.shape} {df.shape} {len(ALL_TOMO_IDS)}\"\n    metrics = compute_metrics(GT_DF, df, mode=\"kaggle\")\n    logger.info(\"%s METRICS:\\n%s\", MODE, metrics)\n\n# Filter based on conf\nif QUANTILE_THRES is None:\n    filter_thres = DECODE_CONF_THRES\nelse:\n    assert 0 <= QUANTILE_THRES <= 1.0\n    pred_confs = df[\"conf\"].values\n    filter_thres = np.quantile(pred_confs, QUANTILE_THRES)\n    # assert 0 <= filter_thres <= 1.0\n\n\nlogger.info(\"USING FILTER THRESHOLD: %f\", filter_thres)\ndf.loc[df[\"conf\"] < filter_thres, [\"motor_z\", \"motor_y\", \"motor_x\"]] = -1\nif MODE in [\"LOCAL\", \"KAGGLE_VAL\"]:\n    kaggle_fbeta, kaggle_precision, kaggle_recall, (tn, fp, fn, tp) = kaggle_score(\n        GT_DF, df\n    )\n    logger.info(\n        \"--------------\\nKAGGLE METRICS AT THRESHOLD=%f:\\n\\nFBETA=%f PRECISION=%f RECALL=%f\\nTN=%d FP=%d FN=%d TP=%d--------------\\n\\n\",\n        filter_thres,\n        kaggle_fbeta,\n        kaggle_precision,\n        kaggle_recall,\n        tn,\n        fp,\n        fn,\n        tp,\n    )\n\ndf.rename(\n    columns={\n        \"motor_z\": \"Motor axis 0\",\n        \"motor_y\": \"Motor axis 1\",\n        \"motor_x\": \"Motor axis 2\",\n    },\n    inplace=True,\n)\ndf = df[[\"tomo_id\", \"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"]]\nlogger.info(\"SUBMISSION DATAFRAME:\\n%s\", df)\n\nif MODE == \"KAGGLE\":\n    logger.warning(f\"CLEARING ALL CONTENTS IN %s\", WORKING_DIR)\n    clear_dir_content(WORKING_DIR)\n\nsubmission_csv_path = os.path.join(WORKING_DIR, \"submission.csv\")\ndf.to_csv(submission_csv_path, index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T00:33:23.804791Z","iopub.execute_input":"2025-06-04T00:33:23.805088Z","iopub.status.idle":"2025-06-04T00:33:23.838735Z","shell.execute_reply.started":"2025-06-04T00:33:23.805065Z","shell.execute_reply":"2025-06-04T00:33:23.837656Z"}},"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":"","metadata":{"trusted":true},"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":"","metadata":{"trusted":true},"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":"","metadata":{"trusted":true},"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":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}