File size: 5,612 Bytes
4d9b003
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
be62f78
4d9b003
 
 
be62f78
 
 
4d9b003
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
be62f78
4d9b003
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
be62f78
4d9b003
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
# 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