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