Download code/tt_diffusion_planner/host/features.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 12.4 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/host/features.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/host/features.py
-
curl -L -o features.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/host/features.py
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) | |
| 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()} | |
| 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) | |