changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
17.6 kB
# SPDX-License-Identifier: Apache-2.0
"""Device-side constants and the precision policy of the ttnn graph (COMPILE parameters, PLAN.md 0.4).
Shapes. The masked attentions run on tile-aligned key counts: a user mask with an unaligned key count lets the
padded keys of the last tile into the softmax (``ttaw.ops.attention``, C20), so the fusion transformer runs on
:data:`TOKENS` = 576 tokens (564 + 12 zero pad tokens, masked as keys) and the decoder on :data:`AGENTS` = 352 agents
(321 + 31 extra rows, masked as self-attention keys). Both counts are what the tiles hold anyway (no extra compute).
Rows past the real counts are computed (finite) and never read back. The cross-attention keys are the real 564
tokens (no mask: the op masks its own padding element by element).
Solver state. ``y = x * mask0`` (the state with the t = 0 columns, the prefix-constrained current state, zeroed) is
the fp32 state of the on-device DPM-Solver++(2M) loop; the decoder input is ``x = y + cs`` (``cs`` = current
states in the t = 0 columns) and the last projection's t = 0 output columns are zeroed in the weights, so the
model output is ``m * mask0`` and every update ``y' = A y - B m_k + C m_(k-1)`` keeps ``y * mask0 = y``. This is
the node's update + ``apply_prefix_constraint`` exactly (``m``'s t = 0 slot is never used: the correction
overwrites it).
Numerics options (:data:`KNOBS`, measured on the device, PORT_LOG.md 5.2): the fused ``ttnn.layer_norm`` loses
the per-entity signal on offset-dominated rows (fp32 decomposition for the mixers and the decoder: ``LN_FP32``); a
device fp32 matmul truncates its operands TF32-like, which the mixers amplify ~100x and which, compounded over the 11
decoder evaluations, moves the plan (split hi / lo matmuls for the mixer inputs and every decoder linear:
``SPLIT_MATMUL``); and the bf16 SDPA kernel is the largest encoder error left on the sensitive nuScenes instants
(fp32 matmul attention in the fusion and the decoder: ``ATTN_MATMUL``). With the defaults the 27 e2e scenes (7
research + 20 nuScenes instants) stay within ego max 0.31 m / mean 0.14 m / neighbour median-max 0.09 m of the fp32
reference (gates 1.0 / 0.3 / 1.5 m); the defaults of the first device round failed the mean on one instant (0.35 m).
Precision (:data:`DEFAULT_PRECISION`, ``ttaw.precision.PrecisionPolicy`` rules, first match wins; override with
``DIFFUSION_PLANNER_PRECISION="dec.*=HiFi2+fp32:a=bf16"``): every matmul / LayerNorm gets an explicit compute
config; ``w=`` is the weight dtype, ``a=`` the dtype of the module's residual stream (sub-layer outputs are produced
in it and LayerNorm keeps it). The ego / neighbour pre-projection island is fp32 and pad-relative (probe P12, a
device fp32 matmul is TF32-like); the decoder pre-projection reads the fp32 solver state with fp32 weights.
"""
from __future__ import annotations
from typing import Dict
from ..reference import config as C
__all__ = ["TOKENS", "TOKENS_REAL", "AGENTS", "AGENTS_REAL", "STATE_COLS", "STATE_COLS_T0", "TILE",
"ISLAND_ROWS", "LANE_AUX_DIM", "NEIGHBOR_AUX_DIM", "POS_AUG_DIM", "DEFAULT_PRECISION", "MIXER_T",
"MIXER_CIN", "KNOBS", "globs"]
TILE = 32
def _aligned(n: int) -> int:
return -(-n // TILE) * TILE
TOKENS_REAL = C.ENCODING_TOKEN_NUM # 564 encoder tokens (the cross-attention keys)
TOKENS = _aligned(TOKENS_REAL) # 576: fusion rows / masked keys
AGENTS_REAL = C.MAX_NUM_AGENTS # 321
AGENTS = _aligned(AGENTS_REAL) # 352: decoder rows / masked self-attention keys
STATE_COLS = C.DIT_INPUT_DIM # 324 = 81 points x 4
STATE_COLS_T0 = C.POSE_DIM # columns 0..3 hold the t = 0 point (prefix constraint)
# pre-projection island rows (the only non-zero time rows after the in-graph truncation, SPEC 3.8)
ISLAND_ROWS = {"ego": tuple(range(C.EGO_HISTORY_KEEP.start, C.EGO_HISTORY_KEEP.stop)),
"neighbor": tuple(range(C.NEIGHBOR_HISTORY_KEEP.start, C.NEIGHBOR_HISTORY_KEEP.stop))}
# mixer token axis length (T) and input channels per category (host features)
MIXER_T = {"ego": len(ISLAND_ROWS["ego"]), "neighbor": len(ISLAND_ROWS["neighbor"]), "lane": C.POINTS_PER_SEGMENT,
"route": C.POINTS_PER_SEGMENT, "polygon": C.POINTS_PER_POLYGON, "line_string": C.POINTS_PER_LINE_STRING}
MIXER_CIN = {"ego": C.POSE_DIM, "neighbor": C.NEIGHBOR_FEATURE_DIM, "lane": C.LANE_FEATURE_DIM,
"route": C.LANE_FEATURE_DIM, "polygon": C.POLYGON_FEATURE_DIM,
"line_string": C.LINE_STRING_FEATURE_DIM}
# exact affine rewrites of the small embeddings (tt.params): the host builds the input columns
NEIGHBOR_AUX_DIM = 3 + 1 # [type one-hot (3), 1] @ [W_type; b_type]
LANE_AUX_DIM = 3 + C.LANE_ATTRIBUTE_DIM + 1 # [speed*has, has, 1-has, attributes (25), 1] @ [w_s; b_s; unk; W_a; b_a]
POS_AUG_DIM = C.POS_FEATURE_DIM + 1 # [pos (14), token valid] @ [W_pos; b_pos]
DEFAULT_PRECISION: Dict[str, str] = {
"enc.island.*": "HiFi4+fp32:w=fp32:a=fp32",
"enc.mixer.*": "HiFi4+fp32:w=bf16:a=fp32",
"enc.fusion*": "HiFi4+fp32:w=bf16:a=fp32",
"enc.*": "HiFi4+fp32:w=bf16:a=fp32",
"dec.preproj.fc1": "HiFi4+fp32:w=fp32:a=fp32",
"dec.*": "HiFi4+fp32:w=bf16:a=fp32",
"turn": "HiFi4+fp32:w=fp32:a=fp32",
}
def _knobs():
from ..ttaw.knobs import Knob, Knobs
return Knobs("DIFFUSION_PLANNER", [
Knob("LN_FP32", "enc.mixer.*,dec.*", "comma-separated module globs whose LayerNorms run as an fp32 "
"decomposition (mean, variance, rsqrt as fp32 SFPU ops) instead of ttnn.layer_norm, whose error on "
"offset-dominated rows (|mean| / std 10-30) is rel-L2 0.03 (PORT_LOG 5.2); 'none' = empty"),
Knob("HIDDEN_FP32", "", "comma-separated module globs whose hidden MLP activations are fp32 (default "
"bf16); 'none' = empty"),
Knob("SPLIT_MATMUL", "enc.island.*,enc.pre.*,dec.*", "comma-separated module globs whose fp32 "
"matmuls are split into bf16 hi / fp32 lo parts (2-3 matmuls, ~1e-5 relative instead of the TF32-like "
"~1e-3; their hidden activations stay fp32): the mixer inputs and every decoder linear; 'none' = "
"empty"),
Knob("ATTN_FP32_ACC", "", "comma-separated module globs whose SDPA runs with fp32 accumulation (probe P7: "
"half the max error, ~1.4x the time); 'none' = empty"),
Knob("ATTN_MATMUL", "enc.fusion.attn,dec.*", "comma-separated module globs (enc.fusion.attn, "
"dec.self_attn, dec.cross_attn) whose attention runs as fp32 matmuls + softmax (C20 attention_matmul) "
"instead of the bf16 SDPA kernel; 'none' = empty"),
Knob("ENC_CH2D", True, "the encoder linears on batched [1, E, T, C] activations run as one 2-D matmul "
"over all rows (free [1, 1, E*T, C] view for T = 64, else a 1-D program config with fuse_batch) "
"instead of the stock per-element tiling on 4-8 cores (OPT round 1 item 1); 0 = the stock call"),
Knob("ATTN_FAST", 1, "fp32 matmul attention variant (tt/attention.py, OPT round 1 item 4a): 0 = "
"ttaw.ops.attention.attention_matmul; 1 = P.V on one output tile per core (bit-identical); 2 = 1 + "
"the scale on Q; 3 = 1 + the scale fused into the mask add (not bit-identical, rejected)",
choices=(0, 1, 2, 3)),
Knob("DEC_MMCFG", True, "explicit 2-D multicast program configs for the decoder's 352-row matmuls "
"(tt/layers.py DEC_MM_CONFIGS, the fastest bit-identical config per shape and split pass of the device "
"sweep, OPT round 1 item 2a); 0 = the auto config"),
Knob("LN_KERNEL", 2, "the fp32 LayerNorm decomposition (LN_FP32 modules) as one fused generic_op "
"program (tt/ln_kernel.py, the same SFPU LLK sequence, bit-identical: OPT round 2 item 3) instead of "
"7-9 stock programs: 1 = per-tile unpacker / SFPU inits, 2 = one init per phase; 0 = the stock "
"decomposition", choices=(0, 1, 2)),
Knob("LN_RESID", True, "with LN_KERNEL: the residual add in front of a fused fp32 LayerNorm (mixer blocks: "
"x + y, x + ch2; decoder: h + gate * a, h + m2b, ...) in the same program, which writes the new stream "
"and the normalised rows (OPT round 2 item 5); 0 = separate add / multiply programs"),
Knob("SPLIT_KCAT", 2, "decoder split matmuls (SPLIT_MATMUL dec.*, tile-aligned K) as one fp32 matmul over "
"the concatenated K axis [x_hi|x_hi|x_lo|1] @ [w_hi;w_lo;w_hi;b] (OPT round 2 item 2): 0 = three "
"passes + adds (the published numerics); 1 = stock-op operand build (typecast / subtract / concat); 2 = "
"the operand built by one generic_op (tt/kcat_kernel.py); not bit-identical to 0 (one fp32 accumulation, "
"exact bias), gates green",
choices=(0, 1, 2)),
Knob("ATTN_SMASK", True, "fp32 matmul attention (ATTN_FAST=1) with a mask: the score scale multiply and the "
"mask add as one generic_op (tt/smask_kernel.py, the same SFPU LLK calls: bit-identical; OPT round 2 "
"item 4b) instead of two binary_ng programs; 0 = the two programs"),
Knob("ATTN_SMSM", 1, "with ATTN_SMASK: the score scale, the mask add and the softmax as one generic_op "
"(tt/smsm_kernel.py: the smask SFPU sequence, then the stock softmax kernel_lib calls on the row in L1; "
"OPT round 3 item 1, bit-identical) instead of the smask program + ttnn.softmax: 1 = the scale as a "
"multiply by a scale tile (mul_binary_tile, as smask), 2 = the same fp32 SFPU multiply by an immediate "
"(mul_unary_tile, no scale tile copy); 0 = the two programs", choices=(0, 1, 2)),
Knob("ATTN_FUSED", True, "with ATTN_SMSM: each fp32 matmul attention (decoder self / cross, fusion) as one "
"generic_op (tt/fattn_kernel.py: Q K^T into DEST, the smsm phases, P V accumulated in DEST, Q / K / V "
"read in place and the heads merged on write; OPT round 3 item 4) instead of head split, Q K^T, smsm, "
"P V and head merge; 0 = those programs"),
Knob("KCAT_EMIT", True, "with SPLIT_KCAT=2: the producers of the decoder's K-concatenated split linears' "
"inputs (the split-row LayerNorms -> qkv / mlp fc1 / cross q, the fused attention -> attn out / cross "
"out) write the split operand [x_hi | x_hi | x_lo | 1] themselves (the kcat LLK calls on their output "
"tile; OPT round 3 item 5) instead of x + a kcat operand program; 0 = the operand programs"),
Knob("LN_TR", True, "with LN_KERNEL and LN_RESID: the two transposes around each mixer block's token-mixing "
"MLP done inside the fused LayerNorm programs (n1 writes its output per-entity transposed, n2 reads the "
"token-mixing output transposed; the stock transpose LLK, exact; OPT round 3 item 6) instead of two "
"ttnn.transpose programs (mixer trunks on the one-core-per-row LN kernel); 0 = the transpose programs"),
Knob("ENC_KCAT", True, "with SPLIT_KCAT=2: the encoder's split linears (SPLIT_MATMUL enc.pre.*, "
"enc.island.*; any K, blocks padded to whole tiles) as one K-concatenated fp32 matmul too (OPT round 3 "
"item 2; one fp32 accumulation, exact bias: a precision change of the SPLIT_KCAT kind); 0 = three "
"passes + adds"),
Knob("LIN_ACT", False, "the GELU of the encoder mixer linears (token / channel MLP fc1, bf16 output) in the "
"matmul epilogue (an explicit copy of the stock auto config + fused_activation, OPT round 3 item 3) "
"instead of a unary program on the rounded bf16 output; a precision change (GELU before the bf16 "
"rounding); 0 = the unary program"),
Knob("KCAT_ACT", True, "with SPLIT_KCAT=2: the GELU of a decoder linear that feeds another K-concatenated "
"split linear (preproj.fc1 -> fc2, mlp fc1 -> fc2) applied inside the next one's operand build (the "
"same SFPU LLK: bit-identical) instead of its own unary program; 0 = the unary program"),
Knob("LN_SPLIT", True, "with LN_KERNEL: a fused fp32 LayerNorm over few tile rows (the decoder's 11) spread "
"over Wt cores per row (tt/ln_kernel.py layer_norm_fp32_split: the root core folds the gathered tiles "
"in order and broadcasts the statistics, bit-identical) instead of one core per row; 0 = one core per "
"row"),
Knob("KCAT_L1", True, "with SPLIT_KCAT=2: the decoder's K = 1024 split operands (mlp fc2, final p4) written "
"L1 block-sharded by the operand build and read in place by a 2-D multicast matmul with in0 sharded "
"(tt/layers.py KCAT_L1_CONFIGS; OPT round 4 item 1, bit-identical) instead of a DRAM round trip of "
"the 4.7 MB operand; 0 = DRAM interleaved"),
Knob("KCAT_ACT_ONCE", False, "with KCAT_ACT: the operand build applies the deferred GELU to one DEST copy of "
"each tile and copies it to the second (copy_dest_values; OPT round 4 item 2, bit-identical) instead "
"of running the GELU on both copies (measured, not kept: +0.06 ms); 0 = twice"),
Knob("ATTN_L1", 2, "with ATTN_FUSED: the fused attention's inputs in L1 (interleaved) instead of DRAM: the "
"decoder qkv / cross q projections write L1, the hoisted cross K / V heads and the masks are L1-resident "
"(OPT round 4 item 3, bit-identical; the 11 query-row cores of a head re-read the same K / V tiles); "
"1 = also the fusion q / kv projections (the stock linear then runs on 6 cores with a separate bias "
"add: +0.3 ms, item 8); 2 = not those; 0 = DRAM", choices=(0, 1, 2)),
Knob("DEC_L1", 2, "the decoder blocks' intermediates (the split-row LN outputs and stream, the fused "
"attention outputs, the out / mlp / cross-out linear outputs) in L1 (interleaved) instead of DRAM (OPT "
"round 4 item 4, bit-identical); 2 = also the pre-projection and the final layer (item 6); 0 = DRAM",
choices=(0, 1, 2)),
Knob("ENC_L1", False, "the encoder mixer blocks' intermediates (the fused LN outputs and stream, the token / "
"channel MLP outputs) in L1 (interleaved) instead of DRAM (OPT round 4 item 5, bit-identical; measured, "
"not kept: +1.67 ms); 0 = DRAM"),
Knob("FUS_L1", False, "the encoder fusion blocks' intermediates (LN outputs, stream adds, attention / out / "
"MLP outputs) in L1 (interleaved) instead of DRAM (OPT round 4 item 7; measured, not kept: +0.39 ms and "
"not bit-identical, the stock ops pick other programs for L1 outputs); 0 = DRAM"),
Knob("COMPACT", 2, "exact compaction (OPT round 5): besides the full-capacity plan, one trace per agent "
"bucket (AGENT_BUCKETS), picked per request from the host masks. 1 = the decoder runs on the first R "
"rows only (the needed rows: ego, every valid self-attention key, every emitted neighbour, every valid "
"neighbour token; the rows past R are masked keys and never read; item 1); 2 = also the encoder's "
"neighbour trunk and head on the first R neighbours, the other neighbour tokens zero (invalid: "
"token_valid zeroes them anyway; item 2). The matmul configs keep their K blocking, so the needed rows "
"are the same values; 0 = the full 352 rows / 320 neighbours always", choices=(0, 1, 2)),
Knob("AGENT_BUCKETS", "32,64,96,128,192", "with COMPACT: the decoder row buckets (multiples of 32 below "
"352), one captured trace each"),
Knob("HOST_FAST", True, "the host pre- and post-processing vectorised (OPT round 5 item 3: the normalisation "
"computes only the rows that are not all-small, the neighbour features only the 6 kept history rows, "
"the trajectory's velocity smoothing / force stop / acceleration as array operations in the same "
"float64 / float32 order, the denormalisation and predicted paths on the emitted rows only); bit-exact "
"(the host suite compares the arrays); 0 = the first port's *_ref functions"),
Knob("INPUT_TRIM", True, "per-request uploads trimmed (OPT round 5 item 4, H5): the solver's current states "
"``cs`` uploaded as one tile column [352, 32] (the 4 t = 0 columns + zeros) and widened to 324 columns in "
"the trace by a tile-aligned concat with a zero block (the same values), and an input whose array is all "
"+0.0 (``y0`` at temperature 0, ``static_x``) not uploaded again while its device buffer already holds "
"zeros from this model's last upload; 0 = every input every plan, ``cs`` at full width"),
Knob("LN_SFPU_BCAST", True, "with LN_KERNEL: the row statistics written to every column by the SFPU row "
"reduce itself (kernels/ln32_sfpu.h, the stock reduce's arithmetic with extra stores: bit-identical) "
"instead of a RISC-V column fill of the packed column-0 tile; 0 = the fill"),
])
KNOBS = _knobs()
def globs(text: str):
"""``"a.*, b"`` -> ``("a.*", "b")``; ``""`` / ``"none"`` -> ``()``."""
t = (text or "").strip()
if t.lower() in ("", "none"):
return ()
return tuple(g.strip() for g in t.split(",") if g.strip())