Skip to content

centroid

sleap_nn.export.wrappers.centroid

Centroid ONNX wrapper.

Classes:

Name Description
CentroidONNXWrapper

ONNX-exportable wrapper for centroid models.

CentroidONNXWrapper

Bases: BaseExportWrapper

ONNX-exportable wrapper for centroid models.

Expects input images as uint8 tensors in [0, 255].

Methods:

Name Description
__init__

Initialize centroid ONNX wrapper.

forward

Run centroid inference and return fixed-size outputs.

Source code in sleap_nn/export/wrappers/centroid.py
class CentroidONNXWrapper(BaseExportWrapper):
    """ONNX-exportable wrapper for centroid models.

    Expects input images as uint8 tensors in [0, 255].
    """

    def __init__(
        self,
        model: nn.Module,
        max_instances: int = 20,
        output_stride: int = 2,
        input_scale: float = 1.0,
        peak_threshold: float = 0.2,
    ):
        """Initialize centroid ONNX wrapper.

        Args:
            model: Centroid detection model.
            max_instances: Maximum number of instances to detect.
            output_stride: Output stride for confidence maps.
            input_scale: Input scaling factor.
            peak_threshold: Minimum confidence for a peak to be considered valid.
        """
        super().__init__(model)
        self.max_instances = max_instances
        self.output_stride = output_stride
        self.input_scale = input_scale
        self.peak_threshold = peak_threshold

    def forward(self, image: torch.Tensor) -> Dict[str, torch.Tensor]:
        """Run centroid inference and return fixed-size outputs."""
        image = self._normalize_uint8(image)
        if self.input_scale != 1.0:
            height = int(image.shape[-2] * self.input_scale)
            width = int(image.shape[-1] * self.input_scale)
            image = F.interpolate(
                image, size=(height, width), mode="bilinear", align_corners=False
            )

        confmaps = self._extract_tensor(self.model(image), ["centroid", "confmap"])
        peaks, values, valid = self._find_topk_peaks(
            confmaps, self.max_instances, self.peak_threshold
        )
        peaks = peaks * (self.output_stride / self.input_scale)

        return {
            "centroids": peaks,
            "centroid_vals": values,
            "instance_valid": valid,
        }

__init__(model, max_instances=20, output_stride=2, input_scale=1.0, peak_threshold=0.2)

Initialize centroid ONNX wrapper.

Parameters:

Name Type Description Default
model Module

Centroid detection model.

required
max_instances int

Maximum number of instances to detect.

20
output_stride int

Output stride for confidence maps.

2
input_scale float

Input scaling factor.

1.0
peak_threshold float

Minimum confidence for a peak to be considered valid.

0.2
Source code in sleap_nn/export/wrappers/centroid.py
def __init__(
    self,
    model: nn.Module,
    max_instances: int = 20,
    output_stride: int = 2,
    input_scale: float = 1.0,
    peak_threshold: float = 0.2,
):
    """Initialize centroid ONNX wrapper.

    Args:
        model: Centroid detection model.
        max_instances: Maximum number of instances to detect.
        output_stride: Output stride for confidence maps.
        input_scale: Input scaling factor.
        peak_threshold: Minimum confidence for a peak to be considered valid.
    """
    super().__init__(model)
    self.max_instances = max_instances
    self.output_stride = output_stride
    self.input_scale = input_scale
    self.peak_threshold = peak_threshold

forward(image)

Run centroid inference and return fixed-size outputs.

Source code in sleap_nn/export/wrappers/centroid.py
def forward(self, image: torch.Tensor) -> Dict[str, torch.Tensor]:
    """Run centroid inference and return fixed-size outputs."""
    image = self._normalize_uint8(image)
    if self.input_scale != 1.0:
        height = int(image.shape[-2] * self.input_scale)
        width = int(image.shape[-1] * self.input_scale)
        image = F.interpolate(
            image, size=(height, width), mode="bilinear", align_corners=False
        )

    confmaps = self._extract_tensor(self.model(image), ["centroid", "confmap"])
    peaks, values, valid = self._find_topk_peaks(
        confmaps, self.max_instances, self.peak_threshold
    )
    peaks = peaks * (self.output_stride / self.input_scale)

    return {
        "centroids": peaks,
        "centroid_vals": values,
        "instance_valid": valid,
    }