# 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