changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
12.4 kB
# SPDX-License-Identifier: Apache-2.0
"""The encoder's in-graph pre-processing and the decoder's masks, computed on the host from normalised inputs.
Everything the v5.0 encoder graph does before its first matmul is tensor plumbing (slices, ``!= 0`` tests, OneHot,
``Atan`` with quadrant ``Where`` s): the TT port computes it here once per plan and uploads the results (PLAN.md 2.12;
SPEC 3.8, 8.2.6: padded lanes produce ``0/0`` headings that only a ``Where`` discards, which an arithmetic device
``where`` would propagate). The CPU reference consumes the same arrays, so the reference-vs-ONNX Runtime test also
proves this module against the graph (``/encoder/*/Where*``, ``/encoder/Concat_4`` / ``Concat_5`` taps).
Semantics (T4M ``model/module/encoder.py``; checked against the exported graph):
- ego: only the 6 OLDEST history rows are kept (rows 6..30 zeroed); the ego token is always valid and its position
feature is the last row of the truncated history, i.e. zeros;
- neighbours: rows 0..24 zeroed; a step is valid if any of dims 0..7 is non-zero, an agent if any step is; the type
one-hot (dims 8..10) and the position feature come from the last step; velocities (dims 4, 5) are zeroed AFTER the
validity test and a valid-step flag is appended (9 channels); invalid agents are zeroed;
- static objects: valid if any of the 10 values is non-zero (always zeros in Autoware, so never valid);
- lanes / route: dims 0..7 per point (zeroed for invalid lanes), attributes = dims 8..32 of point 0, speed limit and
its mask, position feature = point 10 with heading ``atan2(dy, dx)`` exported as ``Atan(dy / dx)`` plus quadrant
``Where`` s;
- polygons / line strings: ``[x, y, type one-hot..., dx, dy]`` with dx, dy = next point minus point (0 for the last
point), valid if any of the first FOUR columns is non-zero, position feature = point 20 / 10 with the NON-geometric
heading ``atan2(col3, col2)``: atan2(dx, is_intersection_area) for polygons, atan2(is_road_border, is_stop_line)
for line strings (SPEC 3.8: port the quirk exactly);
- goal / ego shape / turn indicators: always valid; positions (goal) and (0, 0, 1, 0); turn indicators drop the
current report (``[:, :-1]``, 30 values); the turn token reuses the ego-shape class id 8;
- fusion key mask: invalid tokens are masked with -inf, the ego key is forced valid (``encoder.py:833``);
- decoder: a neighbour is a valid attention key if its current state (``neighbor_agents_past[:, 30, :4]``) is not
all zero; current states = (ego_current_state[:4], neighbour current states) for the prefix constraint.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Dict, Mapping
import numpy as np
from ..reference import config as C
__all__ = ["EncoderFeatures", "DecoderMasks", "encoder_features", "decoder_masks", "atan2_onnx", "line_features"]
PI_F32 = np.float32(3.1415927) # the ONNX constant of the exported atan2 (``/encoder/lane_encoder/Constant_13``)
def atan2_onnx(y: np.ndarray, x: np.ndarray) -> np.ndarray:
"""``torch.atan2`` as exported to ONNX opset 20: ``a = Atan(y / x)``; ``Where(x < 0, Where(y > 0, a + pi, a - pi),
a)`` (float32). Differs from IEEE atan2 only on signed zeros (x = -0, y = 0 -> -pi) and 0/0 (NaN, masked later)."""
y = np.asarray(y, np.float32)
x = np.asarray(x, np.float32)
with np.errstate(divide="ignore", invalid="ignore"):
a = np.arctan((y / x).astype(np.float32)).astype(np.float32)
alt = np.where(y > 0, a + PI_F32, a - PI_F32).astype(np.float32)
return np.where(x < 0, alt, a).astype(np.float32)
def _heading_pos(xy: np.ndarray, y_col: np.ndarray, x_col: np.ndarray) -> np.ndarray:
"""``[x, y, cos(h), sin(h)]`` with ``h = atan2_onnx(y_col, x_col)`` (float32)."""
h = atan2_onnx(y_col, x_col)
with np.errstate(invalid="ignore"):
return np.stack([xy[..., 0], xy[..., 1], np.cos(h), np.sin(h)], axis=-1).astype(np.float32)
def _onehot_pos(pos4: np.ndarray, cls: int) -> np.ndarray:
"""Append the 10-way class one-hot (``add_class_type``)."""
onehot = np.zeros(pos4.shape[:-1] + (C.POS_CLASS_NUM,), np.float32)
onehot[..., cls] = 1.0
return np.concatenate([pos4.astype(np.float32), onehot], axis=-1)
def line_features(x: np.ndarray) -> np.ndarray:
"""``LineEncoder``: ``[points..., dx, dy]`` with dx / dy = next point minus point (0 at the last point)."""
x = np.asarray(x, np.float32)
d = np.zeros(x.shape[:-1] + (2,), np.float32)
d[..., :-1, 0] = x[..., 1:, 0] - x[..., :-1, 0]
d[..., :-1, 1] = x[..., 1:, 1] - x[..., :-1, 1]
return np.concatenate([x, d], axis=-1)
@dataclass
class EncoderFeatures:
"""Host-side inputs of the encoder network (batch 1, no batch dim). ``valid[cat]`` is True for valid entities;
``token_valid`` (564) gates the positional embedding, ``key_valid`` (564) is the fusion key mask (ego forced
valid); ``pos`` (564 x 14) are the position features, with invalid rows set to 0 (the graph's NaN rows there are
discarded by a ``Where``)."""
ego: np.ndarray # [31, 4] truncated ego history
neighbor: np.ndarray # [320, 31, 9]
neighbor_type: np.ndarray # [320, 3]
static: np.ndarray # [5, 10]
lane: np.ndarray # [140, 20, 8]
lane_attr: np.ndarray # [140, 25]
lane_speed: np.ndarray # [140, 1]
lane_has_speed: np.ndarray # [140, 1] bool
route: np.ndarray # [25, 20, 8]
route_attr: np.ndarray # [25, 25]
route_speed: np.ndarray # [25, 1]
route_has_speed: np.ndarray # [25, 1] bool
polygon: np.ndarray # [10, 40, 5]
line_string: np.ndarray # [60, 20, 6]
goal: np.ndarray # [4]
ego_shape: np.ndarray # [3]
turn: np.ndarray # [30]
valid: Dict[str, np.ndarray]
token_valid: np.ndarray # [564] bool
key_valid: np.ndarray # [564] bool
pos: np.ndarray # [564, 14]
def counts(self) -> Dict[str, int]:
return {k: int(v.sum()) for k, v in self.valid.items()}
@dataclass
class DecoderMasks:
agent_valid: np.ndarray # [321] bool: ego + neighbours with a non-zero current state (self-attention keys)
current_states: np.ndarray # [321, 4] normalised (prefix constraint)
def _host_fast() -> bool:
"""``DIFFUSION_PLANNER_HOST_FAST`` (OPT round 5 item 3), read once: the vectorised host functions (bit-exact)
instead of the first port's ``*_ref`` versions."""
from ..tt.config import KNOBS
return bool(KNOBS.read().HOST_FAST)
HOST_FAST = _host_fast()
def _any_nonzero(a: np.ndarray, axes) -> np.ndarray:
return np.any(a != 0, axis=axes)
def encoder_features(norm: Mapping[str, np.ndarray], masks: Mapping[str, np.ndarray]) -> EncoderFeatures:
"""From the normalised inputs (``normalize_inputs``) and the speed masks (``speed_masks``), batch 1."""
f32 = lambda k: np.asarray(norm[k], np.float32)[0] # noqa: E731 (drop the batch dim)
# ego: keep the 6 oldest rows (encoder.py:170-175); the token is always valid, its position row is all zero
ego = np.zeros((C.INPUT_T + 1, C.POSE_DIM), np.float32)
ego[C.EGO_HISTORY_KEEP] = f32("ego_agent_past")[C.EGO_HISTORY_KEEP]
pos = {"ego": _onehot_pos(ego[-1:], C.POS_CLASS["ego"])}
# neighbours: keep the 6 newest rows (encoder.py:176-181, 441-451)
nb_raw = f32("neighbor_agents_past")
if HOST_FAST:
K = C.NEIGHBOR_HISTORY_KEEP # the other rows are zero: computed on the kept ones
x8k = nb_raw[:, K, :8]
svk = _any_nonzero(x8k, -1) # [320, 6] valid steps
nb_valid = svk.any(axis=-1) # [320]
nb_type = nb_raw[:, -1, 8:11].copy() # row 30 is kept
feat = np.zeros(nb_raw.shape[:2] + (9,), np.float32)
fk = feat[:, K]
fk[..., :8] = x8k
fk[..., 8] = svk
fk[..., 4:6] = 0.0 # velocities zeroed after the validity test
fk[~nb_valid] = 0.0
feat[:, K] = fk
pos["neighbor"] = _onehot_pos(nb_raw[:, -1, :4], C.POS_CLASS["neighbor"])
else:
nb = np.zeros_like(nb_raw)
nb[:, C.NEIGHBOR_HISTORY_KEEP] = nb_raw[:, C.NEIGHBOR_HISTORY_KEEP]
nb_type = nb[:, -1, 8:11].copy()
x8 = nb[..., :8]
step_valid = _any_nonzero(x8, -1) # [320, 31]
nb_valid = step_valid.any(axis=-1) # [320]
feat = np.concatenate([x8, step_valid[..., None].astype(np.float32)], axis=-1)
feat[..., 4:6] = 0.0 # velocities zeroed after the validity test
feat[~nb_valid] = 0.0
pos["neighbor"] = _onehot_pos(x8[:, -1, :4], C.POS_CLASS["neighbor"])
static = f32("static_objects")
st_valid = _any_nonzero(static[..., :10], -1)
static = np.where(st_valid[:, None], static, 0.0).astype(np.float32)
pos["static"] = _onehot_pos(f32("static_objects")[:, :4], C.POS_CLASS["static"])
def lanes(key: str, speed_key: str, mask_key: str, cat: str):
x = f32(key)
attr = x[:, 0, C.LANE_FEATURE_DIM:].copy()
x8 = x[..., :C.LANE_FEATURE_DIM]
valid = _any_nonzero(x8, (-1, -2))
p = x8[:, C.LANE_POS_INDEX, :4]
pos[cat] = _onehot_pos(_heading_pos(p, p[:, 3], p[:, 2]), C.POS_CLASS[cat])
x8 = np.where(valid[:, None, None], x8, 0.0).astype(np.float32)
speed = f32(speed_key).reshape(-1, 1)
has = np.asarray(masks[mask_key])[0].reshape(-1, 1).astype(bool)
return x8, attr, speed, has, valid
lane, lane_attr, lane_speed, lane_has, lane_valid = lanes("lanes", "lanes_speed_limit",
"lanes_has_speed_limit", "lane")
route, route_attr, route_speed, route_has, route_valid = lanes("route_lanes", "route_lanes_speed_limit",
"route_lanes_has_speed_limit", "route")
def lines(key: str, cat: str, pos_index: int):
x = line_features(f32(key))
valid = _any_nonzero(x[..., :4], (-1, -2))
p = x[:, pos_index, :4]
pos[cat] = _onehot_pos(_heading_pos(p, p[:, 3], p[:, 2]), C.POS_CLASS[cat])
return np.where(valid[:, None, None], x, 0.0).astype(np.float32), valid
polygon, poly_valid = lines("polygons", "polygon", C.POLYGON_POS_INDEX)
line_string, ls_valid = lines("line_strings", "line_string", C.LINE_STRING_POS_INDEX)
goal = f32("goal_pose").reshape(-1)
pos["goal"] = _onehot_pos(goal[None], C.POS_CLASS["goal"])
unit = np.array([[0.0, 0.0, 1.0, 0.0]], np.float32)
pos["ego_shape"] = _onehot_pos(unit, C.POS_CLASS["ego_shape"])
pos["turn"] = _onehot_pos(unit, C.POS_CLASS["turn"])
turn = np.asarray(norm["turn_indicators"], np.float32)[0, :C.TURN_INDICATOR_HISTORY].copy()
one = np.ones(1, bool)
valid = {"ego": one, "neighbor": nb_valid, "static": st_valid, "lane": lane_valid, "route": route_valid,
"polygon": poly_valid, "line_string": ls_valid, "goal": one, "ego_shape": one.copy(), "turn": one.copy()}
token_valid = np.concatenate([valid[name] for name, _ in C.TOKEN_LAYOUT])
key_valid = token_valid.copy()
key_valid[0] = True
pos_all = np.concatenate([pos[name] for name, _ in C.TOKEN_LAYOUT], axis=0).astype(np.float32)
pos_all[~token_valid] = 0.0
return EncoderFeatures(ego, feat.astype(np.float32), nb_type, static, lane, lane_attr, lane_speed, lane_has,
route, route_attr, route_speed, route_has, polygon, line_string, goal,
f32("ego_shape").reshape(-1).copy(), turn, valid, token_valid, key_valid, pos_all)
def decoder_masks(norm: Mapping[str, np.ndarray]) -> DecoderMasks:
"""Agent key mask of the DiT self-attention (``dit.py:155-156``; the ego is always valid) and the current states
of the prefix constraint (``multi_step_inference.cpp:300-325``), from the normalised inputs."""
nb_now = np.asarray(norm["neighbor_agents_past"], np.float32)[0, :, C.INPUT_T, :C.POSE_DIM]
agent_valid = np.concatenate([np.ones(1, bool), np.any(nb_now != 0, axis=-1)])
cs = np.zeros((C.MAX_NUM_AGENTS, C.POSE_DIM), np.float32)
cs[0] = np.asarray(norm["ego_current_state"], np.float32)[0, :C.POSE_DIM]
cs[1:] = nb_now
return DecoderMasks(agent_valid, cs)