bottomup_multiclass
sleap_nn.export.wrappers.bottomup_multiclass
¶
ONNX wrapper for bottom-up multiclass (supervised ID) models.
Classes:
| Name | Description |
|---|---|
BottomUpMultiClassONNXWrapper |
ONNX-exportable wrapper for bottom-up multiclass (supervised ID) models. |
BottomUpMultiClassONNXWrapper
¶
Bases: BaseExportWrapper
ONNX-exportable wrapper for bottom-up multiclass (supervised ID) models.
This wrapper handles models that output both confidence maps for keypoint detection and class maps for identity classification. Unlike PAF-based bottom-up models, multiclass models use class maps to assign identity to each detected peak, then group peaks by identity.
The wrapper performs: 1. Peak detection in confidence maps (GPU) 2. Class probability sampling at peak locations (GPU) 3. Returns fixed-size tensors for CPU-side grouping
Expects input images as uint8 tensors in [0, 255].
Attributes:
| Name | Type | Description |
|---|---|---|
model |
The underlying PyTorch model. |
|
n_nodes |
Number of keypoint nodes in the skeleton. |
|
n_classes |
Number of identity classes. |
|
max_peaks_per_node |
Maximum number of peaks to detect per node. |
|
cms_output_stride |
Output stride of the confidence map head. |
|
class_maps_output_stride |
Output stride of the class maps head. |
|
input_scale |
Scale factor applied to input images before inference. |
Methods:
| Name | Description |
|---|---|
__init__ |
Initialize the wrapper. |
forward |
Run bottom-up multiclass inference. |
Source code in sleap_nn/export/wrappers/bottomup_multiclass.py
12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 | |
__init__(model, n_nodes, n_classes=2, max_peaks_per_node=20, cms_output_stride=4, class_maps_output_stride=8, input_scale=1.0, peak_threshold=0.2)
¶
Initialize the wrapper.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
model
|
Module
|
The underlying PyTorch model. |
required |
n_nodes
|
int
|
Number of keypoint nodes. |
required |
n_classes
|
int
|
Number of identity classes (e.g., 2 for male/female). |
2
|
max_peaks_per_node
|
int
|
Maximum peaks per node to detect. |
20
|
cms_output_stride
|
int
|
Output stride of confidence maps. |
4
|
class_maps_output_stride
|
int
|
Output stride of class maps. |
8
|
input_scale
|
float
|
Scale factor for input images. |
1.0
|
peak_threshold
|
float
|
Minimum confidence for a peak to be considered valid. |
0.2
|
Source code in sleap_nn/export/wrappers/bottomup_multiclass.py
forward(image)
¶
Run bottom-up multiclass inference.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
image
|
Tensor
|
Input image tensor of shape (batch, channels, height, width). Expected to be uint8 in [0, 255]. |
required |
Returns:
| Type | Description |
|---|---|
Dict[str, Tensor]
|
Dictionary with keys: - "peaks": Detected peak coordinates (batch, n_nodes, max_peaks, 2). Coordinates are in input image space (x, y). - "peak_vals": Peak confidence values (batch, n_nodes, max_peaks). - "peak_mask": Boolean mask for valid peaks (batch, n_nodes, max_peaks). - "class_probs": Class probabilities at each peak location (batch, n_nodes, max_peaks, n_classes). Postprocessing on CPU uses |