{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"},{"sourceId":243550773,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%writefile ek_inference_script.py\n\nfrom pathlib import Path\n\nPREDICTION_PARAMS = {\n    \"depth_window_size\": 128,\n    \"spatial_window_size\": 256,\n    \"depth_window_step\": 64,\n    \"spatial_window_step\": 96,\n    \"depth_overlap\": None,\n    \"spatial_overlap\": None,\n    \"use_weighted_average\": True,\n    \"use_z_flip_tta\": False,\n    \"use_y_flip_tta\": False,\n    \"use_x_flip_tta\": False,\n    \"use_xyz_flip_tta\": False,\n    \"input_scale_factor\": 0.5,\n    \"num_workers\": 2,\n    \"batch_size\": 2,\n}\n\nPOSTPROCESSING_PARAMS = {\n    \"min_score_threshold\": 0.0,\n    \"iou_threshold\": 0.8,\n    \"use_centernet_nms\": True,\n    \"use_gaussian_smoothing\": False,\n    \"use_percentile_threshold\": False,\n}\n\nMODELS_DIR = Path(\"/kaggle/input/v20-avg-build-trt-engine\")\nTRT_MODEL_NAME = \"v20_avg_128x256_ctx.onnx\"\nTRT_CACHE_NAME = \"v20_avg_128x256\"\n\nimport cv2\nimport torch\nimport skimage\nimport numpy as np\nimport dataclasses\nimport math\nimport copy\nimport pandas as pd\nimport einops\nimport gc\nimport itertools\nimport torch.nn.functional as F\nimport multiprocessing as mp\nfrom fire import Fire\nfrom tqdm import tqdm\nfrom collections import defaultdict\nfrom torch import Tensor, nn\nfrom typing import List, Tuple, Union, Iterable, Optional, Callable, Dict, Any\nfrom torch.utils.data import Dataset, DataLoader\nfrom abc import ABC, abstractmethod\nfrom multiprocessing import Pool\nfrom functools import partial\nfrom onnx import TensorProto\n\n\ndef load_slices(image_files: List[Path]) -> List[np.ndarray]:\n    \"\"\"Load all slices from image files using OpenCV and stack them into a volume.\n\n    Args:\n        image_files: List of paths to image files\n\n    Returns:\n        np.ndarray: Stacked volume of shape (depth, height, width)\n    \"\"\"\n    images = [cv2.imread(str(f), cv2.IMREAD_UNCHANGED) for f in image_files]\n    return images\n\n\ndef resize_3d_volume_torch(\n    images: List[np.ndarray] | np.ndarray,\n    scale_factor: Tuple[float, float, float] | float | None,\n    output_shape: Tuple[int, int, int] | np.ndarray | None = None,\n    interpolation_mode=\"trilinear\",\n    device=\"cpu\",\n):\n    if isinstance(images, list):\n        images = np.stack(images, axis=0)\n\n    if output_shape is None:\n        scale_factor = as_tuple_of_3(scale_factor)\n        output_shape = map(int, np.asarray(images.shape) * np.asarray(scale_factor))\n\n    cast_input_fp16 = images.dtype == np.uint8\n\n    with torch.no_grad():\n        volume = torch.from_numpy(images).to(device)\n        if cast_input_fp16:\n            volume = volume.half()\n\n        d, h, w = output_shape\n        # print(\n        #     volume.shape,\n        #     volume.dtype,\n        #     volume.device,\n        #     (d, h, w),\n        #     interpolation_mode,\n        # )\n        volume = torch.nn.functional.interpolate(\n            volume[None, None, :, :, :],\n            size=(d, h, w),\n            mode=interpolation_mode,\n            align_corners=False if interpolation_mode == \"trilinear\" else None,\n        )\n        if cast_input_fp16:\n            volume = volume.to(torch.uint8)\n        volume = volume[0, 0, :, :, :].cpu().numpy()\n\n    return volume\n\n\ndef as_tuple_of_3(value) -> Tuple:\n    if isinstance(value, (int, float)):\n        result = value, value, value\n    else:\n        a, b, c = value\n        result = a, b, c\n\n    return result\n\n\ndef get_shape_of_volume(volume: np.ndarray | List[np.ndarray]) -> Tuple[int, int, int]:\n    if isinstance(volume, np.ndarray):\n        depth, height, width = volume.shape\n    else:\n        depth = len(volume)\n        height, width = volume[0].shape\n\n    return depth, height, width\n\n\ndef gaussian_blur_3d(x: Tensor, kernel_size: int, sigma: float):\n    # build gaussian kernel\n    kd, kh, kw = as_tuple_of_3(kernel_size)\n    z = torch.linspace(-(kd // 2), kd // 2, steps=kd)\n    y = torch.linspace(-(kh // 2), kh // 2, steps=kh)\n    x_ = torch.linspace(-(kw // 2), kw // 2, steps=kw)\n\n    zz, yy, xx = torch.meshgrid(z, y, x_, indexing=\"ij\")\n    kernel_3d = torch.exp(-(xx**2 + yy**2 + zz**2) / (2 * sigma**2))\n    kernel_3d /= kernel_3d.sum()\n    #  normalize\n    kernel_3d = kernel_3d.to(x.device).to(x.dtype)\n\n    C = x.shape[1]\n    kernel_3d = kernel_3d.view(1, 1, *kernel_3d.shape)\n    kernel_3d = kernel_3d.repeat(C, 1, 1, 1, 1)\n\n    # apply gaussian kernel\n    x = torch.nn.functional.conv3d(x, weight=kernel_3d, bias=None, padding=kernel_size // 2, stride=1, groups=C)\n    return x\n\n\ndef flip_volume(volume: np.ndarray | Tensor, 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    if torch.is_tensor(volume):\n        return volume.flip(dim)\n    else:\n        return np.flip(volume, axis=dim)\n\n\ndef flip_offsets(offsets, dim, offset_dim):\n    if torch.is_tensor(offsets):\n        offsets_flip = offsets.flip(dim).clone()\n        offsets_flip[offset_dim] *= -1  # Flip the z-offsets\n    else:\n        offsets_flip = np.flip(offsets, axis=dim).copy()\n        offsets_flip[offset_dim] *= -1\n    return offsets_flip\n\n\ndef z_flip_volume(volume):\n    return flip_volume(volume, -3)\n\n\ndef z_flip_offsets(offsets):\n    return flip_offsets(offsets, -3, 2)\n\n\ndef y_flip_volume(volume):\n    return flip_volume(volume, -2)\n\n\ndef y_flip_offsets(offsets):\n    return flip_offsets(offsets, -2, 1)\n\n\ndef x_flip_volume(volume):\n    return flip_volume(volume, -1)\n\n\ndef x_flip_offsets(offsets):\n    return flip_offsets(offsets, -1, 0)\n\n\ndef compute_tiles_1d(length: int, window_size: int, window_step: int):\n    \"\"\"\n    Compute the slices for a sliding window over a one dimension.\n    Last slice can go outside the length of the dimension.\n\n    Args:\n        length: Length of the dimension\n        window_size: Size of the window\n        window_step: Step size between consecutive windows\n    \"\"\"\n    if window_step > window_size:\n        raise ValueError(f\"{window_step=} must be less than or equal to {window_size=}\")\n    start = 0\n    while True:\n        next_start = start + window_step\n        finish = start + window_size\n        yield slice(start, finish)\n        if (next_start > length) or (finish >= length):\n            break\n        start = next_start\n\n\ndef compute_tiles_for_window_step(\n    volume_shape: Tuple[int, int, int] | np.ndarray,\n    window_size: Union[int, Tuple[int, int, int]],\n    window_step: Union[int, Tuple[int, int, int]],\n):\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 = as_tuple_of_3(window_size)\n    window_step = as_tuple_of_3(window_step)\n    z, y, x = volume_shape\n\n    z_slices = list(compute_tiles_1d(z, window_size[0], window_step[0]))\n    y_slices = list(compute_tiles_1d(y, window_size[1], window_step[1]))\n    x_slices = list(compute_tiles_1d(x, window_size[2], window_step[2]))\n    for z_slice, y_slice, x_slice in itertools.product(z_slices, y_slices, x_slices):\n        yield z_slice, y_slice, x_slice\n\n\ndef compute_tiles_with_overlap_1d(length: int, window_size: int, overlap: float, stride: int = 1) -> Iterable[slice]:\n    \"\"\"\n    Compute the slices for a sliding window over a one dimension.\n    Method will try to distribute tiles to meet the overlap requirements with respect to the window step being divisible by the stride:\n    - Start of each tile is a multiple of stride\n    - First tile starts at 0\n    - Last tile ends at length (Could go outside the length of the dimension to account stride)\n    - The distance between the start of two consecutive tiles is a multiple of stride\n    - The overlap between consecutive tiles is as close as possible to the desired overlap\n\n    Args:\n        length: Length of the dimension\n        window_size: Size of the window\n        overlap: The approximate overlap between consecutive windows (0 to 1)\n        stride: Minimum increment step between consecutive tiles\n\n    Yields:\n        slice: Slice objects representing the window positions\n    \"\"\"\n    # Input validation\n    if stride <= 0:\n        raise ValueError(\"stride must be > 0\")\n    if window_size <= 0:\n        raise ValueError(\"window_size must be > 0\")\n    if not (0 <= overlap < 1):\n        raise ValueError(\"overlap must be between [0; 1)\")\n\n    # Handle edge case where window size is larger than or equal to the length\n    if window_size >= length:\n        yield slice(0, length)\n        return\n\n    # 1) Compute the \"ideal\" step between window starts\n    ideal_step = window_size * (1 - overlap)\n\n    # 2) Round that to the nearest multiple of `stride`\n    k_floor = math.floor(ideal_step / stride)\n    k_ceil = math.ceil(ideal_step / stride)\n    candidates = []\n    if k_floor >= 1:\n        candidates.append(k_floor * stride)\n    candidates.append(k_ceil * stride)\n\n    # Pick the candidate whose actual overlap is closest to desired\n    def overlap_diff(step: int) -> float:\n        actual = (window_size - step) / window_size\n        return abs(actual - overlap)\n\n    window_step = min(candidates, key=overlap_diff)\n\n    # 3) Compute where the final window *must* start so it ends at `length`\n    last_start_raw = length - window_size\n    last_start = math.ceil(last_start_raw / stride) * stride\n\n    # 4) Walk from 0 up to that last_start in increments of window_step,\n    #    then always append the last_start as the final tile.\n    starts = [0]\n    while starts[-1] + window_step < last_start:\n        starts.append(starts[-1] + window_step)\n    if starts[-1] != last_start:\n        starts.append(last_start)\n\n    if len(starts) > 2:\n        starts = np.array(starts)\n        std0 = np.diff(starts, axis=0).std()\n        # Now try redistributing the starts of the tiles (except the first and last one)\n        # Specifically, we subtract the stride as long as starts[1] > starts[0] and std of diff keep decreasing\n        starts2 = starts.copy()\n        while True:\n            starts2[1:-1] -= stride\n            if not (starts2[1] > starts2[0]):\n                break\n\n            std1 = np.diff(starts2, axis=0).std()\n            if std1 < std0:\n                std0 = std1\n                starts = starts2.copy()\n\n    # 5) Yield slices, clamping the end of the final tile to exactly `length`\n    for s in starts:\n        e = s + window_size\n        if e > length:\n            e = length\n        yield slice(s, e)\n\n\ndef compute_tiles_with_overlap(\n    volume_shape: Tuple[int, int, int] | np.ndarray,\n    window_size: Union[int, Tuple[int, int, int]],\n    overlaps: Tuple[float, float, float] | float,\n    stride: int,\n) -> Iterable[Tuple[slice, slice, slice]]:\n    \"\"\"Compute the slices for a sliding window over a volume.\n    Method will try to distribute tiles to meep the overlap requirements with respect to the window step being divisible by the stride.\n    \"\"\"\n    window_size_z, window_size_y, window_size_x = as_tuple_of_3(window_size)\n    overlap_z, overlap_y, overlap_x = as_tuple_of_3(overlaps)\n    z, y, x = volume_shape\n\n    z_slices = compute_tiles_with_overlap_1d(length=z, window_size=window_size_z, overlap=overlap_z, stride=stride)\n    y_slices = compute_tiles_with_overlap_1d(length=y, window_size=window_size_y, overlap=overlap_y, stride=stride)\n    x_slices = compute_tiles_with_overlap_1d(length=x, window_size=window_size_x, overlap=overlap_x, stride=stride)\n    for z_slice, y_slice, x_slice in itertools.product(z_slices, y_slices, x_slices):\n        yield z_slice, y_slice, x_slice\n\n\ndef normalize_volume_div_255(volume: np.ndarray | List[np.ndarray], dtype=np.float16) -> np.ndarray:\n    if isinstance(volume, list):\n        volume = np.stack(volume, axis=0)\n    return np.divide(volume, 255, dtype=dtype)\n\n\ndef load_and_resize_normalize_smart_2d(image_files: List[Path], scale_factor: float, dtype: np.dtype = np.float16) -> np.ndarray:\n    \"\"\"Load only needed slices, resize each with OpenCV, stack & normalize.\n\n    Args:\n        image_files: List of paths to image files\n        scale_factor: Scale factor to resize the volume\n\n    Returns:\n        np.ndarray: Normalized and resized volume of shape (D, H, W) and dtype np.float16\n    \"\"\"\n    # Calculate which slices to load\n    n_slices = len(image_files)\n    new_n_slices = int(n_slices * scale_factor)\n    indices = np.linspace(0, n_slices - 1, new_n_slices, dtype=int)\n\n    # Load and resize selected slices\n    slices = []\n    for idx in indices:\n        img = cv2.imread(str(image_files[idx]), cv2.IMREAD_UNCHANGED)\n        h, w = img.shape\n        new_h, new_w = int(h * scale_factor), int(w * scale_factor)\n        img = cv2.resize(img, (new_w, new_h), interpolation=cv2.INTER_LINEAR)\n        slices.append(img)\n\n    # Stack and normalize\n    volume = np.stack(slices, axis=0)  # [D, H, W]\n    volume = normalize_volume_div_255(volume, dtype=dtype)\n\n    return volume\n\n\ndef pad_empty_tomos_with_dummy_items(submission: pd.DataFrame, solution: pd.DataFrame) -> pd.DataFrame:\n    \"\"\"\n    Pad solution with rows where tomo_id is missing in submission.\n    If not motor is detected a single row with -1 motor coordinates should be added.\n    \"\"\"\n    all_tomo_ids = solution[\"tomo_id\"].unique()\n    present_tomo_ids = submission[\"tomo_id\"].unique()\n    missing_tomo_ids = set(all_tomo_ids) - set(present_tomo_ids)\n    dummy_rows = []\n    for tomo_id in missing_tomo_ids:\n        dummy_rows.append(\n            {\n                \"tomo_id\": tomo_id,\n                \"score\": 0,\n                \"Motor axis 0\": -1,\n                \"Motor axis 1\": -1,\n                \"Motor axis 2\": -1,\n            }\n        )\n    dummy_rows = pd.DataFrame(dummy_rows)\n    return pd.concat([submission, dummy_rows]).sort_values(\"tomo_id\").reset_index(drop=True)\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\nclass TileDataset(Dataset):\n    def __init__(self, *, volume, window_size: Union[int, Tuple[int, int, int]], tiles, torch_dtype):\n        self.volume = volume\n        self.tiles = tiles\n        self.window_size = as_tuple_of_3(window_size)\n        self.torch_dtype = torch_dtype\n\n    @classmethod\n    def with_overlaps(\n        cls,\n        volume,\n        window_size: Union[int, Tuple[int, int, int]],\n        overlaps: Tuple[float, float, float] | float,\n        stride: int,\n        torch_dtype: torch.dtype,\n    ) -> \"TileDataset\":\n        \"\"\"Creates a TileDataset by specifying number of tiles per dimension.\n\n        Args:\n            volume: Input volume of shape (D, H, W)\n            window_size: Size of sliding window as int or (depth, height, width)\n            overlaps: Overlap between consecutive tiles as float or (depth_overlap, height_overlap, width_overlap)\n            stride: Minimum increment step between consecutive tiles\n            torch_dtype: Torch dtype for the output tensors\n\n        Returns:\n            TileDataset configured with specified number of tiles per dimension\n        \"\"\"\n        tiles = list(compute_tiles_with_overlap(volume.shape, window_size=window_size, overlaps=overlaps, stride=stride))\n        return cls(volume=volume, window_size=window_size, tiles=tiles, torch_dtype=torch_dtype)\n\n    @classmethod\n    def with_window_step(\n        cls,\n        volume,\n        window_size: Union[int, Tuple[int, int, int]],\n        window_step: Union[int, Tuple[int, int, int]],\n        torch_dtype: torch.dtype,\n    ) -> \"TileDataset\":\n        \"\"\"Creates a TileDataset by using minimal number of tiles with given window step.\n\n        Args:\n            volume: Input volume of shape (D, H, W)\n            window_size: Size of sliding window as int or (depth, height, width)\n            window_step: Step size for sliding window as int or (depth_step, height_step, width_step)\n            torch_dtype: Torch dtype for the output tensors\n\n        Returns:\n            TileDataset configured with minimal number of tiles for given window step\n        \"\"\"\n        tiles = list(compute_tiles_for_window_step(volume.shape, window_size, window_step))\n        print(\"Input volume shape:\", volume.shape, \"Window size:\", window_size, \"Window step:\", window_step, \"Tiles:\", len(tiles))\n        dataset = cls(volume=volume, window_size=window_size, tiles=tiles, torch_dtype=torch_dtype)\n        return dataset\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_offsets = (tile[0].start, tile[1].start, tile[2].start)\n\n        return torch.from_numpy(tile_volume).unsqueeze(0).to(self.torch_dtype), torch.tensor(tile_offsets).long()\n\n\n@dataclasses.dataclass\nclass PredictionParams:\n    \"\"\"Parameters for model prediction/inference.\n\n    Attributes:\n        depth_window_size: Size of sliding window in depth (z) dimension\n        spatial_window_size: Size of sliding window in spatial (x,y) dimensions\n        depth_overlap: Number of tiles to split volume into along depth dimension\n        spatial_overlap: Number of tiles to split volume into along spatial dimensions\n        use_weighted_average: Whether to use weighted averaging when merging overlapping predictions\n        use_z_flip_tta: Whether to use test-time augmentation by flipping along z-axis\n        use_y_flip_tta: Whether to use test-time augmentation by flipping along y-axis\n        use_x_flip_tta: Whether to use test-time augmentation by flipping along x-axis\n        input_scale_factor: Factor to scale input volume by before prediction\n        batch_size: Batch size for prediction, defaults to 1\n        num_workers: Number of worker processes for data loading, defaults to 0\n    \"\"\"\n\n    depth_window_size: int\n    spatial_window_size: int\n\n    depth_overlap: float | None\n    spatial_overlap: float | None\n\n    depth_window_step: int | None\n    spatial_window_step: int | None\n\n    use_weighted_average: bool\n    use_z_flip_tta: bool\n    use_y_flip_tta: bool\n    use_x_flip_tta: bool\n    use_xyz_flip_tta: bool\n\n    input_scale_factor: float\n    batch_size: int = 1\n    num_workers: int = 0\n    multiprocessing_method: str = \"dataloader\"\n\n    @property\n    def window_size(self) -> Tuple[int, int, int]:\n        return (self.depth_window_size, self.spatial_window_size, self.spatial_window_size)\n\n    @property\n    def window_step(self) -> Tuple[int, int, int]:\n        return (self.depth_window_step, self.spatial_window_step, self.spatial_window_step)\n\n\n@dataclasses.dataclass\nclass PostprocessingParams:\n    min_score_threshold: float\n    iou_threshold: float\n\n    use_centernet_nms: bool\n    use_gaussian_smoothing: bool\n\n    use_percentile_threshold: bool\n\n    class_sigmas: list = dataclasses.field(default_factory=lambda: [10.0])\n    pre_nms_top_k: int = 1\n    centernet_nms_kernel: Union[int, Tuple[int, int, int]] = 3\n    class_map_gaussian_smoothing_kernel: Union[int, Tuple[int, int, int]] = 3\n\n\n@dataclasses.dataclass\nclass AccumulatedObjectDetectionPredictionContainer:\n    scores: Tensor  # Shape: (num_classes, D, H, W)\n    offsets: Tensor  # Shape: (3, D, H, W)\n    counter: Tensor  # Shape: (D, H, W)\n    stride: int\n    window_size: Tuple[int, int, int]\n    use_weighted_average: bool\n    weight_tensor: Optional[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        stride: int,\n        device=\"cpu\",\n        dtype=torch.float32,\n        use_weighted_average: bool = False,\n    ):\n        d, h, w = shape\n\n        def _ceil_div(value, divisor):\n            return int(np.ceil(value / divisor))\n\n        return cls(\n            scores=torch.zeros((num_classes, _ceil_div(d, stride), _ceil_div(h, stride), _ceil_div(w, stride)), device=device, dtype=dtype),\n            offsets=torch.zeros((3, _ceil_div(d, stride), _ceil_div(h, stride), _ceil_div(w, stride)), device=device, dtype=dtype),\n            counter=torch.zeros(_ceil_div(d, stride), _ceil_div(h, stride), _ceil_div(w, stride), device=device, dtype=dtype),\n            stride=stride,\n            window_size=window_size,\n            use_weighted_average=use_weighted_average,\n        )\n\n    def __post_init__(self):\n        if self.use_weighted_average:\n            output_window_size = (\n                self.window_size[0] // self.stride,\n                self.window_size[1] // self.stride,\n                self.window_size[2] // self.stride,\n            )\n            self.weight_tensor = self.compute_weight_matrix(torch.zeros((1, *output_window_size), device=self.scores.device))\n\n    def __iadd__(self, other):\n        if self.stride != other.stride:\n            raise ValueError(\"Stride 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        self.scores += other.scores.to(self.scores.device)\n        self.offsets += other.offsets.to(self.offsets.device)\n        self.counter += other.counter.to(self.counter.device)\n\n        return self\n\n    def accumulate_batch(self, batch_scores: Tensor, batch_offsets: Tensor, batch_tile_coords: List[Tuple[int, int, int]]):\n        \"\"\"\n        Accumulate predictions from a batch of tiles.\n\n        Args:\n            batch_scores: Tensor of shape (B, num_classes, D, H, W)\n            batch_offsets: Tensor of shape (B, 3, D, H, W)\n            batch_tile_coords: List of (z, y, x) coordinates for each tile in batch\n        \"\"\"\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=batch_scores[i],\n                offsets=batch_offsets[i],\n                tile_coords_zyx=tile_coord,\n            )\n\n    def accumulate(self, scores: Tensor, offsets: Tensor, tile_coords_zyx: Tuple[int, int, int]):\n        \"\"\"\n        Accumulate predictions from a single tile.\n\n        Args:\n            scores: Tensor of shape (num_classes, D, H, W)\n            offsets: Tensor of shape (3, D, H, W)\n            tile_coords_zyx: (z, y, x) coordinates of the tile\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_tile_coords_zyx = tuple(c // self.stride for c in tile_coords_zyx)\n        roi = (\n            slice(strided_tile_coords_zyx[0], strided_tile_coords_zyx[0] + scores.shape[1]),\n            slice(strided_tile_coords_zyx[1], strided_tile_coords_zyx[1] + scores.shape[2]),\n            slice(strided_tile_coords_zyx[2], strided_tile_coords_zyx[2] + scores.shape[3]),\n        )\n\n        scores_view = self.scores[:, roi[0], roi[1], roi[2]]\n        offsets_view = self.offsets[:, roi[0], roi[1], roi[2]]\n        counter_view = self.counter[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_view = self.weight_tensor[: scores.shape[1], : scores.shape[2], : scores.shape[3]]\n        else:\n            weight_view = 1\n\n        counter_view += weight_view\n        scores_view += scores * weight_view\n        offsets_view += offsets * weight_view\n\n    @classmethod\n    def compute_weight_matrix(cls, scores_volume: Tensor, sigma=15) -> Tensor:\n        \"\"\"\n        Compute a weight matrix for weighted averaging of predictions.\n\n        Args:\n            scores_volume: Tensor of shape (C, D, H, W)\n            sigma: Standard deviation for Gaussian weighting\n\n        Returns:\n            Tensor of shape (D, H, W) containing weights\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            device=scores_volume.device,\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) -> Tuple[Tensor, Tensor]:\n        \"\"\"\n        Merge accumulated predictions by averaging.\n\n        Returns:\n            Tuple of (scores, offsets) tensors with averaged predictions\n        \"\"\"\n        c = self.counter.unsqueeze(0)\n        zero_mask = c.eq(0)\n        if zero_mask.any():\n            print(\"Warning: Zero mask found in counter tensor. This may indicate that no predictions were made in this region.\")\n\n        scores = self.scores / c\n        scores.masked_fill_(zero_mask, 0.0)\n\n        offsets = self.offsets / c\n        offsets.masked_fill_(zero_mask, 0.0)\n\n        return scores, offsets\n\n\ndef anchors_for_offsets_feature_map(offsets, stride: int):\n    \"\"\"\n    Create anchors for the given offsets feature map\n\n    :param offsets: Offsets feature map of shape (B, 3, D, H, W). Channels are in XYZ order.\n\n    \"\"\"\n    _, _, d, h, w = offsets.shape\n    z, y, x = torch.meshgrid(\n        torch.arange(d, device=offsets.device),\n        torch.arange(h, device=offsets.device),\n        torch.arange(w, 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 + 1e-7)\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 strides: Stride of the logits\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(scores, kernel: Union[int, Tuple[int, int, int]] = 3):\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    *,\n    scores: Tensor,\n    offsets: Tensor,\n    stride: int,\n    postprocess_hparams: PostprocessingParams,\n    use_single_label_per_anchor=True,\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 stride: Strides of the network\n    :param postprocess_hparams: Postprocessing parameters\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.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(postprocess_hparams.min_score_threshold, 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 postprocess_hparams.use_gaussian_smoothing:\n        scores = gaussian_blur_3d(scores.unsqueeze(0), kernel_size=3, sigma=1.0).squeeze(0)\n\n    if postprocess_hparams.use_centernet_nms:\n        scores = centernet_heatmap_nms(scores.unsqueeze(0), kernel=postprocess_hparams.centernet_nms_kernel).squeeze(0)\n\n    scores, centers, _ = decode_detections(scores.unsqueeze(0), offsets.unsqueeze(0), stride)\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(postprocess_hparams.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 postprocess_hparams.pre_nms_top_k is not None and len(class_scores) > postprocess_hparams.pre_nms_top_k:\n            class_scores, sort_idx = torch.topk(class_scores, postprocess_hparams.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        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 > postprocess_hparams.iou_threshold\n            suppressed |= high_iou_mask.to(suppressed.device)\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    return final_centers, final_labels, final_scores\n\n\n@dataclasses.dataclass\nclass AnnotatedSample:\n    tomo_id: str\n    images_dir: Path\n    volume_shape: np.ndarray  # (3) Z Y X\n    motor_coordinates_zyx: np.ndarray  # Shape (N, 3), Z Y X\n    angstroms_per_voxel: float\n    radius: np.ndarray\n\n    original_shape: np.ndarray\n    original_angstroms_per_voxel: float\n\n    scale_factor: float\n    transforms: Tuple = tuple()\n\n    @property\n    def motor_coordinates_xyz(self):\n        return self.motor_coordinates_zyx[:, [2, 1, 0]].copy()\n\n    @property\n    def labels(self):\n        \"\"\"\n        Dummy labels\n        :return:\n        \"\"\"\n        return np.zeros(len(self.motor_coordinates_zyx), dtype=np.int32)\n\n    def _apply_transforms(self, volume):\n        for name, arg in self.transforms:\n            if name == \"flip\":\n                volume = np.flip(volume, axis=arg)\n            elif name == \"rot90\":\n                volume = np.rot90(volume, k=arg, axes=(1, 2))\n            else:\n                return ValueError(f\"Unknown transform: {name}:{arg}\")\n        return volume\n\n    def get_volume(self, cache: bool = True, dtype=np.float32):\n        if not cache:\n            return self.get_uncached_volume(dtype=dtype)\n\n        # Available cache resolutions in angstroms per voxel\n        target_resolution = self.angstroms_per_voxel\n\n        # Find closest cached resolution\n        closest_resolution = min(CACHE_RESOLUTIONS, key=lambda x: abs(x - target_resolution))\n        # print(f\"Closest resolution for {self.tomo_id} is {closest_resolution:.3f}A/pix\")\n        cached_path = self.images_dir / f\"volume_{closest_resolution:.3f}A_per_px.npy\"\n        cached_file_exists = cached_path.exists()\n\n        if cache and cached_file_exists:\n            volume = np.load(cached_path)\n        else:\n            logger.warning(\n                f\"No cached volume ({cached_path}) found for {self.tomo_id} for resolution {closest_resolution}. Requested resolution {target_resolution} and scale {self.scale_factor}. Expected volume shape {self.volume_shape}\"\n            )\n            image_files = list(sorted(self.images_dir.glob(\"*.jpg\")))\n            if len(image_files) == 0:\n                raise ValueError(f\"No image files found in {self.images_dir}. Expected at least one .jpg file.\")\n\n            volume = load_slices(image_files)\n\n        # Resize to target shape if needed\n        if get_shape_of_volume(volume) != tuple(self.volume_shape):\n            # print(f\"Volume shape {volume.shape} to target shape {self.volume_shape}\")\n            volume = resize_3d_volume_torch(volume, scale_factor=None, output_shape=self.volume_shape)\n\n        volume = normalize_volume_div_255(volume, dtype=dtype)\n        return self._apply_transforms(volume)\n\n    def rot90(self, k: int):\n        k = k % 4\n        if k == 0:\n            return self\n\n        depth, height, width = self.volume_shape\n        motor_coordinates = self.motor_coordinates_zyx.copy()\n        transforms = ((\"rot90\", k),)\n\n        if k % 4 == 1:\n            motor_coordinates[:, [1, 2]] = motor_coordinates[:, [2, 1]].copy()\n            motor_coordinates[:, 1] = width - motor_coordinates[:, 1] - 1\n            rotated_shape = (depth, width, height)\n        elif k % 4 == 2:\n            motor_coordinates[:, 1] = height - motor_coordinates[:, 1] - 1\n            motor_coordinates[:, 2] = width - motor_coordinates[:, 2] - 1\n            rotated_shape = (depth, height, width)\n        elif k % 4 == 3:\n            motor_coordinates[:, [1, 2]] = motor_coordinates[:, [2, 1]].copy()\n            motor_coordinates[:, 2] = height - motor_coordinates[:, 2] - 1\n            rotated_shape = (depth, width, height)\n        else:\n            rotated_shape = (depth, height, width)\n\n        return AnnotatedSample(\n            tomo_id=self.tomo_id + f\"_rot90_{k}\",\n            volume_shape=np.array(rotated_shape, dtype=int),\n            motor_coordinates_zyx=motor_coordinates,\n            radius=self.radius.copy(),\n            angstroms_per_voxel=self.angstroms_per_voxel,\n            images_dir=self.images_dir,\n            transforms=self.transforms + transforms,\n            scale_factor=self.scale_factor,\n            original_angstroms_per_voxel=self.original_angstroms_per_voxel,\n            original_shape=self.original_shape,\n        )\n\n    def flip(self, flip_x, flip_y, flip_z):\n        if not flip_z and not flip_y and not flip_x:\n            return self\n\n        depth, height, width = self.volume_shape\n\n        motor_coordinates = self.motor_coordinates_zyx.copy()\n        transforms = []\n\n        if flip_z:\n            # volume = np.flip(volume, axis=0)\n            motor_coordinates[:, 0] = (depth - 1) - motor_coordinates[:, 0]\n            transforms += [(\"flip\", 0)]\n\n        if flip_y:\n            # volume = np.flip(volume, axis=1)\n            motor_coordinates[:, 1] = (height - 1) - motor_coordinates[:, 1]\n            transforms += [(\"flip\", 1)]\n\n        if flip_x:\n            # volume = np.flip(volume, axis=2)\n            motor_coordinates[:, 2] = (width - 1) - motor_coordinates[:, 2]\n            transforms += [(\"flip\", 2)]\n\n        tomo_id_name = self.tomo_id\n        if flip_x:\n            tomo_id_name += \"_flip_x\"\n        if flip_y:\n            tomo_id_name += \"_flip_y\"\n        if flip_z:\n            tomo_id_name += \"_flip_z\"\n\n        return AnnotatedSample(\n            tomo_id=tomo_id_name,\n            motor_coordinates_zyx=motor_coordinates,\n            radius=self.radius.copy(),\n            angstroms_per_voxel=self.angstroms_per_voxel,\n            scale_factor=self.scale_factor,\n            transforms=self.transforms + tuple(transforms),\n            volume_shape=self.volume_shape,\n            images_dir=self.images_dir,\n            original_angstroms_per_voxel=self.original_angstroms_per_voxel,\n            original_shape=self.original_shape,\n        )\n\n    def scale(self, scale_factor: float):\n        if scale_factor == 1:\n            return self\n        return AnnotatedSample(\n            tomo_id=self.tomo_id,\n            images_dir=self.images_dir,\n            motor_coordinates_zyx=self.motor_coordinates_zyx * scale_factor,\n            angstroms_per_voxel=self.angstroms_per_voxel / scale_factor,\n            volume_shape=(self.volume_shape * scale_factor).astype(int),\n            radius=self.radius * scale_factor,\n            scale_factor=self.scale_factor * scale_factor,\n            transforms=self.transforms,\n            original_angstroms_per_voxel=self.original_angstroms_per_voxel,\n            original_shape=self.original_shape,\n        )\n\n    def scale_for_resolution(self, angstroms_per_voxel: float):\n        \"\"\"\n        Scale the sample to a new resolution\n        :param angstroms_per_voxel: Target resolution in Angstroms per voxel\n        :return:\n        \"\"\"\n        src_angstroms_per_voxel = self.angstroms_per_voxel\n        scale_factor = src_angstroms_per_voxel / angstroms_per_voxel\n        new_angstroms_per_voxel = self.angstroms_per_voxel / scale_factor\n        if math.fabs(new_angstroms_per_voxel - angstroms_per_voxel) > 1e-5:\n            raise ValueError(\n                f\"Computed A/px {new_angstroms_per_voxel=} is not equal to {angstroms_per_voxel=}. Source tomo {self.tomo_id} has {self.angstroms_per_voxel=}\"\n            )\n\n        return AnnotatedSample(\n            tomo_id=f\"{self.tomo_id}_{angstroms_per_voxel:.1f}A_pix\",\n            images_dir=self.images_dir,\n            motor_coordinates_zyx=self.motor_coordinates_zyx * scale_factor,\n            angstroms_per_voxel=angstroms_per_voxel,\n            volume_shape=(self.volume_shape * scale_factor).astype(int),\n            radius=self.radius * scale_factor,\n            scale_factor=self.scale_factor * scale_factor,\n            transforms=self.transforms,\n            original_angstroms_per_voxel=self.original_angstroms_per_voxel,\n            original_shape=self.original_shape,\n        )\n\n    def build_cache(self, device=\"cpu\") -> None:\n        \"\"\"Build cache files for standard resolutions.\n\n        Creates cached .npy files for standard resolutions (6-18Å in 2Å steps).\n        Each cache file is named as volume_<resolution>A_per_px.npy.\n        The volumes are resized to match the target resolution while preserving the physical size.\n        Caches raw volumes without normalization.\n        \"\"\"\n        # Load raw volume at full resolution\n        image_files = list(sorted(self.images_dir.glob(\"*.jpg\")))\n        volume = load_slices(image_files)\n        volume = np.stack(volume, axis=0)\n        original_shape = np.array(volume.shape)\n\n        for target_resolution in CACHE_RESOLUTIONS:\n            # Calculate scale factor and target shape\n            # sample_scaled = self.scale_for_resolution(target_resolution)\n            scale = self.angstroms_per_voxel / target_resolution\n            target_shape = tuple(map(int, original_shape * scale))\n            if np.prod(target_shape) == 0:\n                logger.warning(f\"Skipping cache for {self.tomo_id} at {target_resolution:.3f}A/pix due to zero shape.\")\n                continue\n\n            # Resize volume to target resolution\n            cached_volume = resize_3d_volume_torch(volume, scale_factor=None, output_shape=target_shape, device=device)\n\n            # Save cache file\n            cache_path = self.images_dir / f\"volume_{target_resolution:.3f}A_per_px.npy\"\n            np.save(cache_path, cached_volume, allow_pickle=False)\n            # print(f\"Saved cache for {self.tomo_id} at {target_resolution:.3f}A/pix {cached_volume.shape=} to {cache_path}\")\n\n    def show(self):\n        import matplotlib.pyplot as plt\n\n        volume = self.get_uncached_volume()\n        num_motors = len(self.motor_coordinates_zyx)\n        fig, axes = plt.subplots(num_motors, 3, figsize=(15, 5 * num_motors))\n\n        for i, coord in enumerate(self.motor_coordinates_zyx):\n            z, y, x = map(int, coord)\n            r = self.radius[i]\n\n            xy_slice = volume[z, :, :]\n            xz_slice = volume[:, y, :]\n            yz_slice = volume[:, :, x]\n\n            slices = [xy_slice, xz_slice, yz_slice]\n\n            for j, (slice_img, title) in enumerate(zip(slices, [\"XY\", \"XZ\", \"YZ\"])):\n                ax = axes[i, j] if num_motors > 1 else axes[j]\n                ax.imshow(slice_img, cmap=\"gray\")\n                if j == 0:\n                    circle = plt.Circle((coord[2], coord[1]), radius=r, edgecolor=\"red\", facecolor=\"none\")\n                elif j == 1:\n                    circle = plt.Circle((coord[2], coord[0]), radius=r, edgecolor=\"red\", facecolor=\"none\")\n                else:\n                    circle = plt.Circle((coord[1], coord[0]), radius=r, edgecolor=\"red\", facecolor=\"none\")\n\n                ax.add_patch(circle)\n                ax.set_title([\"XY\", \"XZ\", \"YZ\"][j] + f\" projection at motor {i}\")\n                ax.axis(\"off\")\n\n        plt.tight_layout()\n        plt.title(f\"{self.tomo_id} with {num_motors} motors\")\n        plt.show()\n\n    def get_uncached_volume(self, dtype=np.float32):\n        image_files = list(sorted(self.images_dir.glob(\"*.jpg\")))\n        volume = load_and_resize_normalize_smart_2d(image_files, self.scale_factor, dtype=dtype)\n        return self._apply_transforms(volume)\n\n\n@dataclasses.dataclass\nclass BasicSample:\n    tomo_id: str\n    images_dir: Path\n    scale_factor: float = 1.0\n\n    def scale(self, scale_factor):\n        return BasicSample(\n            tomo_id=self.tomo_id,\n            images_dir=self.images_dir,\n            scale_factor=self.scale_factor * scale_factor,\n        )\n\n    def get_uncached_volume(self, dtype=np.float32):\n        image_files = list(sorted(self.images_dir.glob(\"*.jpg\")))\n        return load_and_resize_normalize_smart_2d(image_files, self.scale_factor, dtype=dtype)\n\n    def load_volume_shape(self) -> Tuple[int, int, int]:\n        image_files = list(sorted(self.images_dir.glob(\"*.jpg\")))\n\n        depth = len(image_files)\n        height, width = cv2.imread(str(image_files[0]), cv2.IMREAD_UNCHANGED).shape\n        return depth, height, width\n\n\ndef parse_test_samples(images_dir) -> List[BasicSample]:\n    images_dir = Path(images_dir)\n\n    # find all dirs\n    tomo_dirs = list(sorted(images_dir.glob(\"*\")))\n\n    samples = []\n    for tomo_dir in tomo_dirs:\n        samples.append(\n            BasicSample(\n                tomo_id=tomo_dir.name,\n                images_dir=tomo_dir,\n            )\n        )\n\n    return samples\n\n\nclass BaseModelRunner(ABC):\n    \"\"\"\n    Abstract base class for model runners.\n    All model runners should inherit from this class and implement the __call__ method.\n    \"\"\"\n\n    @abstractmethod\n    def __call__(self, volume: Union[Tensor, np.ndarray]) -> Tuple[Tensor, Tensor]:\n        \"\"\"\n        Perform inference on the input volume.\n\n        Args:\n            volume: Input volume tensor or numpy array of shape (B, C, D, H, W)\n\n        Returns:\n            Tuple containing:\n            - List of score tensors, one for each output stride\n            - List of offset tensors, one for each output stride\n        \"\"\"\n        pass\n\n    def __enter__(self):\n        return self\n\n    def __exit__(self, exc_type, exc_val, exc_tb):\n        \"\"\"\n        Clean up resources when exiting the context manager.\n        \"\"\"\n        pass\n\n\ndef load_trt_libraries():\n    from ctypes import cdll\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\nclass OnnxRuntimeRunner(BaseModelRunner):\n    \"\"\"\n    This class is a more runner that uses CUDAExecutionProvider of OnnxRuntime to execute a model.\n    It provides methods to load the model, perform inference, and manage the input and output tensors.\n    \"\"\"\n\n    def __init__(self, model_path: str | Path, torch_device: str):\n        \"\"\"\n        Initialize the OnnxRuntimeTensorRTModel with the path to the ONNX model.\n\n        Args:\n            model_path (str): The path to the ONNX model file.\n        \"\"\"\n        import onnxruntime as ort\n\n        sess_options = ort.SessionOptions()\n        sess_options.enable_cpu_mem_arena = False\n        sess_options.enable_mem_pattern = False\n        sess_options.enable_mem_reuse = False\n        sess_options.enable_profiling = False\n\n        device_id = int(torch_device.split(\":\")[-1]) if isinstance(torch_device, str) else torch_device.index\n\n        self.device_id = device_id\n        self.model_path = model_path\n        self.providers = [(\"CUDAExecutionProvider\", {\"device_id\": device_id})]\n        self.session = ort.InferenceSession(model_path, sess_options=sess_options)\n\n        self.input_name = self.session.get_inputs()[0].name\n        self.output_names = [o.name for o in self.session.get_outputs()]\n\n    def __enter__(self):\n        \"\"\"\n        Start the inference session.\n        \"\"\"\n        self.session.set_providers(self.providers)\n        return self\n\n    def __exit__(self, exc_type, exc_val, exc_tb):\n        \"\"\"\n        Clean up resources when exiting the context manager.\n        \"\"\"\n        # https://github.com/microsoft/onnxruntime/issues/17142\n        self.session.set_providers([])\n\n    def __call__(self, volume):\n        \"\"\"\n        Perform inference on the input data.\n\n        Args:\n            volume: The input data for the model.\n\n        Returns:\n            The output of the model.\n        \"\"\"\n        if torch.is_tensor(volume):\n            volume = volume.detach().cpu().numpy()\n\n        scores, offsets = self.session.run(self.output_names, {self.input_name: volume})\n        scores = torch.from_numpy(scores)\n        offsets = torch.from_numpy(offsets)\n        return scores, offsets\n\n\nclass JITRunner(BaseModelRunner):\n    \"\"\"\n    This class is a wrapper for the ONNX Runtime TensorRT model without IO binding.\n    \"\"\"\n\n    def __init__(self, model_path: str, torch_device, torch_dtype):\n        \"\"\"\n        Initialize the OnnxRuntimeTensorRTModel with the path to the ONNX model.\n\n        Args:\n            model_path (str): The path to the ONNX model file.\n            torch_device: The torch device to use for computation\n            torch_dtype: The torch dtype to use for computation\n        \"\"\"\n        self.model = torch.jit.load(model_path, map_location=torch_device)\n        self.torch_device = torch_device\n        self.torch_dtype = torch_dtype\n\n    def __call__(self, volume: Union[Tensor, np.ndarray]) -> Tuple[Tensor, Tensor]:\n        \"\"\"\n        Perform inference on the input data.\n\n        Args:\n            volume: The input data for the model.\n\n        Returns:\n            Tuple containing:\n            - List of score tensors, one for each output stride\n            - List of offset tensors, one for each output stride\n        \"\"\"\n        # Ensure volume is on CPU and in the right format for ONNX Runtime\n        volume = volume.to(device=self.torch_device, dtype=self.torch_dtype)\n        with torch.amp.autocast(device_type=\"cuda\", dtype=self.torch_dtype):\n            scores, offsets = self.model(volume)\n        return scores, offsets\n\n\nclass TensorRTRunner(BaseModelRunner):\n    \"\"\"\n    This class is a wrapper for the ONNX Runtime TensorRT model without IO binding.\n    \"\"\"\n\n    def __init__(self, model_path: str, torch_device, torch_dtype, volume_shape: Tuple[int, int, int], trt_cache_path=None):\n        \"\"\"\n        Initialize the OnnxRuntimeTensorRTModel with the path to the ONNX model.\n\n        Args:\n            model_path (str): The path to the ONNX model file.\n            torch_device: The torch device to use for computation\n            torch_dtype: The torch dtype to use for computation\n            volume_shape: The expected shape of the input volume\n        \"\"\"\n        import onnxruntime as ort\n\n        self.model_path = model_path\n        self.torch_device = torch_device\n        self.torch_dtype = torch_dtype\n        self.volume_shape = volume_shape\n        self.np_dtype = {\n            torch.float16: np.float16,\n            torch.float32: np.float32,\n        }[torch_dtype]\n\n        device_id = int(torch_device.split(\":\")[-1]) if isinstance(torch_device, str) else torch_device.index\n        trt_kwargs = {\n            \"device_id\": device_id,\n            \"trt_fp16_enable\": (torch_dtype == torch.float16),\n            \"trt_builder_optimization_level\": 5,\n            \"trt_max_workspace_size\": 15 * 1073741824,\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=str(trt_cache_path / \"trt_timing_cache\"),\n                trt_engine_cache_path=str(trt_cache_path / \"trt_engine_cache\"),\n            )\n\n        providers = [(\"TensorrtExecutionProvider\", trt_kwargs)]\n\n        self.device_id = device_id\n        self.session = ort.InferenceSession(model_path, providers=providers)\n        self.input_name = self.session.get_inputs()[0].name\n        self.output_names = [output.name for output in self.session.get_outputs()]\n        if len(self.output_names) != 2:\n            raise ValueError(f\"Expected 2 outputs, but got {len(self.output_names)}: {self.output_names}\")\n\n    def __call__(self, volume: Union[Tensor, np.ndarray]) -> Tuple[List[Tensor], List[Tensor]]:\n        \"\"\"\n        Perform inference on the input data.\n\n        Args:\n            volume: The input data for the model.\n\n        Returns:\n            Tuple containing:\n            - List of score tensors, one for each output stride\n            - List of offset tensors, one for each output stride\n        \"\"\"\n        if torch.is_tensor(volume):\n            volume = volume.detach().cpu().numpy()\n\n        # Ensure volume is on CPU and in the right format for ONNX Runtime\n        volume = volume.astype(self.np_dtype, copy=False)\n\n        # Run inference\n        outputs = self.session.run(self.output_names, {self.input_name: volume})\n        if len(outputs) != 2:\n            raise ValueError(f\"Expected 2 outputs, but got {len(outputs)}: {outputs}\")\n\n        # Convert outputs back to torch tensors on the desired device\n        scores = torch.from_numpy(outputs[0])\n        offsets = torch.from_numpy(outputs[1])\n        return scores, offsets\n\n\nclass TensorRTWithIOBindingRunner(BaseModelRunner):\n    \"\"\"\n    This class is a wrapper for the ONNX Runtime TensorRT model.\n    It provides methods to load the model, perform inference, and manage the input and output tensors.\n    \"\"\"\n\n    def __init__(\n        self,\n        model_path: str,\n        torch_device,\n        torch_dtype,\n        volume_shape: Tuple[int, int, int],\n        output_stride: int,\n        trt_cache_path=None,\n    ):\n        \"\"\"\n        Initialize the OnnxRuntimeTensorRTModel with the path to the ONNX model.\n\n        Args:\n            model_path (str): The path to the ONNX model file.\n            torch_device: The torch device to use for computation\n            torch_dtype: The torch dtype to use for computation\n            volume_shape: The expected shape of the input volume\n        \"\"\"\n        import onnxruntime as ort\n\n        self.model_path = model_path\n        self.torch_device = torch_device\n        self.torch_dtype = torch_dtype\n        self.volume_shape = volume_shape\n        self.np_dtype = {torch.float16: np.float16, torch.float32: np.float32, torch.bfloat16: TensorProto.BFLOAT16}[torch_dtype]\n\n        device_id = int(torch_device.split(\":\")[-1]) if isinstance(torch_device, str) else torch_device.index\n        trt_kwargs = {\n            \"device_id\": device_id,\n            \"trt_fp16_enable\": (torch_dtype == torch.float16),\n            \"trt_builder_optimization_level\": 5,\n            \"trt_max_workspace_size\": 15 * 1073741824,\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=str(trt_cache_path / \"trt_timing_cache\"),\n                trt_engine_cache_path=str(trt_cache_path / \"trt_engine_cache\"),\n            )\n\n        providers = [(\"TensorrtExecutionProvider\", trt_kwargs)]\n\n        self.device_id = device_id\n        self.session = ort.InferenceSession(model_path, providers=providers)\n        self.input_name = self.session.get_inputs()[0].name\n        self.output_names = [output.name for output in self.session.get_outputs()]\n\n        outputs = self.session.get_outputs()\n        self.scores_shape = (\n            1,\n            1,\n            volume_shape[0] // output_stride,\n            volume_shape[1] // output_stride,\n            volume_shape[2] // output_stride,\n        )\n        self.offsets_shape = (\n            1,\n            3,\n            volume_shape[0] // output_stride,\n            volume_shape[1] // output_stride,\n            volume_shape[2] // output_stride,\n        )\n        self.scores_name = outputs[0].name\n        self.offsets_name = outputs[1].name\n\n        if len(self.output_names) != 2:\n            raise ValueError(f\"Expected 2 outputs, but got {len(self.output_names)}: {self.output_names}\")\n\n    def __call__(self, volume: Union[Tensor, np.ndarray]) -> Tuple[Tensor, Tensor]:\n        \"\"\"\n        Perform inference on the input data.\n\n        Args:\n            volume: The input data for the model.\n\n        Returns:\n            Tuple containing:\n            - List of score tensors, one for each output stride\n            - List of offset tensors, one for each output stride\n        \"\"\"\n        if isinstance(volume, np.ndarray):\n            volume = torch.from_numpy(volume)\n\n        volume: Tensor = volume.contiguous().to(device=self.torch_device, dtype=self.torch_dtype)\n\n        # Get output shapes and names from the model\n\n        # Create output tensors\n        scores = torch.empty(self.scores_shape, dtype=self.torch_dtype, device=self.torch_device)\n        offsets = torch.empty(self.offsets_shape, dtype=self.torch_dtype, device=self.torch_device)\n\n        io_binding = self.session.io_binding()\n        io_binding.bind_input(\n            name=self.input_name,\n            device_type=\"cuda\",\n            device_id=volume.device.index,\n            element_type=self.np_dtype,\n            shape=tuple(volume.shape),\n            buffer_ptr=volume.data_ptr(),\n        )\n\n        io_binding.bind_output(\n            self.scores_name,\n            device_type=\"cuda\",\n            device_id=volume.device.index,\n            element_type=self.np_dtype,\n            shape=tuple(scores.shape),\n            buffer_ptr=scores.data_ptr(),\n        )\n        io_binding.bind_output(\n            self.offsets_name,\n            device_type=\"cuda\",\n            device_id=volume.device.index,\n            element_type=self.np_dtype,\n            shape=tuple(offsets.shape),\n            buffer_ptr=offsets.data_ptr(),\n        )\n\n        self.session.run_with_iobinding(io_binding)\n\n        return scores, offsets\n\n\ndef get_train_df_for_evaluation(train_csv_path, num_samples=100):\n    df = pd.read_csv(train_csv_path)\n    df[\"Has motor\"] = df[\"Number of motors\"].apply(lambda x: 1 if x > 0 else 0)\n    one_or_zero_motor_mask = (df[\"Number of motors\"] == 1) | (df[\"Number of motors\"] == 0)\n    df = df[one_or_zero_motor_mask].reset_index(drop=True).head(num_samples)\n    return df\n\n\ndef load_sample(sample: BasicSample | AnnotatedSample, *, input_scale_factor, dtype):\n    return sample.tomo_id, sample.scale(input_scale_factor).get_uncached_volume(dtype=dtype)\n\n\ndef filler_bucket_assignment(costs: np.ndarray, num_buckets: int) -> np.ndarray:\n    order = np.argsort(-costs)\n    current_buckets_cost = np.zeros(num_buckets)\n    assignment = np.zeros_like(costs, dtype=int)\n    for element_index in order:\n        bucket_with_min_cost = np.argmin(current_buckets_cost)\n        assignment[element_index] = bucket_with_min_cost\n        current_buckets_cost[bucket_with_min_cost] += costs[element_index]\n\n    return assignment\n\n\ndef split_test_samples(samples: List[BasicSample], prediction_params: PredictionParams, local_rank: int, world_size: int, output_stride: int):\n    samples_shapes = np.array([sample.load_volume_shape() for sample in samples])  # [N, 3]\n    samples_shapes = (samples_shapes * prediction_params.input_scale_factor).astype(int)\n    window_size = prediction_params.window_size\n\n    if prediction_params.depth_window_step:\n        window_step = prediction_params.window_step\n        tiles_for_samples = [list(compute_tiles_for_window_step(shape, window_size, window_step)) for shape in samples_shapes]\n    elif prediction_params.depth_overlap is not None and prediction_params.spatial_overlap is not None:\n        tiles_for_samples = [\n            list(\n                compute_tiles_with_overlap(\n                    shape,\n                    window_size=window_size,\n                    overlaps=(\n                        prediction_params.depth_overlap,\n                        prediction_params.spatial_overlap,\n                        prediction_params.spatial_overlap,\n                    ),\n                    stride=output_stride,\n                )\n            )\n            for shape in samples_shapes\n        ]\n    else:\n        raise ValueError(\"Either depth_window_step or depth_overlap and spatial_overlap must be specified\")\n\n    # This is our cost\n    num_tiles = np.array([len(tiles) for tiles in tiles_for_samples])\n\n    print(\"Total tiles before balancing (naive assignment)\", sum(num_tiles[local_rank::world_size]), \"at rank\", local_rank)\n    assignment = filler_bucket_assignment(num_tiles, world_size)\n    this_rank_mask = assignment == local_rank\n\n    print(\"Total tiles after balancing\", sum(num_tiles[this_rank_mask]), \"at rank\", local_rank)\n    return [s for s, keep in zip(samples, this_rank_mask) if keep]\n\n\nclass LoaderDataset(Dataset):\n    def __init__(self, samples, input_scale_factor):\n        self.samples = samples\n        self.input_scale_factor = input_scale_factor\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        return load_sample(sample=sample, input_scale_factor=self.input_scale_factor, dtype=np.float16)\n\n\n@torch.no_grad()\ndef predict_scores_offsets_from_volume(\n    *,\n    volume: np.ndarray,\n    model_runner: BaseModelRunner,\n    output_stride: int,\n    tomo_id: str,\n    prediction_params: PredictionParams,\n    torch_dtype: Union[str, torch.dtype],\n    torch_device: Union[str, torch.device],\n):\n    def _predict_fn(inp):\n        return predict_scores_offsets_from_volume_no_tta(\n            volume=inp,\n            model_runner=model_runner,\n            output_stride=output_stride,\n            tomo_id=tomo_id,\n            prediction_params=prediction_params,\n            torch_dtype=torch_dtype,\n            torch_device=torch_device,\n        )\n\n    scores, offsets = _predict_fn(volume)\n\n    if prediction_params.use_xyz_flip_tta:\n        # Apply flips in X->Y->Z order to input\n        volume_flip = z_flip_volume(y_flip_volume(x_flip_volume(volume)))\n        scores_flip, offsets_flip = _predict_fn(volume_flip)\n        # Deaugment in reverse order (Z->Y->X) to get back to original coordinates\n        scores_flip = x_flip_volume(y_flip_volume(z_flip_volume(scores_flip)))\n        offsets_flip = x_flip_offsets(y_flip_offsets(z_flip_offsets(offsets_flip)))\n        scores += scores_flip\n        offsets += offsets_flip\n\n    if prediction_params.use_z_flip_tta:\n        scores_flip, offsets_flip = _predict_fn(z_flip_volume(volume))\n        scores += z_flip_volume(scores_flip)\n        offsets += z_flip_offsets(offsets_flip)\n\n    if prediction_params.use_y_flip_tta:\n        scores_flip, offsets_flip = _predict_fn(y_flip_volume(volume))\n        scores += y_flip_volume(scores_flip)\n        offsets += y_flip_offsets(offsets_flip)\n\n    if prediction_params.use_x_flip_tta:\n        scores_flip, offsets_flip = _predict_fn(x_flip_volume(volume))\n        scores += x_flip_volume(scores_flip)\n        offsets += x_flip_offsets(offsets_flip)\n\n    divisor = 1 + sum(\n        [\n            prediction_params.use_x_flip_tta,\n            prediction_params.use_y_flip_tta,\n            prediction_params.use_z_flip_tta,\n            prediction_params.use_xyz_flip_tta,\n        ]\n    )\n    if divisor > 1:\n        scores /= divisor\n        offsets /= divisor\n\n    return scores, offsets\n\n\n@torch.no_grad()\ndef predict_scores_offsets_from_volume_no_tta(\n    *,\n    volume: np.ndarray,\n    model_runner: BaseModelRunner,\n    output_stride: int,\n    tomo_id: str,\n    prediction_params: PredictionParams,\n    torch_dtype: Union[str, torch.dtype],\n    torch_device: Union[str, torch.device],\n):\n    torch.cuda.empty_cache()\n    container = None\n\n    if prediction_params.depth_window_step and prediction_params.spatial_window_step:\n        dataset = TileDataset.with_window_step(\n            volume=volume,\n            window_size=prediction_params.window_size,\n            window_step=prediction_params.window_step,\n            torch_dtype=torch_dtype,\n        )\n    elif prediction_params.depth_overlap is not None and prediction_params.spatial_overlap is not None:\n        dataset = TileDataset.with_overlaps(\n            volume,\n            window_size=prediction_params.window_size,\n            overlaps=(\n                prediction_params.depth_overlap,\n                prediction_params.spatial_overlap,\n                prediction_params.spatial_overlap,\n            ),\n            torch_dtype=torch_dtype,\n            stride=output_stride,\n        )\n    else:\n        raise ValueError(\"Either depth_window_step or depth_overlap and spatial_overlap must be specified\")\n\n    loader = DataLoader(\n        dataset,\n        batch_size=prediction_params.batch_size,\n        drop_last=False,\n        pin_memory=True,\n    )\n\n    with model_runner:\n        for tile_volume, tile_offsets in loader:\n            scores, offsets = model_runner(tile_volume)\n\n            if container is None:\n                num_classes = infer_num_classes_from_logits(scores)\n                container = AccumulatedObjectDetectionPredictionContainer.from_shape(\n                    shape=volume.shape,\n                    num_classes=num_classes,\n                    window_size=prediction_params.window_size,\n                    use_weighted_average=prediction_params.use_weighted_average,\n                    stride=output_stride,\n                    device=scores.device,\n                    dtype=scores.dtype,\n                )\n\n            container.accumulate_batch(scores, offsets, tile_offsets)\n\n        scores, offsets = container.merge_()\n    return scores, offsets\n\n\ndef postprocess_scores_offsets_into_submission(\n    *,\n    scores,\n    offsets,\n    output_stride: int,\n    tomo_id: str,\n    scale_predictions_to_input_scale: bool,\n    prediction_params: PredictionParams,\n    postprocess_hparams: PostprocessingParams,\n):\n    topk_coords_px, topk_clses, topk_scores = decode_detections_with_nms(\n        scores=scores,\n        offsets=offsets,\n        stride=output_stride,\n        postprocess_hparams=postprocess_hparams,\n    )\n    topk_scores = topk_scores.float().cpu().numpy()\n    top_coords = topk_coords_px.float().cpu().numpy()\n    topk_clses = topk_clses.cpu().numpy()\n\n    if scale_predictions_to_input_scale:\n        top_coords = top_coords / prediction_params.input_scale_factor\n\n    submission = {\n        \"tomo_id\": [],\n        \"score\": [],\n        \"Motor axis 0\": [],\n        \"Motor axis 1\": [],\n        \"Motor axis 2\": [],\n    }\n\n    if len(topk_clses) == 0:\n        submission[\"tomo_id\"].append(tomo_id)\n        submission[\"score\"].append(float(0))\n        submission[\"Motor axis 0\"].append(float(-1))\n        submission[\"Motor axis 1\"].append(float(-1))\n        submission[\"Motor axis 2\"].append(float(-1))\n    else:\n        for coord, score in zip(top_coords, topk_scores):\n            submission[\"tomo_id\"].append(tomo_id)\n            submission[\"score\"].append(float(score))\n            submission[\"Motor axis 0\"].append(float(coord[2]))\n            submission[\"Motor axis 1\"].append(float(coord[1]))\n            submission[\"Motor axis 2\"].append(float(coord[0]))\n    submission = pd.DataFrame.from_dict(submission)\n    return submission\n\n\ndef predict_volume(\n    *,\n    volume: np.ndarray,\n    model_runner: BaseModelRunner,\n    output_stride: int,\n    tomo_id: str,\n    prediction_params: PredictionParams,\n    postprocess_hparams: PostprocessingParams,\n    scale_predictions_to_input_scale: bool,\n    torch_device: Union[str, torch.device],\n    torch_dtype: Union[str, torch.dtype],\n    prediction_callback: Optional[Callable[[str, np.ndarray, np.ndarray], None]] = None,\n) -> pd.DataFrame:\n    scores, offsets = predict_scores_offsets_from_volume(\n        volume=volume,\n        model_runner=model_runner,\n        output_stride=output_stride,\n        tomo_id=tomo_id,\n        prediction_params=prediction_params,\n        torch_dtype=torch_dtype,\n        torch_device=torch_device,\n    )\n\n    if prediction_callback is not None:\n        prediction_callback(tomo_id, scores, offsets)\n\n    submission = postprocess_scores_offsets_into_submission(\n        scores=scores,\n        offsets=offsets,\n        output_stride=output_stride,\n        tomo_id=tomo_id,\n        scale_predictions_to_input_scale=scale_predictions_to_input_scale,\n        prediction_params=prediction_params,\n        postprocess_hparams=postprocess_hparams,\n    )\n    return submission\n\n\ndef predict_samples(\n    samples: List[BasicSample],\n    model_runner: BaseModelRunner,\n    output_stride: int,\n    prediction_params: PredictionParams,\n    postprocessing_params: PostprocessingParams,\n    torch_device=\"cuda\",\n    torch_dtype=torch.float16,\n    prediction_callback: Optional[Callable[[str, np.ndarray, np.ndarray], None]] = None,\n):\n    if postprocessing_params.use_percentile_threshold:\n        # Monkey patch postprocessing params to disable filtering\n        postprocessing_params = copy.deepcopy(postprocessing_params)\n        postprocessing_params.min_score_threshold = 0.0\n\n    # Run predictions\n    submissions = []\n    if prediction_params.multiprocessing_method == \"dataloader\":\n        loader = DataLoader(\n            LoaderDataset(samples, prediction_params.input_scale_factor),\n            batch_size=None,\n            num_workers=prediction_params.num_workers,\n            prefetch_factor=2 if prediction_params.num_workers > 0 else None,\n        )\n    elif prediction_params.multiprocessing_method == \"map\":\n        import multiprocessing as mp\n\n        load_sample_fn = partial(load_sample, input_scale_factor=prediction_params.input_scale_factor, dtype=np.float16)\n        loader = mp.Pool(processes=prediction_params.num_workers).imap_unordered(load_sample_fn, samples, chunksize=1)\n    else:\n        raise ValueError(\"prediction_params.multiprocessing_method must be one of 'dataloader' or 'map'\")\n\n    for tomo_id, volume in tqdm(loader, total=len(samples), desc=\"Predicting\"):\n        submission = predict_volume(\n            volume=volume,\n            model_runner=model_runner,\n            output_stride=output_stride,\n            tomo_id=tomo_id,\n            scale_predictions_to_input_scale=True,\n            prediction_params=prediction_params,\n            postprocess_hparams=postprocessing_params,\n            torch_device=torch_device,\n            torch_dtype=torch_dtype,\n            prediction_callback=prediction_callback,\n        )\n        submissions.append(submission)\n        del volume\n        gc.collect()\n\n    del loader\n\n    submission = pd.concat(submissions).sort_values(by=\"tomo_id\").reset_index(drop=True)\n    return submission\n\n\nDATA_DIR = Path(\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025\")\nTEST_IMAGES_DIR = DATA_DIR / \"test\"\nTRAIN_IMAGES_DIR = DATA_DIR / \"train\"\nTRAIN_CSV = DATA_DIR / \"train_labels.csv\"\n\n\ndef inference_main(device_id: int, world_size: int = 2, predict_train: bool = False):\n    torch.cuda.set_device(device_id)\n    torch_device = f\"cuda:{device_id}\"\n    torch_dtype = torch.float16\n\n    load_trt_libraries()\n\n    prediction_params = PredictionParams(**PREDICTION_PARAMS)\n    postprocessing_params = PostprocessingParams(**POSTPROCESSING_PARAMS)\n    volume_shape = (128, 256, 256)\n\n    if predict_train:\n        df = get_train_df_for_evaluation(TRAIN_CSV, 100)\n        train_samples = parse_test_samples(TRAIN_IMAGES_DIR)\n        train_samples = [sample for sample in train_samples if sample.tomo_id in df[\"tomo_id\"].values]\n        samples = split_test_samples(train_samples, prediction_params, device_id, world_size, output_stride=32)\n    else:\n        test_samples = parse_test_samples(TEST_IMAGES_DIR)\n        samples = split_test_samples(test_samples, prediction_params, device_id, world_size, output_stride=32)\n\n    model_runner = TensorRTRunner(\n        model_path=MODELS_DIR / TRT_MODEL_NAME,\n        trt_cache_path=MODELS_DIR / TRT_CACHE_NAME,\n        torch_device=torch_device,\n        torch_dtype=torch_dtype,\n        volume_shape=volume_shape,\n    )\n\n    submission = predict_samples(\n        samples,\n        model_runner,\n        output_stride=32,\n        prediction_params=prediction_params,\n        postprocessing_params=postprocessing_params,\n        torch_device=torch_device,\n        torch_dtype=torch_dtype,\n    )\n\n    submission.to_csv(f\"ek_submission_shard_{device_id}.csv\", index=False)\n\n\nif __name__ == \"__main__\":\n    Fire(inference_main)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T18:39:33.17442Z","iopub.execute_input":"2025-05-30T18:39:33.175199Z","iopub.status.idle":"2025-05-30T18:39:33.214232Z","shell.execute_reply.started":"2025-05-30T18:39:33.175175Z","shell.execute_reply":"2025-05-30T18:39:33.213219Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport cv2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T18:42:33.070331Z","iopub.execute_input":"2025-05-30T18:42:33.070991Z","iopub.status.idle":"2025-05-30T18:42:33.075642Z","shell.execute_reply.started":"2025-05-30T18:42:33.070965Z","shell.execute_reply":"2025-05-30T18:42:33.074658Z"}},"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\", \"ek_inference_script.py\", \"--device_id\", \"0\"], cwd=\"/kaggle/working/\")\nprocess2 = subprocess.Popen([\"python\", \"ek_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(\"Processes finished in\", elapsed, \"seconds\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T18:39:46.663812Z","iopub.execute_input":"2025-05-30T18:39:46.664101Z","iopub.status.idle":"2025-05-30T18:42:10.103258Z","shell.execute_reply.started":"2025-05-30T18:39:46.66408Z","shell.execute_reply":"2025-05-30T18:42:10.102603Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ek_submission_shard_0 = pd.read_csv(\"ek_submission_shard_0.csv\")\nek_submission_shard_1 = pd.read_csv(\"ek_submission_shard_1.csv\")\nsubmission = pd.concat([ek_submission_shard_0, ek_submission_shard_1], ignore_index=True)\nsubmission.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T18:42:20.312922Z","iopub.execute_input":"2025-05-30T18:42:20.313689Z","iopub.status.idle":"2025-05-30T18:42:20.323801Z","shell.execute_reply.started":"2025-05-30T18:42:20.313666Z","shell.execute_reply":"2025-05-30T18:42:20.322811Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Blending logic","metadata":{}},{"cell_type":"code","source":"import os\n\nif os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n    TH = submission['score'].quantile(0.55)\nelse:\n    TH = 0.5\n\nTH","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T19:05:43.330603Z","iopub.execute_input":"2025-05-30T19:05:43.330923Z","iopub.status.idle":"2025-05-30T19:05:43.336902Z","shell.execute_reply.started":"2025-05-30T19:05:43.330902Z","shell.execute_reply":"2025-05-30T19:05:43.336078Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mask = submission['score']<TH\nsubmission.loc[mask,['Motor axis 0','Motor axis 1','Motor axis 2']] = -1\nsubmission","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T19:05:50.75861Z","iopub.execute_input":"2025-05-30T19:05:50.759212Z","iopub.status.idle":"2025-05-30T19:05:50.772407Z","shell.execute_reply.started":"2025-05-30T19:05:50.759192Z","shell.execute_reply":"2025-05-30T19:05:50.7716Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission[['tomo_id','Motor axis 0','Motor axis 1','Motor axis 2']].to_csv('submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T19:05:57.997776Z","iopub.execute_input":"2025-05-30T19:05:57.998606Z","iopub.status.idle":"2025-05-30T19:05:58.007224Z","shell.execute_reply.started":"2025-05-30T19:05:57.998571Z","shell.execute_reply":"2025-05-30T19:05:58.006514Z"}},"outputs":[],"execution_count":null}]}