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