changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
5.61 kB
# SPDX-License-Identifier: Apache-2.0
"""Host packing of one plan's :class:`host.pipeline.Prepared` into the persistent trace inputs (numpy, no ttnn).
Every array is float32 in the 4-D shape of its device input (TILE, uploaded each plan): the mixer inputs keep only
the time rows the export reads (ego 0..5, neighbours 25..30: the pad-relative island needs nothing else), the aux
columns of the exact embedding rewrites (``tt/params.py``), the 576-row token arrays (564 real + 12 zero pad
tokens), the key-bias rows (0 / -inf; extra keys masked) and the solver's ``cs`` / ``y0`` on 352 rows.
"""
from __future__ import annotations
from typing import Dict, Optional
import numpy as np
from ..host.pipeline import Prepared
from ..reference import config as C
from ..ttaw.ops.attention import key_bias_row
from . import config as T
from .params import lane_aux_features, pad_rows
__all__ = ["plan_inputs", "INPUT_SPECS", "warmup_inputs", "decoder_state", "BF16_INPUTS", "CS_COLS"]
f32 = np.float32
# INPUT_TRIM: the current states as one tile column (widened to STATE_COLS in the trace)
CS_COLS = T.TILE if T.KNOBS.read().INPUT_TRIM else T.STATE_COLS
# name -> device shape (all fp32 TILE except the bf16 key-bias rows)
INPUT_SPECS: Dict[str, tuple] = {
"ego_x": (1, 1, T.MIXER_T["ego"], T.MIXER_CIN["ego"]),
"neighbor_x": (1, C.MAX_NUM_NEIGHBORS, T.MIXER_T["neighbor"], T.MIXER_CIN["neighbor"]),
"neighbor_aux": (1, 1, C.MAX_NUM_NEIGHBORS, T.NEIGHBOR_AUX_DIM),
"static_x": (1, 1, C.NUM_STATIC_OBJECTS, C.STATIC_OBJECT_DIM),
"lane_x": (1, C.NUM_SEGMENTS_IN_LANE, T.MIXER_T["lane"], T.MIXER_CIN["lane"]),
"lane_aux": (1, 1, C.NUM_SEGMENTS_IN_LANE, T.LANE_AUX_DIM),
"route_x": (1, C.NUM_SEGMENTS_IN_ROUTE, T.MIXER_T["route"], T.MIXER_CIN["route"]),
"route_aux": (1, 1, C.NUM_SEGMENTS_IN_ROUTE, T.LANE_AUX_DIM),
"polygon_x": (1, C.NUM_POLYGONS, T.MIXER_T["polygon"], T.MIXER_CIN["polygon"]),
"line_string_x": (1, C.NUM_LINE_STRINGS, T.MIXER_T["line_string"], T.MIXER_CIN["line_string"]),
"goal_x": (1, 1, 1, C.POSE_DIM),
"ego_shape_x": (1, 1, 1, C.EGO_SHAPE_DIM),
"turn_x": (1, 1, 1, C.TURN_INDICATOR_HISTORY),
"token_valid": (1, 1, T.TOKENS, 1),
"pos_aug": (1, 1, T.TOKENS, T.POS_AUG_DIM),
"fusion_key_row": (1, 1, 1, T.TOKENS),
"agent_key_row": (1, 1, 1, T.AGENTS),
"cs": (1, 1, T.AGENTS, CS_COLS),
"y0": (1, 1, T.AGENTS, T.STATE_COLS),
}
BF16_INPUTS = ("fusion_key_row", "agent_key_row")
def decoder_state(x: np.ndarray, current_states: Optional[np.ndarray] = None) -> np.ndarray:
"""A ``[321, 81, 4]`` state -> ``[1, 1, 352, 324]`` float32 (extra rows zero). With ``current_states`` the t = 0
columns are replaced by them (the ``cs`` input); with None they are zeroed (the ``y`` state)."""
flat = np.asarray(x, f32).reshape(C.MAX_NUM_AGENTS, T.STATE_COLS)
out = pad_rows(flat, T.AGENTS).astype(f32)
out[:, :T.STATE_COLS_T0] = 0.0
if current_states is not None:
out[:C.MAX_NUM_AGENTS, :T.STATE_COLS_T0] = np.asarray(current_states, f32)
return out[None, None]
def plan_inputs(prep: Prepared) -> Dict[str, np.ndarray]:
"""``{input name: array}`` for one plan (see :data:`INPUT_SPECS`)."""
f = prep.features
ego_rows, nb_rows = list(T.ISLAND_ROWS["ego"]), list(T.ISLAND_ROWS["neighbor"])
tv = pad_rows(np.asarray(f.token_valid, bool), T.TOKENS)
nb_type = np.asarray(f.neighbor_type, f32)
out = {
"ego_x": np.asarray(f.ego, f32)[ego_rows][None, None],
"neighbor_x": np.asarray(f.neighbor, f32)[:, nb_rows][None],
"neighbor_aux": np.concatenate([nb_type, np.ones((nb_type.shape[0], 1), f32)], 1)[None, None],
"static_x": np.asarray(f.static, f32)[None, None],
"lane_x": np.asarray(f.lane, f32)[None],
"lane_aux": lane_aux_features(f.lane_speed, f.lane_has_speed, f.lane_attr)[None, None],
"route_x": np.asarray(f.route, f32)[None],
"route_aux": lane_aux_features(f.route_speed, f.route_has_speed, f.route_attr)[None, None],
"polygon_x": np.asarray(f.polygon, f32)[None],
"line_string_x": np.asarray(f.line_string, f32)[None],
"goal_x": np.asarray(f.goal, f32).reshape(1, 1, 1, -1),
"ego_shape_x": np.asarray(f.ego_shape, f32).reshape(1, 1, 1, -1),
"turn_x": np.asarray(f.turn, f32).reshape(1, 1, 1, -1),
"token_valid": tv.astype(f32).reshape(1, 1, -1, 1),
"pos_aug": np.concatenate([pad_rows(np.asarray(f.pos, f32), T.TOKENS), tv.astype(f32)[:, None]],
1)[None, None],
"fusion_key_row": key_bias_row(pad_rows(np.asarray(f.key_valid, bool), T.TOKENS)),
"agent_key_row": key_bias_row(pad_rows(np.asarray(prep.decoder.agent_valid, bool), T.AGENTS)),
"cs": decoder_state(np.zeros((C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM), f32),
prep.decoder.current_states)[..., :CS_COLS],
"y0": decoder_state(prep.x_T),
}
for name, arr in out.items():
if arr.shape != INPUT_SPECS[name]:
raise AssertionError(f"input {name}: shape {arr.shape}, expected {INPUT_SPECS[name]}")
return out
def warmup_inputs() -> Dict[str, np.ndarray]:
"""Valid-looking initial contents (every token / agent a valid key: no fully masked softmax row)."""
out = {name: np.zeros(shape, f32) for name, shape in INPUT_SPECS.items()}
out["fusion_key_row"] = key_bias_row(np.arange(T.TOKENS) < T.TOKENS_REAL)
out["agent_key_row"] = key_bias_row(np.arange(T.AGENTS) < T.AGENTS_REAL)
out["token_valid"][..., :T.TOKENS_REAL, 0] = 1.0
return out