{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":10599591,"sourceType":"datasetVersion","datasetId":6560827},{"sourceId":219682124,"sourceType":"kernelVersion"}],"dockerImageVersionId":30840,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#!pip install -q zarr einops onnx onnxconverter-common onnxruntime_gpu~=1.19.1 fire\n#nvidia-cudnn-cu12 nvidia-cuda-runtime-cu12 nvidia-cufft-cu12 nvidia-cublas-cu12 tensorrt-cu12==10.5 tensorrt-lean-cu12 tensorrt-dispatch-cu12 fire","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T15:02:49.356009Z","iopub.execute_input":"2025-01-29T15:02:49.356284Z","iopub.status.idle":"2025-01-29T15:03:20.012676Z","shell.execute_reply.started":"2025-01-29T15:02:49.356261Z","shell.execute_reply":"2025-01-29T15:03:20.011713Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!pip install -vv tensorrt-cu12==10.5","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T15:21:44.523229Z","iopub.execute_input":"2025-01-29T15:21:44.523676Z","iopub.status.idle":"2025-01-29T15:21:44.528103Z","shell.execute_reply.started":"2025-01-29T15:21:44.523642Z","shell.execute_reply":"2025-01-29T15:21:44.527002Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!pip install --extra-index-url https://pypi.nvidia.com tensorrt-cu12_libs==10.5 \n#!pip install --extra-index-url https://pypi.nvidia.com tensorrt-cu12_bindings==10.5 \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T15:21:45.349128Z","iopub.execute_input":"2025-01-29T15:21:45.349422Z","iopub.status.idle":"2025-01-29T15:21:45.353053Z","shell.execute_reply.started":"2025-01-29T15:21:45.349399Z","shell.execute_reply":"2025-01-29T15:21:45.352136Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorrt\nprint(tensorrt.__version__)\nassert tensorrt.Builder(tensorrt.Logger())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T15:04:24.722039Z","iopub.execute_input":"2025-01-29T15:04:24.722357Z","iopub.status.idle":"2025-01-29T15:04:27.367390Z","shell.execute_reply.started":"2025-01-29T15:04:24.722321Z","shell.execute_reply":"2025-01-29T15:04:27.366573Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from ctypes import *\n\ncudnn_libc = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/nvidia/cudnn/lib/libcudnn.so.9\")\ncublas_libc = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/nvidia/cublas/lib/libcublas.so.12\")\ncublaslt_libc = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/nvidia/cublas/lib/libcublasLt.so.12\")\ncudart_libc = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/nvidia/cuda_runtime/lib/libcudart.so.12\")\ncufft_libc = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/nvidia/cufft/lib/libcufft.so.11\")\ntrt = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/tensorrt_libs/libnvinfer.so.10\")\ntrt_libnvonnxparser = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/tensorrt_libs/libnvonnxparser.so.10\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T15:04:47.492147Z","iopub.execute_input":"2025-01-29T15:04:47.492500Z","iopub.status.idle":"2025-01-29T15:04:47.995305Z","shell.execute_reply.started":"2025-01-29T15:04:47.492462Z","shell.execute_reply":"2025-01-29T15:04:47.994530Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference script\n\nModified from JIT version https://www.kaggle.com/code/bloodaxe/v4-segresnet-dynunet but uses ONNXRuntime with TensorRT backend","metadata":{}},{"cell_type":"code","source":"%%writefile inference_script.py\n\nfrom ctypes import *\n\nimport os\nimport torch\nimport dataclasses\nimport zarr\nimport math\n\nfrom collections import defaultdict\nfrom pathlib import Path\nfrom fire import Fire\n\nimport numpy as np\nimport pandas as pd\nimport torch.jit\nfrom tqdm import tqdm\n\nimport typing\nimport einops\n\nfrom torch import Tensor, nn\nfrom typing import List, Tuple, Union, Any, Iterable, Optional\nfrom torch.utils.data import Dataset, DataLoader\nimport onnxruntime as ort\n\nANGSTROMS_IN_PIXEL = 10.012\n\nTARGET_CLASSES = (\n    {\n        \"name\": \"apo-ferritin\",\n        \"label\": 0,\n        \"color\": [0, 117, 255],\n        \"radius\": 60,\n        \"map_threshold\": 0.0418,\n    },\n    {\n        \"name\": \"beta-galactosidase\",\n        \"label\": 1,\n        \"color\": [176, 0, 192],\n        \"radius\": 90,\n        \"map_threshold\": 0.0578,\n    },\n    {\n        \"name\": \"ribosome\",\n        \"label\": 2,\n        \"color\": [0, 92, 49],\n        \"radius\": 150,\n        \"map_threshold\": 0.0374,\n    },\n    {\n        \"name\": \"thyroglobulin\",\n        \"label\": 3,\n        \"color\": [43, 255, 72],\n        \"radius\": 130,\n        \"map_threshold\": 0.0278,\n    },\n    {\n        \"name\": \"virus-like-particle\",\n        \"label\": 4,\n        \"color\": [255, 30, 53],\n        \"radius\": 135,\n        \"map_threshold\": 0.201,\n    },\n    {\"name\": \"beta-amylase\", \"label\": 5, \"color\": [153, 63, 0, 128], \"radius\": 65, \"map_threshold\": 0.035},\n)\n\nCLASS_LABEL_TO_CLASS_NAME = {c[\"label\"]: c[\"name\"] for c in TARGET_CLASSES}\nTARGET_SIGMAS = [c[\"radius\"] / ANGSTROMS_IN_PIXEL for c in TARGET_CLASSES]\n\n\ndef normalize_volume_to_unit_range(volume):\n    volume = volume - volume.min()\n    volume = volume / volume.max()\n    return volume.astype(np.float32)\n\n\ndef as_tuple_of_3(value) -> Tuple:\n    if isinstance(value, int):\n        result = value, value, value\n    else:\n        a, b, c = value\n        result = a, b, c\n\n    return result\n\n\ndef compute_better_tiles_1d(length: int, window_size: int, num_tiles: int):\n    \"\"\"\n    Compute the slices for a sliding window over a one dimension.\n    Method distribute tiles evenly over the length such that first tile is [0, window_size), and last tile is [length-window_size, length).\n    \"\"\"\n    last_tile_start = length - window_size\n\n    starts = np.linspace(0, last_tile_start, num_tiles, dtype=int)\n    ends = starts + window_size\n    for start, end in zip(starts, ends):\n        yield slice(start, end)\n\n\ndef compute_better_tiles_with_num_tiles(\n    volume_shape: Tuple[int, int, int],\n    window_size: Union[int, Tuple[int, int, int]],\n    num_tiles: Tuple[int, int, int],\n) -> Iterable[Tuple[slice, slice, slice]]:\n    \"\"\"Compute the slices for a sliding window over a volume.\n    A method can output a last slice that is smaller than the window size.\n    \"\"\"\n    window_size_z, window_size_y, window_size_x = as_tuple_of_3(window_size)\n    num_z_tiles, num_y_tiles, num_x_tiles = as_tuple_of_3(num_tiles)\n    z, y, x = volume_shape\n\n    for z_slice in compute_better_tiles_1d(z, window_size_z, num_z_tiles):\n        for y_slice in compute_better_tiles_1d(y, window_size_y, num_y_tiles):\n            for x_slice in compute_better_tiles_1d(x, window_size_x, num_x_tiles):\n                yield (\n                    z_slice,\n                    y_slice,\n                    x_slice,\n                )\n\n\nclass TileDataset(Dataset):\n    def __init__(\n        self, volume, window_size: Union[int, Tuple[int, int, int]], tiles_per_dim: Tuple[int, int, int], dtype, return_tensors\n    ):\n        self.volume = volume.astype(dtype)\n        self.tiles = list(compute_better_tiles_with_num_tiles(volume.shape, window_size, tiles_per_dim))\n        self.window_size = as_tuple_of_3(window_size)\n        self.dtype = dtype\n        self.return_tensors = return_tensors\n\n    def __len__(self):\n        return len(self.tiles)\n\n    def __getitem__(self, index):\n        tile = self.tiles[index]\n        tile_volume = self.volume[tile[0], tile[1], tile[2]]\n\n        pad_z = self.window_size[0] - tile_volume.shape[0]\n        pad_y = self.window_size[1] - tile_volume.shape[1]\n        pad_x = self.window_size[2] - tile_volume.shape[2]\n\n        tile_volume = np.pad(\n            tile_volume,\n            ((0, pad_z), (0, pad_y), (0, pad_x)),\n            mode=\"constant\",\n            constant_values=0,\n        )\n\n        tile_volume = tile_volume[None, :, :, :]  # Add channel dim\n        tile_offsets = np.array([tile[0].start, tile[1].start, tile[2].start], dtype=int)\n\n        if self.return_tensors == \"pt\":\n            tile_volume = torch.from_numpy(tile_volume)\n            tile_offsets = torch.from_numpy(tile_offsets).long()\n\n        return tile_volume, tile_offsets\n\n\ndef get_volume(\n    root_dir: str | Path,\n    study_name: str,\n    mode: str = \"denoised\",\n    split: str = \"train\",\n    voxel_spacing_str: str = \"VoxelSpacing10.000\",\n):\n    \"\"\"\n    Opens a Zarr store for the specified study and mode (e.g. denoised, isonetcorrected),\n    returns it as a NumPy array (fully loaded).\n\n    :param root_dir: Base directory (e.g., /path/to/czii-cryo-et-object-identification).\n    :param study_name: For example, \"TS_5_4\".\n    :param mode: Which volume mode to load, e.g. \"denoised\", \"isonetcorrected\", \"wbp\", etc.\n    :param split: \"train\" or \"test\".\n    :param voxel_spacing_str: Typically \"VoxelSpacing10.000\" from your structure.\n    :return: A 3D NumPy array of the volume data.\n    \"\"\"\n    # Example path:\n    #   /.../train/static/ExperimentRuns/TS_5_4/VoxelSpacing10.000/denoised.zarr\n    zarr_path = os.path.join(\n        str(root_dir),\n        split,\n        \"static\",\n        \"ExperimentRuns\",\n        study_name,\n        voxel_spacing_str,\n        f\"{mode}.zarr\",\n    )\n\n    # Open the top-level Zarr group\n    store = zarr.DirectoryStore(zarr_path)\n    zgroup = zarr.open(store, mode=\"r\")\n\n    #\n    # Typically, you'll see something like zgroup[0][0][0] or zgroup['0']['0']['0']\n    # for the actual volume data, but it depends on how your Zarr store is structured.\n    # Let’s assume the final data is at zgroup[0][0][0].\n    #\n    # You may need to inspect your actual Zarr structure and adjust accordingly.\n    #\n    volume = zgroup[0]  # read everything into memory\n\n    return np.asarray(volume)\n\n\n@dataclasses.dataclass\nclass AccumulatedObjectDetectionPredictionContainer:\n    scores: List[Tensor]\n    offsets: List[Tensor]\n    counter: List[Tensor]\n    strides: List[int]\n    window_size: Tuple[int, int, int]\n    use_weighted_average: bool\n    weight_tensor: List[Tensor] = None\n\n    @classmethod\n    def from_shape(\n        cls,\n        shape: Tuple[int, int, int],\n        window_size: Tuple[int, int, int],\n        num_classes: int,\n        strides: List[int],\n        device=\"cpu\",\n        dtype=torch.float32,\n        use_weighted_average: bool = False,\n    ):\n        d, h, w = shape\n\n        # fmt: off\n        return cls(\n            scores=[torch.zeros((num_classes, d // stride, h // stride, w // stride), device=device, dtype=dtype) for stride in strides],\n            offsets=[torch.zeros((3, d // stride, h // stride, w // stride), device=device, dtype=dtype) for stride in strides],\n            counter=[torch.zeros(d // stride, h // stride, w // stride, device=device, dtype=dtype) for stride in strides],\n            strides=list(strides),\n            window_size=window_size,\n            use_weighted_average=use_weighted_average,\n        )\n        # fmt: on\n\n    def __post_init__(self):\n        if self.use_weighted_average:\n            output_window_sizes = [\n                (self.window_size[0] // s, self.window_size[1] // s, self.window_size[2] // s) for s in self.strides\n            ]\n            self.weight_tensors = [\n                self.compute_weight_matrix(torch.zeros((1, *s), device=self.scores[0].device)) for s in output_window_sizes\n            ]\n\n    def __iadd__(self, other):\n        if self.strides != other.strides:\n            raise ValueError(\"Strides mismatch\")\n        if self.use_weighted_average != other.use_weighted_average:\n            raise ValueError(\"use_weighted_average mismatch\")\n        if self.window_size != other.window_size:\n            raise ValueError(\"Window size mismatch\")\n\n        for i in range(len(self.scores)):\n            self.scores[i] += other.scores[i].to(self.scores[i].device)\n            self.offsets[i] += other.offsets[i].to(self.offsets[i].device)\n            self.counter[i] += other.counter[i].to(self.counter[i].device)\n\n        return self\n\n    def accumulate_batch(self, batch_scores, batch_offsets, batch_tile_coords):\n        batch_size = len(batch_tile_coords)\n        for i in range(batch_size):\n            tile_coord = batch_tile_coords[i]\n            self.accumulate(\n                scores_list=[s[i] for s in batch_scores],\n                offsets_list=[o[i] for o in batch_offsets],\n                tile_coords_zyx=tile_coord,\n            )\n\n    def accumulate(self, scores_list: List[Tensor], offsets_list: List[Tensor], tile_coords_zyx):\n        if len(scores_list) != len(self.scores):\n            raise ValueError(\n                f\"Number of feature maps mismatch. Scores list has size {len(scores_list)}. Number of accumulator scores {len(self.scores)}\"\n            )\n        if not isinstance(scores_list, list):\n            raise ValueError(\"Scores should be a list of tensors\")\n        if not isinstance(offsets_list, list):\n            raise ValueError(\"Offsets should be a list of tensors\")\n\n        num_feature_maps = len(self.scores)\n\n        for i in range(num_feature_maps):\n            stride = self.strides[i]\n            scores = scores_list[i]\n            offsets = offsets_list[i]\n\n            if scores.ndim != 4 or offsets.ndim != 4:\n                raise ValueError(\"Scores and offsets should have shape (C, D, H, W)\")\n\n            strided_offsets_zyx = tuple(map(int, tile_coords_zyx // stride))\n\n            roi = (\n                slice(strided_offsets_zyx[0], strided_offsets_zyx[0] + scores.shape[1]),\n                slice(strided_offsets_zyx[1], strided_offsets_zyx[1] + scores.shape[2]),\n                slice(strided_offsets_zyx[2], strided_offsets_zyx[2] + scores.shape[3]),\n            )\n\n            scores_view = self.scores[i][:, roi[0], roi[1], roi[2]]\n            offsets_view = self.offsets[i][:, roi[0], roi[1], roi[2]]\n            counter_view = self.counter[i][roi[0], roi[1], roi[2]]\n\n            # Crop tile_scores to the view shape\n            scores = scores[:, : scores_view.shape[1], : scores_view.shape[2], : scores_view.shape[3]]\n            offsets = offsets[:, : offsets_view.shape[1], : offsets_view.shape[2], : offsets_view.shape[3]]\n\n            if self.use_weighted_average:\n                weight_matrix = self.weight_tensors[i]\n                weight_view = weight_matrix[\n                    : scores.shape[1], : scores.shape[2], : scores.shape[3]\n                ]  # Crop weight matrix to shape of predicted tensor\n            else:\n                weight_view = 1\n\n            counter_view += weight_view\n            scores_view += scores.to(scores_view.device) * weight_view\n            offsets_view += offsets.to(offsets_view.device) * weight_view\n\n    @classmethod\n    def compute_weight_matrix(self, scores_volume: Tensor, sigma=15):\n        \"\"\"\n        :param scores_volume: Tensor of shape (C, D, H, W)\n        :return: Tensor of shape (D, H, W)\n        \"\"\"\n        center = torch.tensor(\n            [\n                scores_volume.shape[1] / 2,\n                scores_volume.shape[2] / 2,\n                scores_volume.shape[3] / 2,\n            ]\n        )\n\n        i = torch.arange(scores_volume.shape[1], device=scores_volume.device)\n        j = torch.arange(scores_volume.shape[2], device=scores_volume.device)\n        k = torch.arange(scores_volume.shape[3], device=scores_volume.device)\n\n        I, J, K = torch.meshgrid(i, j, k, indexing=\"ij\")\n        distances = torch.sqrt((I - center[0]) ** 2 + (J - center[1]) ** 2 + (K - center[2]) ** 2)\n        weight = torch.exp(-distances / (sigma**2))\n\n        # I just like the look of heatmap\n        return weight**3\n\n    def merge_(self):\n        num_feature_maps = len(self.scores)\n        for i in range(num_feature_maps):\n            c = self.counter[i].unsqueeze(0)\n            zero_mask = c.eq(0)\n\n            self.scores[i] /= c\n            self.scores[i].masked_fill_(zero_mask, 0.0)\n\n            self.offsets[i] /= c\n            self.offsets[i].masked_fill_(zero_mask, 0.0)\n\n        return self.scores, self.offsets\n\n\ndef anchors_for_offsets_feature_map(offsets, stride):\n    z, y, x = torch.meshgrid(\n        torch.arange(offsets.size(-3), device=offsets.device),\n        torch.arange(offsets.size(-2), device=offsets.device),\n        torch.arange(offsets.size(-1), device=offsets.device),\n        indexing=\"ij\",\n    )\n    anchors = torch.stack([x, y, z], dim=0)\n    anchors = anchors.float().add_(0.5).mul_(stride)\n\n    anchors = anchors[None, ...].repeat(offsets.size(0), 1, 1, 1, 1)\n    return anchors\n\n\ndef keypoint_similarity(pts1, pts2, sigmas):\n    \"\"\"\n    Compute similarity between two sets of keypoints\n    :param pts1: ...x3\n    :param pts2: ...x3\n    \"\"\"\n    d = ((pts1 - pts2) ** 2).sum(dim=-1, keepdim=False)  # []\n    e: Tensor = d / (2 * sigmas**2)\n    iou = torch.exp(-e)\n    return iou\n\n\ndef decode_detections(logits: Tensor | List[Tensor], offsets: Tensor | List[Tensor], strides: int | List[int]):\n    \"\"\"\n    Decode detections from logits and offsets\n    :param logits: Predicted logits B C D H W\n    :param offsets: Predicted offsets B 3 D H W\n    :param anchors: Stride of the network\n\n    :return: Tuple of probas and centers:\n             probas - B N C\n             centers - B N 3\n\n    \"\"\"\n    if torch.is_tensor(logits):\n        logits = [logits]\n    if torch.is_tensor(offsets):\n        offsets = [offsets]\n    if isinstance(strides, int):\n        strides = [strides]\n\n    anchors = [anchors_for_offsets_feature_map(offset, s) for offset, s in zip(offsets, strides)]\n\n    logits_flat = []\n    centers_flat = []\n    anchors_flat = []\n\n    for logit, offset, anchor in zip(logits, offsets, anchors):\n        centers = anchor + offset\n\n        logits_flat.append(einops.rearrange(logit, \"B C D H W -> B (D H W) C\"))\n        centers_flat.append(einops.rearrange(centers, \"B C D H W -> B (D H W) C\"))\n        anchors_flat.append(einops.rearrange(anchor, \"B C D H W -> B (D H W) C\"))\n\n    logits_flat = torch.cat(logits_flat, dim=1)\n    centers_flat = torch.cat(centers_flat, dim=1)\n    anchors_flat = torch.cat(anchors_flat, dim=1)\n\n    return logits_flat, centers_flat, anchors_flat\n\n\ndef centernet_heatmap_nms(\n    scores,\n    kernel: Union[int, Tuple[int, int, int]] = 3,\n    # kernel: Union[int, Tuple[int, int, int]] = (5,3,3)\n):\n    kernel = as_tuple_of_3(kernel)\n    pad = (kernel[0] - 1) // 2, (kernel[1] - 1) // 2, (kernel[2] - 1) // 2\n\n    maxpool = torch.nn.functional.max_pool3d(scores, kernel_size=kernel, padding=pad, stride=1)\n\n    mask = scores == maxpool\n    peaks = scores * mask\n    return peaks\n\n\n@torch.no_grad()\ndef decode_detections_with_nms(\n    scores: List[Tensor],\n    offsets: List[Tensor],\n    strides: List[int],\n    min_score: Union[float, List[float]],\n    class_sigmas: List[float],\n    iou_threshold: float = 0.25,\n    use_single_label_per_anchor: bool = True,\n    use_centernet_nms: bool = False,\n    pre_nms_top_k: Optional[int] = None,\n) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:\n    \"\"\"\n    Decode detections from scores and centers with NMS\n\n    :param scores: Predicted scores of shape (C, D, H, W)\n    :param offsets: Predicted offsets of shape (3, D, H, W)\n    :param min_score: Minimum score to consider a detection\n    :param class_sigmas: Class sigmas (class radius for NMS), length = number of classes\n    :param iou_threshold: Threshold above which detections are suppressed\n\n    :return:\n        - final_centers [N, 3] (x, y, z)\n        - final_labels [N]\n        - final_scores [N]\n    \"\"\"\n\n    # Number of classes is the second dimension of `scores`\n    # e.g. scores shape = (C, D, H, W)\n    num_classes = scores[0].shape[0]  # the 'C' dimension\n\n    # Allow min_score to be a single value or a list of values\n    min_score = np.asarray(min_score, dtype=np.float32).reshape(-1)\n    if len(min_score) == 1:\n        min_score = np.full(num_classes, min_score[0], dtype=np.float32)\n\n    if use_centernet_nms:\n        scores = [centernet_heatmap_nms(s.unsqueeze(0)).squeeze(0) for s in scores]\n\n    scores, centers, _ = decode_detections([s.unsqueeze(0) for s in scores], [o.unsqueeze(0) for o in offsets], strides)\n    scores = scores.squeeze(0)\n    centers = centers.squeeze(0)\n\n    labels_of_max_score = scores.argmax(dim=1)\n\n    # Prepare final outputs\n    final_labels_list = []\n    final_scores_list = []\n    final_centers_list = []\n\n    # NMS per class\n    for class_index in range(num_classes):\n        sigma_value = float(class_sigmas[class_index])  # Get the sigma for this class\n        score_threshold = float(min_score[class_index])\n        score_mask = scores[:, class_index] >= score_threshold  # Filter out low-scoring detections\n\n        if use_single_label_per_anchor:\n            class_mask = labels_of_max_score.eq(class_index)  # Pick out only detections of this class\n            mask = class_mask & score_mask\n        else:\n            mask = score_mask\n\n        if not mask.any():\n            continue\n\n        class_scores = scores[mask, class_index]  # shape: [Nc]\n        class_centers = centers[mask]  # shape: [Nc, 3]\n\n        if pre_nms_top_k is not None and len(class_scores) > pre_nms_top_k:\n            class_scores, sort_idx = torch.topk(class_scores, pre_nms_top_k, largest=True, sorted=True)\n            class_centers = class_centers[sort_idx]\n        else:\n            class_scores, sort_idx = class_scores.sort(descending=True)\n            class_centers = class_centers[sort_idx]\n\n        # Run a simple “greedy” NMS\n        keep_indices = []\n        suppressed = torch.zeros_like(class_scores, dtype=torch.bool)  # track suppressed\n\n        # print(f\"Predictions for class {class_index}: \", torch.count_nonzero(class_mask).item())\n\n        for i in range(class_scores.size(0)):\n            if suppressed[i]:\n                continue\n            # Keep this detection\n            keep_indices.append(i)\n\n            # Suppress detections whose IoU with i is above threshold\n            iou = keypoint_similarity(class_centers[i : i + 1, :], class_centers, sigma_value)\n\n            high_iou_mask = iou > iou_threshold\n            suppressed |= high_iou_mask.to(suppressed.device)\n\n        print(f\"Predictions for class {class_index} after NMS\", len(keep_indices))\n\n        # Gather kept detections for this class\n        keep_indices = torch.as_tensor(keep_indices, dtype=torch.long, device=class_scores.device)\n        final_labels_list.append(torch.full((keep_indices.numel(),), class_index, dtype=torch.long))\n        final_scores_list.append(class_scores[keep_indices])\n        final_centers_list.append(class_centers[keep_indices])\n\n    # Concatenate from all classes\n    final_labels = torch.cat(final_labels_list, dim=0) if final_labels_list else torch.empty((0,), dtype=torch.long)\n    final_scores = torch.cat(final_scores_list, dim=0) if final_scores_list else torch.empty((0,))\n    final_centers = torch.cat(final_centers_list, dim=0) if final_centers_list else torch.empty((0, 3))\n\n    print(f\"Final predictions after NMS: {final_centers.size(0)}\")\n    return final_centers, final_labels, final_scores\n\n\ndef infer_num_classes_from_logits(logits):\n    if not torch.is_tensor(logits):\n        logits = logits[0]\n\n    b, c, d, h, w = logits.size()\n    return int(c)\n\n\ndef flip_volume(volume, dim):\n    \"\"\"\n    Flip the volume along the specified dimension.\n    :param volume: Volume to flip. B C D H W shape\n    :param dim: Dimension to flip\n    \"\"\"\n    return volume.flip(dim)\n\n\ndef flip_offsets(offsets, dim, offset_dim):\n    offsets_flip = torch.flip(offsets, [dim]).clone()\n    offsets_flip[:, offset_dim] *= -1  # Flip the z-offsets\n    return offsets_flip\n\n\ndef z_flip_volume(volume):\n    return flip_volume(volume, 2)\n\n\ndef z_flip_offsets(offsets):\n    return flip_offsets(offsets, 2, 2)\n\n\ndef y_flip_volume(volume):\n    return flip_volume(volume, 3)\n\n\ndef y_flip_offsets(offsets):\n    return flip_offsets(offsets, 3, 1)\n\n\ndef x_flip_volume(volume):\n    return flip_volume(volume, 4)\n\n\ndef x_flip_offsets(offsets):\n    return flip_offsets(offsets, 4, 0)\n\n\ndef ortvalue_to_torch_tensor(ortvalue):\n    # ._dlpack() is a Python API in onnxruntime >= 1.10\n    # If using older versions, use ortvalue.dlpack()\n    return torch.utils.dlpack.from_dlpack(ortvalue._dlpack())\n\n\n@torch.no_grad()\ndef predict_scores_offsets_from_volume_using_ort(\n    session: ort.InferenceSession,\n    volume: np.ndarray,\n    batch_size,\n    torch_device,\n    num_workers,\n    output_strides,\n    study_name,\n    torch_dtype,\n    use_weighted_average,\n    window_size: Tuple[int, int, int],\n    tiles_per_dim: Tuple[int, int, int],\n    use_z_flip_tta: bool,\n    use_y_flip_tta: bool,\n    use_x_flip_tta: bool,\n    device_id: int,\n):\n    container = None\n    volume = normalize_volume_to_unit_range(volume)\n    ds = TileDataset(volume, window_size, tiles_per_dim, dtype=np.float16, return_tensors=\"np\")\n\n    container = AccumulatedObjectDetectionPredictionContainer.from_shape(\n        shape=volume.shape,\n        num_classes=6,  # Hard-coded\n        window_size=window_size,\n        use_weighted_average=use_weighted_average,\n        strides=output_strides,\n        device=torch_device,\n        dtype=torch_dtype,\n    )\n\n    tile_scores = torch.zeros(\n        [batch_size, 6, window_size[0] // 2, window_size[1] // 2, window_size[2] // 2], dtype=torch_dtype, device=torch_device\n    )\n    tile_offsets = torch.zeros(\n        [batch_size, 3, window_size[0] // 2, window_size[1] // 2, window_size[2] // 2], dtype=torch_dtype, device=torch_device\n    )\n\n    for tile_index in tqdm(range(len(ds)), desc=f\"{study_name} {volume.shape}\"):\n        sample = ds[tile_index]\n        tile_volume, tile_coords = sample\n\n        tile_volume = torch.from_numpy(tile_volume[None, :, :, :]).contiguous().to(torch_device)  # Add batch dimension\n\n        io_binding = session.io_binding()\n        io_binding.bind_input(\n            name=\"volume\",\n            device_type=\"cuda\",  # \"cuda\" means input is on GPU\n            device_id=device_id,  # GPU device index\n            element_type=np.float16,  # dtype must match the model\n            shape=tuple(tile_volume.shape),\n            buffer_ptr=tile_volume.data_ptr(),  # pointer to GPU buffer\n        )\n\n        # (5) Bind outputs to CUDA so ORT will write them directly to GPU memory\n        io_binding.bind_output(\n            \"scores\",\n            device_type=\"cuda\",\n            device_id=device_id,  # GPU device index\n            element_type=np.float16,  # dtype must match the model\n            shape=tuple(tile_scores.shape),\n            buffer_ptr=tile_scores.data_ptr(),  # pointer to GPU buffer\n        )\n        io_binding.bind_output(\n            \"offsets\",\n            device_type=\"cuda\",\n            device_id=device_id,  # GPU device index\n            element_type=np.float16,  # dtype must match the model\n            shape=tuple(tile_offsets.shape),\n            buffer_ptr=tile_offsets.data_ptr(),  # pointer to GPU buffer\n        )\n\n        session.run_with_iobinding(io_binding)\n\n        probas = [tile_scores]\n        offsets = [tile_offsets]\n\n        container.accumulate_batch(probas, offsets, [tile_coords])\n\n    scores, offsets = container.merge_()\n    return scores, offsets\n\n\ndef postprocess_scores_offsets_into_submission(\n    scores,\n    offsets,\n    iou_threshold,\n    output_strides,\n    score_thresholds,\n    study_name,\n    use_centernet_nms,\n    use_single_label_per_anchor,\n    pre_nms_top_k: int,\n):\n    topk_coords_px, topk_clses, topk_scores = decode_detections_with_nms(\n        scores=scores,\n        offsets=offsets,\n        strides=output_strides,\n        class_sigmas=TARGET_SIGMAS,\n        min_score=score_thresholds,\n        iou_threshold=iou_threshold,\n        use_centernet_nms=use_centernet_nms,\n        use_single_label_per_anchor=use_single_label_per_anchor,\n        pre_nms_top_k=pre_nms_top_k,\n    )\n    topk_scores = topk_scores.float().cpu().numpy()\n    top_coords = topk_coords_px.float().cpu().numpy() * ANGSTROMS_IN_PIXEL\n    topk_clses = topk_clses.cpu().numpy()\n    submission = dict(\n        experiment=[],\n        particle_type=[],\n        score=[],\n        x=[],\n        y=[],\n        z=[],\n    )\n    for cls, coord, score in zip(topk_clses, top_coords, topk_scores):\n        submission[\"experiment\"].append(study_name)\n        submission[\"particle_type\"].append(CLASS_LABEL_TO_CLASS_NAME[int(cls)])\n        submission[\"score\"].append(float(score))\n        submission[\"x\"].append(float(coord[0]))\n        submission[\"y\"].append(float(coord[1]))\n        submission[\"z\"].append(float(coord[2]))\n    submission = pd.DataFrame.from_dict(submission)\n    return submission\n\n\n@torch.no_grad()\n@torch.jit.optimized_execution(False)\ndef predict_volume(\n    *,\n    volume: np.ndarray,\n    session: ort.InferenceSession,\n    output_strides: List[int],\n    window_size: Tuple[int, int, int],\n    tiles_per_dim: Tuple[int, int, int],\n    study_name: str,\n    score_thresholds: Union[float, List[float]],\n    iou_threshold,\n    batch_size,\n    num_workers,\n    use_weighted_average,\n    use_centernet_nms,\n    use_single_label_per_anchor,\n    device_id: int,\n    torch_device: str,\n    torch_dtype,\n    pre_nms_top_k,\n    use_z_flip_tta: bool,\n    use_y_flip_tta: bool,\n    use_x_flip_tta: bool,\n):\n    scores, offsets = predict_scores_offsets_from_volume_using_ort(\n        volume=volume,\n        session=session,\n        output_strides=output_strides,\n        window_size=window_size,\n        tiles_per_dim=tiles_per_dim,\n        batch_size=batch_size,\n        num_workers=num_workers,\n        torch_device=torch_device,\n        torch_dtype=torch_dtype,\n        study_name=study_name,\n        use_weighted_average=use_weighted_average,\n        use_z_flip_tta=use_z_flip_tta,\n        use_y_flip_tta=use_y_flip_tta,\n        use_x_flip_tta=use_x_flip_tta,\n        device_id=device_id,\n    )\n\n    submission = postprocess_scores_offsets_into_submission(\n        scores=scores,\n        offsets=offsets,\n        iou_threshold=iou_threshold,\n        output_strides=output_strides,\n        score_thresholds=score_thresholds,\n        study_name=study_name,\n        use_centernet_nms=use_centernet_nms,\n        use_single_label_per_anchor=use_single_label_per_anchor,\n        pre_nms_top_k=pre_nms_top_k,\n    )\n    return submission\n\n\ndef main_inference_entry_point(\n    *,\n    ensemble: str,\n    trt_cache_path: str,\n    score_thresholds,\n    device_id: int,\n    world_size: int = 2,\n    tiles_per_dim=(1, 9, 9),\n    output_strides: List[int] = (2,),\n    window_size=(192, 128, 128),\n    iou_threshold: float = 0.85,\n    use_weighted_average: bool = True,\n    use_centernet_nms: bool = True,\n    use_single_label_per_anchor: bool = False,\n    use_z_flip_tta: bool = False,\n    use_y_flip_tta: bool = False,\n    use_x_flip_tta: bool = False,\n    pre_nms_top_k: int = 16536,\n    batch_size=1,\n    num_workers=0,\n    torch_dtype=torch.float16,\n    data_path=\"/kaggle/input/czii-cryo-et-object-identification\",\n    split: str = \"test\",\n):\n    device_id = int(device_id)\n    torch_device = torch.device(f\"cuda:{device_id}\")  # Build torch device that matches device_id\n\n    trt_kwargs = {\n        \"device_id\": device_id,\n        \"trt_fp16_enable\": True,\n        \"trt_max_workspace_size\": 12 * 1073741824,\n        \"trt_builder_optimization_level\": 4,\n    }\n\n    if trt_cache_path is not None:\n        trt_kwargs.update(\n            trt_force_timing_cache=True,\n            trt_timing_cache_enable=True,\n            trt_engine_cache_enable=True,\n            trt_timing_cache_path=os.path.join(trt_cache_path, \"trt_timing_cache\"),\n            trt_engine_cache_path=os.path.join(trt_cache_path, \"trt_engine_cache\"),\n            # trt_ep_context_file_path = trt_cache_path,\n        )\n\n    sess_options = ort.SessionOptions()\n    session = ort.InferenceSession(\n        path_or_bytes=ensemble, providers=[(\"TensorrtExecutionProvider\", trt_kwargs)], sess_options=sess_options\n    )\n\n    path = Path(data_path)\n    studies_path = path / split / \"static\" / \"ExperimentRuns\"\n\n    studies = list(sorted(os.listdir(studies_path)))\n    studies = studies[device_id::world_size]  # Hopefully this is correct\n    print(\"Process got\", len(studies), \"to process\")\n\n    submissions = []\n\n    for study_name in tqdm(studies):\n        study_volume = get_volume(\n            root_dir=path,\n            study_name=study_name,\n            mode=\"denoised\",\n            split=split,\n        )\n\n        study_sub = predict_volume(\n            session=session,\n            volume=study_volume,\n            study_name=study_name,\n            output_strides=output_strides,\n            window_size=window_size,\n            tiles_per_dim=tiles_per_dim,\n            use_weighted_average=use_weighted_average,\n            use_centernet_nms=use_centernet_nms,\n            use_single_label_per_anchor=use_single_label_per_anchor,\n            pre_nms_top_k=pre_nms_top_k,\n            use_z_flip_tta=use_z_flip_tta,\n            use_y_flip_tta=use_y_flip_tta,\n            use_x_flip_tta=use_x_flip_tta,\n            score_thresholds=score_thresholds,\n            iou_threshold=iou_threshold,\n            batch_size=batch_size,\n            num_workers=num_workers,\n            device_id=device_id,\n            torch_device=torch_device,\n            torch_dtype=torch_dtype,\n        )\n\n        submissions.append(study_sub)\n\n    submission = pd.concat(submissions)\n    return submission\n\n\ndef main(device_id):\n\n    cudnn_libc = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/nvidia/cudnn/lib/libcudnn.so.9\")\n    cublas_libc = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/nvidia/cublas/lib/libcublas.so.12\")\n    cublaslt_libc = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/nvidia/cublas/lib/libcublasLt.so.12\")\n    cudart_libc = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/nvidia/cuda_runtime/lib/libcudart.so.12\")\n    cufft_libc = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/nvidia/cufft/lib/libcufft.so.11\")\n    trt = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/tensorrt_libs/libnvinfer.so.10\")\n    trt_libnvonnxparser = cdll.LoadLibrary(\"/usr/local/lib/python3.10/dist-packages/tensorrt_libs/libnvonnxparser.so.10\")\n\n    import tensorrt\n\n    print(tensorrt.__version__)\n    assert tensorrt.Builder(tensorrt.Logger())\n\n    submission = main_inference_entry_point(\n        ensemble=\"/kaggle/input/compile-and-cache-ensembe/v4_segresnet_dynunet_ensemble_1x192x128x128_ctx.onnx\",\n        trt_cache_path=\"/kaggle/input/compile-and-cache-ensembe/v4_segresnet_dynunet_ensemble_1x192x128x128\",\n        tiles_per_dim=(1, 11, 11),\n        score_thresholds=[\n            0.255,\n            0.235,\n            0.16,\n            0.205,\n            0.225,\n            0.5,\n        ],  # LB: 784 V4 OOF Computed CV score: 0.8295528641195601 std: 0.01879723638715648\n        device_id=device_id,\n    )\n\n    submission.to_csv(f\"submission_shard_{device_id}.csv\", index=False)\n\n\nif __name__ == \"__main__\":\n    Fire(main)\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-01-29T15:04:51.439955Z","iopub.execute_input":"2025-01-29T15:04:51.440244Z","iopub.status.idle":"2025-01-29T15:04:51.451801Z","shell.execute_reply.started":"2025-01-29T15:04:51.440222Z","shell.execute_reply":"2025-01-29T15:04:51.450982Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# For debugigng\n# submission = main_inference_entry_point(\n#     ensemble=\"/kaggle/input/compile-and-cache-ensembe/v4_segresnet_dynunet_ensemble_1x192x128x128_ctx.onnx\",\n#     trt_cache_path=\"/kaggle/input/compile-and-cache-ensembe/v4_segresnet_dynunet_ensemble_1x192x128x128\",\n    \n#     #ensemble=\"/kaggle/input/compile-and-cache-ensembe/v4_segresnet_dynunet_ensemble_1x192x128x128.onnx\",\n#     #trt_cache_path=None,\n    \n#     #trt_cache_path=\"/kaggle/input/compile-and-cache-ensembe/v4_segresnet_dynunet\",\n    \n#     window_size=(192, 128, 128),\n#     tiles_per_dim=(1, 9, 9),\n#     score_thresholds=[0.255,0.235,0.16 ,0.205,0.225, 0.5], # LB: 784 V4 OOF Computed CV score: 0.8295528641195601 std: 0.01879723638715648\n#     device_id=0,\n#     world_size=1,\n# )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T15:04:51.639193Z","iopub.execute_input":"2025-01-29T15:04:51.639420Z","iopub.status.idle":"2025-01-29T15:04:51.642769Z","shell.execute_reply.started":"2025-01-29T15:04:51.639401Z","shell.execute_reply":"2025-01-29T15:04:51.642030Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess\nimport pandas as pd\nimport cv2\n\nstart = cv2.getTickCount()\n\nprocess1 = subprocess.Popen([\"python\", \"inference_script.py\", \"--device_id\", \"0\"], cwd=\"/kaggle/working/\")\nprocess2 = subprocess.Popen([\"python\", \"inference_script.py\", \"--device_id\", \"1\"], cwd=\"/kaggle/working/\")\n\n# Wait for both processes to finish\nprocess1.wait()\nprocess2.wait()\n\nfinish = cv2.getTickCount()\nelapsed = (finish - start) / cv2.getTickFrequency()\n\nprint(\"2 studies processed in\", elapsed, \"seconds\", elapsed / 2, \"seconds/study\")\n\nestimated_submission_time_hours = ((250 * elapsed / 2) / 3600)\n\nprint(\"Estimated submission time\", estimated_submission_time_hours, \"h\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T15:04:54.586496Z","iopub.execute_input":"2025-01-29T15:04:54.586828Z","iopub.status.idle":"2025-01-29T15:08:12.474355Z","shell.execute_reply.started":"2025-01-29T15:04:54.586789Z","shell.execute_reply":"2025-01-29T15:08:12.473601Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission_shard_0 = pd.read_csv(\"submission_shard_0.csv\")\nsubmission_shard_0.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T15:08:12.475655Z","iopub.execute_input":"2025-01-29T15:08:12.476179Z","iopub.status.idle":"2025-01-29T15:08:12.500613Z","shell.execute_reply.started":"2025-01-29T15:08:12.476155Z","shell.execute_reply":"2025-01-29T15:08:12.499963Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission_shard_1 = pd.read_csv(\"submission_shard_1.csv\")\nsubmission_shard_1.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T15:08:12.501872Z","iopub.execute_input":"2025-01-29T15:08:12.502080Z","iopub.status.idle":"2025-01-29T15:08:12.514216Z","shell.execute_reply.started":"2025-01-29T15:08:12.502062Z","shell.execute_reply":"2025-01-29T15:08:12.513422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = pd.concat([submission_shard_0, submission_shard_1], ignore_index=True).drop(columns=[\"score\"])\n\nsubmission[\"id\"] = range(len(submission))\nsubmission.to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T15:08:12.515125Z","iopub.execute_input":"2025-01-29T15:08:12.515415Z","iopub.status.idle":"2025-01-29T15:08:12.538682Z","shell.execute_reply.started":"2025-01-29T15:08:12.515379Z","shell.execute_reply":"2025-01-29T15:08:12.537822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T15:08:12.539485Z","iopub.execute_input":"2025-01-29T15:08:12.539937Z","iopub.status.idle":"2025-01-29T15:08:12.548740Z","shell.execute_reply.started":"2025-01-29T15:08:12.539916Z","shell.execute_reply":"2025-01-29T15:08:12.548050Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}