changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
22.6 kB
# SPDX-License-Identifier: Apache-2.0
"""C16 decode: the dense CenterHead decoder of Autoware's lidar_centerpoint (also PointPainting's
image_projection_based_fusion), and the mapping of decoded boxes to Autoware ``DetectedObject`` fields.
``decode_centerhead`` is ``generateBoxes3D_kernel`` + the ``is_score_keep`` compaction
(``autoware_lidar_centerpoint/lib/postprocess/postprocess_kernel.cu:52-178,220-226``), evaluated per BEV cell in
float32 as the kernel does (S:centerpoint:281-287):
- ``score_c = sigmoid(heatmap[c])``; ``label`` = first class with the largest score (strict ``>`` scan);
- ``x = (voxel_x * ds) * (xi + reg[0]) + min_x``, ``y`` likewise with ``yi``: **no +0.5 cell offset**;
- the distance bin is the first ``i`` with ``sqrt(x^2 + y^2) < upper[i]``; beyond the last bin the cell is dropped;
- the cell is dropped when ``score < thresholds[bin][label]`` or ``sqrt(rot0^2 + rot1^2) < yaw_norm[label]``, and a
zero score is never kept;
- box: ``z = height[0]`` (box centre), ``width = exp(dim[0])``, ``length = exp(dim[1])``, ``height = exp(dim[2])``
(the deployed ONNX regresses (w, l, h)), ``yaw = atan2(rot[0], rot[1])``, ``vel = vel[0:2]``.
No max-pool peak finding and no top-K: every cell over threshold is a candidate, as in Autoware (S:centerpoint:310).
Thresholds outside [0, 1) are coerced to 0 (``centerpoint_config.hpp:79-93``): :meth:`CenterHeadDecodeConfig.create`
applies that. The decoder works on head maps from any source (ONNX Runtime, the torch reference, a device readback).
Order: rows come out in cell order (``yi * W + xi``); :func:`sort_by_score` is the ``thrust::sort`` by descending
score, made deterministic (stable: ties keep the ascending cell order; thrust's sort is not stable).
``decode_cells`` evaluates the same decoder at given cells and drops none: it keeps each cell's gate inputs
(distance bin, score threshold, yaw norm) and outcomes, so two head-map sources (device and fp32 reference) can be
compared at the same cells and every cell one keeps and the other drops can be explained. ``head_values_at`` stores
the maps at a set of cells (compact agreement goldens).
``to_detected_objects`` is ``box3DToDetectedObject`` (``ros_utils.cpp:29-87``): the class name maps to an
``autoware_perception_msgs/ObjectClassification`` label (``getSemanticType``), ``yaw_ros = -yaw - pi/2`` (float),
dimensions (length, width, height), ``orientation_availability`` SIGN_UNKNOWN for car-like labels, and the twist in
the object frame when ``has_twist``.
numpy only; no side effects on import.
"""
from __future__ import annotations
import math
from dataclasses import dataclass, field, fields, replace
from typing import Any, Dict, List, Mapping, Optional, Sequence
import numpy as np
__all__ = [
"AUTOWARE_LABELS",
"LABEL_IDS",
"SIGN_UNKNOWN",
"UNAVAILABLE",
"semantic_label",
"is_car_like",
"CenterHeadDecodeConfig",
"DecodedBoxes",
"sigmoid_f32",
"decode_centerhead",
"sort_by_score",
"head_values_at",
"DecodedCells",
"decode_cells",
"DetectedObjects",
"to_detected_objects",
]
# autoware_perception_msgs/msg/ObjectClassification.msg label values (S:centerpoint:290)
AUTOWARE_LABELS = ("UNKNOWN", "CAR", "TRUCK", "BUS", "TRAILER", "MOTORCYCLE", "BICYCLE", "PEDESTRIAN")
LABEL_IDS = {name: i for i, name in enumerate(AUTOWARE_LABELS)}
# autoware_perception_msgs/msg/DetectedObjectKinematics.msg orientation_availability
UNAVAILABLE, SIGN_UNKNOWN = 0, 1
_SEMANTIC = {"CAR": "CAR", "TRUCK": "TRUCK", "BUS": "BUS", "TRAILER": "TRAILER", "BICYCLE": "BICYCLE",
"MOTORBIKE": "MOTORCYCLE", "PEDESTRIAN": "PEDESTRIAN"}
def semantic_label(class_name: str) -> int:
"""``getSemanticType`` (``ros_utils.cpp:89-108``): network class name -> ObjectClassification label; MOTORBIKE ->
MOTORCYCLE; anything else -> UNKNOWN."""
return LABEL_IDS[_SEMANTIC.get(class_name, "UNKNOWN")]
def is_car_like(label: Any) -> Any:
"""``object_recognition_utils::isCarLikeVehicle``: CAR, TRUCK, BUS, TRAILER."""
lab = np.asarray(label)
return np.isin(lab, [LABEL_IDS["CAR"], LABEL_IDS["TRUCK"], LABEL_IDS["BUS"], LABEL_IDS["TRAILER"]])
def _coerce01(values: Any) -> np.ndarray:
v = np.asarray(values, dtype=np.float32)
return np.where((v >= 0.0) & (v < 1.0), v, np.float32(0.0)).astype(np.float32)
@dataclass(frozen=True)
class CenterHeadDecodeConfig:
"""Decoder parameters. ``score_thresholds`` is (num_bins, num_classes), the node's ``[bin][class]`` layout
(``node.cpp:129-167``; read as ``thresholds[bin * C + label]``, ``postprocess_kernel.cu:114``)."""
class_names: tuple
voxel_size_xy: tuple
range_min_xy: tuple
downsample_factor: int
distance_bin_upper_limits: tuple
score_thresholds: np.ndarray
yaw_norm_thresholds: np.ndarray
has_variance: bool = False
has_twist: bool = False
@classmethod
def create(cls, *, class_names: Sequence[str], voxel_size_xy: Sequence[float], range_min_xy: Sequence[float],
downsample_factor: int, distance_bin_upper_limits: Sequence[float], score_thresholds: Any,
yaw_norm_thresholds: Sequence[float], has_variance: bool = False,
has_twist: bool = False) -> "CenterHeadDecodeConfig":
"""Validated config with Autoware's coercions: thresholds outside [0, 1) -> 0, ascending bin limits, one
yaw-norm threshold per class. ``score_thresholds`` may be a scalar, a per-class list, or (bins, classes)."""
names = tuple(class_names)
bins = tuple(float(v) for v in distance_bin_upper_limits)
if list(bins) != sorted(bins):
raise ValueError("distance_bin_upper_limits must be ascending (centerpoint_config.hpp:67-70)")
thr = np.asarray(score_thresholds, dtype=np.float32)
if thr.ndim == 0:
thr = np.full((len(bins), len(names)), thr, np.float32)
elif thr.ndim == 1 and thr.shape[0] == len(names):
thr = np.tile(thr[None, :], (len(bins), 1))
if thr.shape != (len(bins), len(names)):
raise ValueError(f"score_thresholds must be (bins={len(bins)}, classes={len(names)}), got {thr.shape}")
yaw = np.asarray(yaw_norm_thresholds, dtype=np.float64)
if yaw.shape != (len(names),):
raise ValueError("yaw_norm_thresholds needs one value per class (node.cpp:82-85)")
return cls(names, tuple(float(v) for v in voxel_size_xy), tuple(float(v) for v in range_min_xy),
int(downsample_factor), bins, _coerce01(thr), _coerce01(yaw), bool(has_variance), bool(has_twist))
def with_score_threshold(self, value: Optional[float]) -> "CenterHeadDecodeConfig":
"""The same config with every class / bin threshold set to ``value`` (coerced like Autoware); None keeps it."""
if value is None:
return self
return replace(self, score_thresholds=_coerce01(np.full_like(self.score_thresholds, value)))
@property
def num_classes(self) -> int:
return len(self.class_names)
@dataclass
class DecodedBoxes:
"""Decoded candidates (struct of arrays, all of length N). ``yaw`` is the network ("mmdet3d") yaw."""
cell: np.ndarray # int64 cell index yi * W + xi
label: np.ndarray # int32 class index into class_names
score: np.ndarray # float32 max sigmoid
x: np.ndarray # float32 box centre, model frame
y: np.ndarray
z: np.ndarray
length: np.ndarray # float32 exp(dim[1])
width: np.ndarray # float32 exp(dim[0])
height: np.ndarray # float32 exp(dim[2])
yaw: np.ndarray # float32 atan2(rot0, rot1)
vel_x: np.ndarray
vel_y: np.ndarray
def __len__(self) -> int:
return int(self.score.shape[0])
def take(self, idx: Any) -> "DecodedBoxes":
idx = np.asarray(idx, dtype=np.int64)
return DecodedBoxes(**{f.name: getattr(self, f.name)[idx] for f in fields(self)})
def to_dict(self) -> Dict[str, np.ndarray]:
return {f.name: getattr(self, f.name) for f in fields(self)}
def sigmoid_f32(x: Any) -> np.ndarray:
"""``1.0f / (1.0f + expf(-x))`` in float32 (``postprocess_kernel.cu:46-49``)."""
v = np.asarray(x, dtype=np.float32)
with np.errstate(over="ignore"):
return (np.float32(1.0) / (np.float32(1.0) + np.exp(-v))).astype(np.float32)
def decode_centerhead(heads: Mapping[str, Any], cfg: CenterHeadDecodeConfig) -> DecodedBoxes:
"""Decode the six head maps ``heatmap (C, H, W)``, ``reg (2, H, W)``, ``height (1, H, W)``, ``dim (3, H, W)``,
``rot (2, H, W)``, ``vel (2, H, W)`` (a leading batch axis of 1 is accepted) -> candidates in cell order."""
if cfg.has_variance:
raise NotImplementedError("variance heads (CenterPoint-sigma) are not supported by this decoder")
def get(name: str, channels: int) -> np.ndarray:
a = np.asarray(heads[name], dtype=np.float32)
if a.ndim == 4:
if a.shape[0] != 1:
raise ValueError(f"{name}: batch {a.shape[0]} != 1")
a = a[0]
if a.ndim != 3 or a.shape[0] < channels:
raise ValueError(f"{name}: expected ({channels}, H, W), got {a.shape}")
return a
hm = get("heatmap", cfg.num_classes)[:cfg.num_classes]
C, H, W = hm.shape
reg, hei, dim, rot = get("reg", 2), get("height", 1), get("dim", 3), get("rot", 2)
vel = get("vel", 2) if "vel" in heads else np.zeros((2, H, W), np.float32)
for name, a in (("reg", reg), ("height", hei), ("dim", dim), ("rot", rot), ("vel", vel)):
if a.shape[1:] != (H, W):
raise ValueError(f"{name} is {a.shape[1:]}, heatmap is {(H, W)}")
scores = sigmoid_f32(hm)
nan = np.isnan(scores)
if nan.any(): # NaN never wins the strict '>' scan; an all-NaN cell keeps label -1 and score 0
scores = np.where(nan, np.float32(-np.inf), scores)
label = np.argmax(scores, axis=0).astype(np.int32) # first max == strict '>' from -1
max_score = np.take_along_axis(scores, label[None].astype(np.int64), 0)[0]
found = np.isfinite(max_score)
max_score = np.where(found, max_score, np.float32(0.0)).astype(np.float32)
yi, xi = np.meshgrid(np.arange(H, dtype=np.float32), np.arange(W, dtype=np.float32), indexing="ij")
f32 = np.float32
sx = f32(f32(cfg.voxel_size_xy[0]) * f32(cfg.downsample_factor))
sy = f32(f32(cfg.voxel_size_xy[1]) * f32(cfg.downsample_factor))
x = (sx * (xi + reg[0]) + f32(cfg.range_min_xy[0])).astype(np.float32)
y = (sy * (yi + reg[1]) + f32(cfg.range_min_xy[1])).astype(np.float32)
radial = np.sqrt(x * x + y * y).astype(np.float32)
bucket = np.full((H, W), -1, dtype=np.int64)
for i in range(len(cfg.distance_bin_upper_limits) - 1, -1, -1): # first upper limit above the distance
bucket = np.where(radial < f32(cfg.distance_bin_upper_limits[i]), i, bucket)
thr = cfg.score_thresholds[np.clip(bucket, 0, None), np.clip(label, 0, None)]
yaw_norm = np.sqrt(rot[0] * rot[0] + rot[1] * rot[1]).astype(np.float32)
yaw_thr = cfg.yaw_norm_thresholds[np.clip(label, 0, None)]
keep = found & (bucket >= 0) & ~(max_score < thr) & (yaw_norm >= yaw_thr) & (max_score > 0.0)
cell = np.nonzero(keep.ravel())[0]
def at(a: np.ndarray) -> np.ndarray:
return a.reshape(-1)[cell]
return DecodedBoxes(
cell=cell.astype(np.int64), label=at(label).astype(np.int32), score=at(max_score), x=at(x), y=at(y),
z=at(hei[0]), length=np.exp(at(dim[1])).astype(np.float32), width=np.exp(at(dim[0])).astype(np.float32),
height=np.exp(at(dim[2])).astype(np.float32), yaw=np.arctan2(at(rot[0]), at(rot[1])).astype(np.float32),
vel_x=at(vel[0]), vel_y=at(vel[1]))
def sort_by_score(boxes: DecodedBoxes) -> DecodedBoxes:
"""Descending score, stable (ties keep cell order): a deterministic ``thrust::sort(..., score_greater())``."""
return boxes.take(np.argsort(-boxes.score, kind="stable"))
# ------------------------------------------------------------------ the decoder at given cells (agreement metrics)
def head_values_at(heads: Mapping[str, Any], cells: Any) -> Dict[str, np.ndarray]:
"""The channels of every head map ``(C, H, W)`` (a leading batch axis of 1 is accepted) at the BEV cells
``cells`` (``yi * W + xi``, any order) -> ``{name: (C, K) float32}``: an exact, compact form of the maps at the
cells that matter (agreement goldens). :func:`decode_cells` decodes it."""
idx = np.asarray(cells, dtype=np.int64).reshape(-1)
out: Dict[str, np.ndarray] = {}
for name, a in heads.items():
a = np.asarray(a, dtype=np.float32)
if a.ndim == 4:
if a.shape[0] != 1:
raise ValueError(f"{name}: batch {a.shape[0]} != 1")
a = a[0]
if a.ndim != 3:
raise ValueError(f"{name}: expected (C, H, W), got {a.shape}")
flat = a.reshape(a.shape[0], -1)
if idx.size and (idx.min() < 0 or idx.max() >= flat.shape[1]):
raise IndexError(f"{name}: cells outside [0, {flat.shape[1]})")
out[name] = np.ascontiguousarray(flat[:, idx])
return out
@dataclass
class DecodedCells:
""":func:`decode_centerhead` evaluated at given cells with **no cell dropped** (struct of arrays, length K). The
box fields are those of :class:`DecodedBoxes`, bit-identical to ``decode_centerhead``'s row for every cell it
keeps (``keep``); the gate inputs and outcomes are kept per cell, so a cell one implementation keeps and another
drops can be explained (score threshold vs yaw-norm gate)."""
cell: np.ndarray # int64 cell index yi * W + xi
label: np.ndarray # int32 first-max class; -1 when every class score is NaN
score: np.ndarray # float32 max sigmoid (0 when every class score is NaN)
x: np.ndarray
y: np.ndarray
z: np.ndarray
length: np.ndarray
width: np.ndarray
height: np.ndarray
yaw: np.ndarray # float32 network yaw atan2(rot0, rot1)
vel_x: np.ndarray
vel_y: np.ndarray
yaw_norm: np.ndarray # float32 sqrt(rot0^2 + rot1^2): the input of the yaw-norm gate
distance_bin: np.ndarray # int64 first bin with radial distance < its upper limit; -1 beyond the last
score_threshold: np.ndarray # float32 thresholds[bin][label] (bin and label clipped to 0, as the decoder)
yaw_norm_threshold: np.ndarray # float32 yaw-norm threshold of the label
passes_score: np.ndarray # bool: a label, inside a bin, not (score < threshold), score > 0
passes_yaw_norm: np.ndarray # bool: yaw_norm >= the label's yaw-norm threshold
@property
def keep(self) -> np.ndarray:
"""The cells ``decode_centerhead`` keeps (both gates pass)."""
return self.passes_score & self.passes_yaw_norm
def __len__(self) -> int:
return int(self.score.shape[0])
def take(self, idx: Any) -> "DecodedCells":
idx = np.asarray(idx, dtype=np.int64)
return DecodedCells(**{f.name: getattr(self, f.name)[idx] for f in fields(self)})
def to_boxes(self) -> DecodedBoxes:
"""The :class:`DecodedBoxes` fields of every cell (gates not applied: ``.take(np.nonzero(keep)[0])`` first
for the decoder's rows)."""
return DecodedBoxes(**{f.name: getattr(self, f.name) for f in fields(DecodedBoxes)})
def decode_cells(values: Mapping[str, Any], cells: Any, grid_w: int, cfg: CenterHeadDecodeConfig) -> DecodedCells:
"""The decoder at the cells ``cells`` of a head grid ``grid_w`` cells wide, from their channel values ``values``
(``{name: (C, K)}``, :func:`head_values_at`; ``vel`` optional), in the float32 arithmetic of
:func:`decode_centerhead` and with the gate outcomes instead of the drop (:class:`DecodedCells`)."""
if cfg.has_variance:
raise NotImplementedError("variance heads (CenterPoint-sigma) are not supported by this decoder")
idx = np.asarray(cells, dtype=np.int64).reshape(-1)
k = idx.size
def get(name: str, channels: int) -> np.ndarray:
a = np.asarray(values[name], dtype=np.float32)
if a.ndim != 2 or a.shape[0] < channels or a.shape[1] != k:
raise ValueError(f"{name}: expected ({channels}, {k}), got {a.shape}")
return a
hm = get("heatmap", cfg.num_classes)[:cfg.num_classes]
reg, hei, dim, rot = get("reg", 2), get("height", 1), get("dim", 3), get("rot", 2)
vel = get("vel", 2) if "vel" in values else np.zeros((2, k), np.float32)
scores = sigmoid_f32(hm)
nan = np.isnan(scores)
if nan.any():
scores = np.where(nan, np.float32(-np.inf), scores)
label = np.argmax(scores, axis=0).astype(np.int32) if k else np.zeros(0, np.int32)
max_score = np.take_along_axis(scores, label[None].astype(np.int64), 0)[0] if k else np.zeros(0, np.float32)
found = np.isfinite(max_score)
max_score = np.where(found, max_score, np.float32(0.0)).astype(np.float32)
f32 = np.float32
w = int(grid_w)
yi, xi = (idx // w).astype(np.float32), (idx % w).astype(np.float32)
sx = f32(f32(cfg.voxel_size_xy[0]) * f32(cfg.downsample_factor))
sy = f32(f32(cfg.voxel_size_xy[1]) * f32(cfg.downsample_factor))
x = (sx * (xi + reg[0]) + f32(cfg.range_min_xy[0])).astype(np.float32)
y = (sy * (yi + reg[1]) + f32(cfg.range_min_xy[1])).astype(np.float32)
radial = np.sqrt(x * x + y * y).astype(np.float32)
bucket = np.full(k, -1, dtype=np.int64)
for i in range(len(cfg.distance_bin_upper_limits) - 1, -1, -1):
bucket = np.where(radial < f32(cfg.distance_bin_upper_limits[i]), i, bucket)
thr = cfg.score_thresholds[np.clip(bucket, 0, None), np.clip(label, 0, None)].astype(np.float32)
yaw_norm = np.sqrt(rot[0] * rot[0] + rot[1] * rot[1]).astype(np.float32)
yaw_thr = cfg.yaw_norm_thresholds[np.clip(label, 0, None)].astype(np.float32)
return DecodedCells(
cell=idx.copy(), label=np.where(found, label, -1).astype(np.int32), score=max_score, x=x, y=y,
z=hei[0].copy(), length=np.exp(dim[1]).astype(np.float32), width=np.exp(dim[0]).astype(np.float32),
height=np.exp(dim[2]).astype(np.float32), yaw=np.arctan2(rot[0], rot[1]).astype(np.float32),
vel_x=vel[0].copy(), vel_y=vel[1].copy(), yaw_norm=yaw_norm, distance_bin=bucket, score_threshold=thr,
yaw_norm_threshold=yaw_thr, passes_score=found & (bucket >= 0) & ~(max_score < thr) & (max_score > 0.0),
passes_yaw_norm=yaw_norm >= yaw_thr)
@dataclass
class DetectedObjects:
"""Autoware ``DetectedObject`` fields of N objects (struct of arrays). Positions and dimensions hold the float32
values of the decoder (the message stores them as double); ``yaw`` is the ROS yaw (float32 arithmetic)."""
label: np.ndarray # uint8 ObjectClassification label
existence_probability: np.ndarray # float32
x: np.ndarray
y: np.ndarray
z: np.ndarray
yaw: np.ndarray
length: np.ndarray
width: np.ndarray
height: np.ndarray
orientation_availability: np.ndarray # uint8: SIGN_UNKNOWN for car-like labels, else UNAVAILABLE
twist_x: Optional[np.ndarray] = None # object-frame twist (has_twist only)
twist_y: Optional[np.ndarray] = None
source_index: Optional[np.ndarray] = None # row of the DecodedBoxes each object came from
label_before_remap: Optional[np.ndarray] = None
extra: Dict[str, np.ndarray] = field(default_factory=dict)
def __len__(self) -> int:
return int(self.existence_probability.shape[0])
def take(self, idx: Any) -> "DetectedObjects":
idx = np.asarray(idx, dtype=np.int64)
out = {}
for f in fields(self):
v = getattr(self, f.name)
if f.name == "extra":
out[f.name] = {k: a[idx] for k, a in v.items()}
else:
out[f.name] = None if v is None else v[idx]
return DetectedObjects(**out)
def boxes_xyzlwh_yaw(self) -> np.ndarray:
"""(N, 7) float32 x, y, z, length, width, height, yaw (the ``ttaw.outputs.Detections3D`` layout)."""
return np.stack([self.x, self.y, self.z, self.length, self.width, self.height, self.yaw],
axis=1).astype(np.float32)
def to_records(self, labels: Sequence[str] = AUTOWARE_LABELS) -> List[Dict[str, Any]]:
out = []
for i in range(len(self)):
d = {"label": labels[int(self.label[i])], "existence_probability": float(self.existence_probability[i]),
"x": float(self.x[i]), "y": float(self.y[i]), "z": float(self.z[i]), "yaw": float(self.yaw[i]),
"length": float(self.length[i]), "width": float(self.width[i]), "height": float(self.height[i]),
"orientation_availability": int(self.orientation_availability[i])}
if self.label_before_remap is not None and self.label_before_remap[i] != self.label[i]:
d["label_before_remap"] = labels[int(self.label_before_remap[i])]
out.append(d)
return out
def to_detected_objects(boxes: DecodedBoxes, class_names: Sequence[str], *, has_twist: bool = False) -> DetectedObjects:
"""``box3DToDetectedObject`` for every row (module docstring)."""
table = np.array([semantic_label(n) for n in class_names], dtype=np.uint8)
lab = np.asarray(boxes.label, dtype=np.int64)
valid = (lab >= 0) & (lab < len(class_names))
label = np.where(valid, table[np.clip(lab, 0, len(class_names) - 1)], LABEL_IDS["UNKNOWN"]).astype(np.uint8)
yaw = (-np.asarray(boxes.yaw, dtype=np.float32).astype(np.float64) - math.pi / 2.0).astype(np.float32)
orient = np.where(is_car_like(label), SIGN_UNKNOWN, UNAVAILABLE).astype(np.uint8)
tx = ty = None
if has_twist:
c, s = np.cos(yaw), np.sin(yaw) # float32, as std::cos(float)
tx = (c * boxes.vel_x + s * boxes.vel_y).astype(np.float32)
ty = (-s * boxes.vel_x + c * boxes.vel_y).astype(np.float32)
return DetectedObjects(label=label, existence_probability=np.asarray(boxes.score, np.float32),
x=boxes.x.copy(), y=boxes.y.copy(), z=boxes.z.copy(), yaw=yaw,
length=boxes.length.copy(), width=boxes.width.copy(), height=boxes.height.copy(),
orientation_availability=orient, twist_x=tx, twist_y=ty,
source_index=np.arange(len(boxes), dtype=np.int64), label_before_remap=label.copy(),
extra={"cell": np.asarray(boxes.cell, dtype=np.int64)})