changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
7.01 kB
# SPDX-License-Identifier: Apache-2.0
"""The host halves of one plan, shared by the Python API, the HTTP server and the CPU reference.
``prepare(raw, normalization)`` -> :class:`Prepared`: the raw ONNX-named tensors (already checked against
``config.INPUT_SCHEMA``) normalised like the node, the speed masks, the encoder's host features, the decoder masks
and the solver's initial ``x``. The network (device trace or CPU reference) turns it into ``final_x0`` (normalised
``[321, 81, 4]``, prefix constraint applied) and the turn-indicator logits; ``make_output(...)`` applies the node's
post-processing and returns the ``ttaw.outputs.Trajectory`` that ``model(...)`` and ``POST /predict`` return.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Dict, Mapping, Optional, Sequence
import numpy as np
from ..reference import config as C
from ..ttaw.io import InputError, encode_array
from ..ttaw.outputs import Trajectory
from .features import DecoderMasks, EncoderFeatures, decoder_masks, encoder_features
from .normalize import normalize_inputs, speed_masks
from .postprocess import HOST_FAST
from .postprocess import (TurnIndicatorManager, denoising_steps_ego, denormalize, predicted_paths,
trajectory_from_poses)
__all__ = ["Prepared", "prepare", "make_output", "TRAJECTORY_COLUMNS", "PREDICTED_AGENT_COLUMNS", "RUNTIME_PARAMS"]
TRAJECTORY_COLUMNS = ("x", "y", "yaw", "cos", "sin", "velocity", "acceleration")
PREDICTED_AGENT_COLUMNS = ("x", "y", "yaw", "cos", "sin")
# RT-host parameters (request params / call kwargs): name -> (type, min, max, default); the node's YAML defaults
RUNTIME_PARAMS = {
"velocity_smoothing_window": (int, 1, C.OUTPUT_T - 1, C.VELOCITY_SMOOTHING_WINDOW),
"stopping_threshold": (float, 0.0, None, C.STOPPING_THRESHOLD),
"turn_indicator_keep_offset": (float, None, None, C.TURN_INDICATOR_KEEP_OFFSET),
"return_denoising_steps": (bool, None, None, False),
}
@dataclass
class Prepared:
"""Host state of one plan."""
raw: Dict[str, np.ndarray] # the validated raw tensors (batch 1)
norm: Dict[str, np.ndarray] # normalised tensors + the two speed masks
features: EncoderFeatures
decoder: DecoderMasks
x_T: np.ndarray # [321, 81, 4] initial solver state (normalised; correction not yet applied)
neighbor_rows: np.ndarray # indices of the non-empty neighbour rows (the node's emitted agents)
enable_force_stop: bool # ego vx > DBL_EPSILON (diffusion_planner_core.cpp:628-629)
prev_report: int # the latest TurnIndicatorsReport (turn_indicators[30])
def prepare(raw: Mapping[str, np.ndarray], observation: Mapping[str, Any]) -> Prepared:
"""``raw``: the 15 tensors of ``config.INPUT_SCHEMA`` (float32, batch 1, ego frame, before normalization);
``observation``: ``Normalization.observation``."""
missing = sorted(set(C.INPUT_NAMES) - set(raw))
if missing:
raise InputError(f"inputs: missing {missing}")
raw = {k: np.ascontiguousarray(np.asarray(raw[k], np.float32)) for k in C.INPUT_NAMES}
for k, shape in C.INPUT_SHAPES.items():
if raw[k].shape != shape:
raise InputError(f"inputs[{k!r}] has shape {raw[k].shape}, expected {shape}")
try:
norm = normalize_inputs(raw, observation)
except (KeyError, ValueError) as e: # a normalizer problem is a weights-file problem, not a client mistake
raise RuntimeError(f"normalization failed: {e}") from e
norm.update(speed_masks(norm))
feats = encoder_features(norm, norm)
dec = decoder_masks(norm)
nb = raw["neighbor_agents_past"][0].reshape(C.MAX_NUM_NEIGHBORS, -1)
rows = np.flatnonzero(np.any(nb != 0, axis=1))
vx = float(raw["ego_current_state"][0, 4])
prev_report = int(round(float(raw["turn_indicators"][0, C.INPUT_T])))
x_T = norm["sampled_trajectories"][0].copy()
return Prepared(raw, norm, feats, dec, x_T, rows, vx > np.finfo(np.float64).eps, prev_report)
def make_output(final_x0: np.ndarray, logit: np.ndarray, prepared: Prepared, normalization: Any,
params: Mapping[str, Any], *, model: str = "", denoising_steps: Optional[Sequence[np.ndarray]] = None,
timing_ms: Optional[Dict[str, float]] = None, meta: Optional[Dict[str, Any]] = None) -> Trajectory:
"""The node's post-processing of one plan -> ``Trajectory`` (base_link).
``final_x0``: ``[321, 81, 4]`` normalised (t = 0 = current state); ``logit``: ``[5]``; ``params``: the validated
``RUNTIME_PARAMS``; ``denoising_steps``: the 11 iterates ``[321, 81, 4]`` when ``return_denoising_steps``."""
x0 = np.asarray(final_x0, np.float32).reshape(C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM)
mean, std = normalization.state()
rows = prepared.neighbor_rows
if HOST_FAST: # only the rows read below: the ego and the emitted neighbours (element-wise: the same values)
idx = np.concatenate([[0], np.asarray(rows, np.int64) + 1])
pick = lambda v: v[idx] if v.shape[0] == C.MAX_NUM_AGENTS else v # noqa: E731
denorm = (x0[idx, 1:] * pick(std) + pick(mean)).astype(np.float32) # [1 + n, 80, 4]
else:
denorm = denormalize(x0, mean, std) # [321, 80, 4]
traj = trajectory_from_poses(denorm[0], (0.0, 0.0, 0.0),
velocity_smoothing_window=int(params["velocity_smoothing_window"]),
enable_force_stop=prepared.enable_force_stop,
stopping_threshold=float(params["stopping_threshold"]))
manager = TurnIndicatorManager(keep_offset=float(params["turn_indicator_keep_offset"]))
decision = manager.evaluate(np.asarray(logit, np.float32).reshape(-1), 0.0, prepared.prev_report)
paths = (predicted_paths(denorm, rows.size) if HOST_FAST else predicted_paths(denorm, C.MAX_NUM_NEIGHBORS)[rows]) \
if rows.size else np.zeros(
(0, C.OUTPUT_T, len(PREDICTED_AGENT_COLUMNS)), np.float32)
out_meta: Dict[str, Any] = {"predicted_agent_columns": list(PREDICTED_AGENT_COLUMNS),
"predicted_agent_rows": [int(r) for r in rows],
"force_stop": bool(traj.force_stop),
"time_from_start_s": [round(float(t), 3) for t in traj.time_from_start],
"valid_counts": prepared.features.counts()}
if params.get("return_denoising_steps") and denoising_steps is not None:
ego_steps = denoising_steps_ego(np.stack(denoising_steps), mean, std)
out_meta["denoising_steps"] = encode_array(ego_steps, key="denoising_steps")
out_meta.update(meta or {})
return Trajectory(traj.as_columns(), columns=TRAJECTORY_COLUMNS, turn_indicator=decision.to_dict(),
predicted_agents=paths, model=model, frame_id="base_link", timing_ms=dict(timing_ms or {}),
meta=out_meta)