# SPDX-License-Identifier: Apache-2.0 """``TtDiffusionPlanner``: the whole plan of Diffusion Planner v5.0 in one trace on a Blackhole p150. Variant ``plan`` (``ttaw.trace.TraceRunner``; inputs :data:`tt.inputs.INPUT_SPECS`, written every plan): encoder + fusion (:mod:`.encoder`) -> cross K / V of the 3 DiT blocks hoisted once -> 11 x (DiT evaluation with the per-step folded adaLN rows + fp32 DPM-Solver++(2M) update + prefix constraint) (:mod:`.decoder`) -> turn head -> one packed readback (``final_x0`` ``[352, 324]`` fp32, ``logit`` ``[5]``, ``ego_steps`` ``[11, 324]``). No host fallback inside the plan, so there are no trace segments. Debug variants (``debug=True``; tests only, also replayed from traces): - ``encoder_taps``: the encoder of ``plan`` returning every encoder tap of ``reference.model.TAP_NAMES``; - ``decode_once``: one decoder evaluation on teacher-forced inputs (``dbg_x`` ``[352, 324]``, ``dbg_enc`` ``[576, 256]``: the reference encoding + 12 zero pad-token rows) with the per-step rows as inputs (``dbg..``, ``dbg.final.g`` / ``.b``), so one trace serves all 11 evaluation times. Everything here runs under the model lock of ``api.DiffusionPlanner`` (one chip, batch 1). """ from __future__ import annotations import time from typing import Any, Dict, List, Optional, Sequence import numpy as np from ..reference import config as C from ..ttaw.ops import attention as A from ..ttaw.trace import TraceRunner, pack_outputs from . import config as T from . import inputs as I from . import params as P from .decoder import STEP_KEYS, TtDecoder, TtTurnHead from .encoder import TtEncoder from .layers import Build, policy __all__ = ["TtDiffusionPlanner", "agent_buckets", "needed_rows", "bucket_rows"] def agent_buckets() -> tuple: """``COMPACT``: the agent buckets of ``AGENT_BUCKETS`` (multiples of 32 below 352, sorted); () when off.""" k = T.KNOBS.read() if not k.COMPACT: return () out = sorted({int(v) for v in str(k.AGENT_BUCKETS).replace(" ", "").split(",") if v}) bad = [v for v in out if v % T.TILE or not 0 < v < T.AGENTS] if bad: raise ValueError(f"DIFFUSION_PLANNER_AGENT_BUCKETS: {bad} are not multiples of {T.TILE} below {T.AGENTS}") return tuple(out) def needed_rows(prepared: Any) -> int: """1 + the last decoder row a plan reads or attends to: the ego (row 0), the valid self-attention keys (``agent_valid``) and the emitted neighbours (``neighbor_rows`` + 1). The rows past it are masked keys whose outputs are never read, so dropping them leaves the needed rows' values unchanged.""" valid = np.flatnonzero(np.asarray(prepared.decoder.agent_valid, bool)) last = int(valid.max()) if valid.size else 0 tok = np.flatnonzero(np.asarray(prepared.features.valid["neighbor"], bool)) # valid neighbour tokens if tok.size: last = max(last, int(tok.max()) + 1) emitted = np.asarray(prepared.neighbor_rows) if emitted.size: last = max(last, int(emitted.max()) + 1) return last + 1 def bucket_rows(prepared: Any, buckets: Sequence[int]) -> int: n = needed_rows(prepared) return next((int(b) for b in sorted(buckets) if b >= n), T.AGENTS) class TtDiffusionPlanner: """Weights on the device, the ``TraceRunner`` with its persistent inputs and variants, and the host glue. ``weights``: ``reference.weights.PlannerWeights``; ``precision``: extra policy rules (``"dec.*=HiFi2+fp32"``) on top of ``DIFFUSION_PLANNER_PRECISION``; ``ln_fp32`` / ``hidden_fp32`` / ``split`` / ``attn_fp32_acc`` / ``attn_matmul``: module globs overriding the knobs of the same names (``tt/config.py``); ``debug``: also register ``encoder_taps`` / ``decode_once``.""" def __init__(self, device: Any, weights: Any, *, debug: bool = False, precision: Optional[str] = None, ln_fp32: Optional[Sequence[str]] = None, hidden_fp32: Optional[Sequence[str]] = None, split: Optional[Sequence[str]] = None, attn_fp32_acc: Optional[Sequence[str]] = None, attn_matmul: Optional[Sequence[str]] = None, steps: int = C.DPM_SOLVER_STEPS, num_command_queues: Optional[int] = None): import ttnn t0 = time.perf_counter() self.device = device self.build = Build(device, policy(spec=precision), ln_fp32=ln_fp32, hidden_fp32=hidden_fp32, split=split, attn_fp32_acc=attn_fp32_acc, attn_matmul=attn_matmul) p = weights.params self.tables = P.step_tables(p, steps) self.buckets = agent_buckets() if hasattr(device, "compute_with_storage_grid_size") else () self.encoder = TtEncoder(self.build, p, nb_rows=self.buckets if self.build.compact_enc else ()) self.decoder = TtDecoder(self.build, p, self.tables, rows=(T.AGENTS,) + self.buckets) self.turn = TtTurnHead(self.build, p) self.runner = TraceRunner(device, num_command_queues=num_command_queues, name="diffusion-planner") warm = I.warmup_inputs() for name, shape in I.INPUT_SPECS.items(): dtype = "bfloat16" if name in I.BF16_INPUTS else "float32" self.runner.add_input(name, init=warm[name], dtype=dtype, layout=ttnn.TILE_LAYOUT) self.cs_pad = {} # INPUT_TRIM: zero columns that widen cs to STATE_COLS, per row count if I.CS_COLS < T.STATE_COLS: for r in (T.AGENTS,) + self.buckets: self.cs_pad[r] = self.build.upload(np.zeros((r, T.STATE_COLS - I.CS_COLS), np.float32), "float32") self.trim = I.CS_COLS < T.STATE_COLS self._zero_inputs: set = set() # INPUT_TRIM: inputs whose device buffer holds +0.0 (our last upload) self.runner.add_variant("plan", self._plan) for r in self.buckets: # COMPACT: one trace per agent bucket (the decoder on r rows) self.runner.add_variant(f"plan_r{r}", lambda ctx, r=r: self._plan(ctx, r)) self.debug = bool(debug) if self.debug: self.runner.add_input("dbg_x", init=np.zeros((1, 1, T.AGENTS, T.STATE_COLS), np.float32), dtype="float32", layout=ttnn.TILE_LAYOUT) self.runner.add_input("dbg_enc", init=np.zeros((1, 1, T.TOKENS, C.HIDDEN_DIM), np.float32), dtype=self.encoder.fstream, layout=ttnn.TILE_LAYOUT) for name in self._dbg_row_names(): self.runner.add_input(name, init=np.zeros((1, 1, 1, C.HIDDEN_DIM), np.float32), dtype="float32", layout=ttnn.TILE_LAYOUT) self.runner.add_variant("encoder_taps", self._encoder_taps) self.runner.add_variant("decode_once", self._decode_once) self.build_ms = (time.perf_counter() - t0) * 1e3 # ---- traced functions -------------------------------------------------------------------------------------- def _plan(self, ctx, rows: int = T.AGENTS): """The whole plan; ``rows`` < 352 (``COMPACT``): the decoder on the first ``rows`` agents (the inputs' leading rows, sliced in the trace).""" import ttnn enc = self.encoder.forward(ctx, nb=rows if rows != T.AGENTS else None) kv = self.decoder.cross_kv(enc) y0, cs, key = ctx["y0"], ctx["cs"], ctx["agent_key_row"] if rows != T.AGENTS: y0 = ttnn.slice(y0, [0, 0, 0, 0], [1, 1, rows, T.STATE_COLS]) cs = ttnn.slice(cs, [0, 0, 0, 0], [1, 1, rows, int(cs.shape[-1])]) key = ttnn.slice(key, [0, 0, 0, 0], [1, 1, 1, rows]) if int(cs.shape[-1]) < T.STATE_COLS: # INPUT_TRIM: widen the one-tile cs (tile-aligned concat) cs = ttnn.concat([cs, self.cs_pad[rows]], dim=-1) self_mask = A.expand_key_bias(key, rows, memory_config=self.build.attn_mem()) final, ego = self.decoder.solve(y0, cs, kv, self_mask) logit = self.turn(final, enc) return pack_outputs({"final_x0": final, "logit": logit, "ego_steps": ego}) def _encoder_taps(self, ctx): taps: Dict[str, Any] = {} self.encoder.forward(ctx, taps) return taps @staticmethod def _dbg_row_names() -> List[str]: return [f"dbg.{i}.{k}" for i in range(C.DIT_DEPTH) for k in STEP_KEYS] + ["dbg.final.g", "dbg.final.b"] def _decode_once(self, ctx): rows = {"blocks": [{k: ctx[f"dbg.{i}.{k}"] for k in STEP_KEYS} for i in range(C.DIT_DEPTH)], "g": ctx["dbg.final.g"], "b": ctx["dbg.final.b"]} kv = self.decoder.cross_kv(ctx["dbg_enc"]) self_mask = A.expand_key_bias(ctx["agent_key_row"], T.AGENTS) return self.decoder.evaluate(ctx["dbg_x"], rows, kv, self_mask) # ---- host side ------------------------------------------------------------------------------------------- def capture(self) -> None: self.runner.capture() def filter_inputs(self, inputs: Dict[str, Any]) -> Dict[str, Any]: """``INPUT_TRIM``: drop the inputs that are all +0.0 while their device buffer already holds the zeros of this model's previous upload (a value test, not a cache of the previous request: any non-zero array is always uploaded). Every input passed on is recorded as uploaded.""" if not self.trim: return inputs out = {} for name, a in inputs.items(): arr = a if isinstance(a, np.ndarray) else None zero = arr is not None and arr.dtype == np.float32 and not arr.view(np.uint32).any() if zero and name in self._zero_inputs: continue out[name] = a if zero: self._zero_inputs.add(name) else: self._zero_inputs.discard(name) return out def _run(self, variant: str, inputs: Dict[str, Any], eager: bool): inputs = self.filter_inputs(inputs) return self.runner.run_eager(variant, inputs=inputs) if eager else self.runner(variant, inputs=inputs) def rows_for(self, prepared: Any) -> int: """The decoder rows a plan needs (``COMPACT``): the smallest agent bucket holding the ego, every valid self-attention key and every emitted neighbour row; 352 without buckets or when none is large enough.""" return bucket_rows(prepared, self.buckets) def variant_for(self, prepared: Any) -> str: r = self.rows_for(prepared) return "plan" if r == T.AGENTS else f"plan_r{r}" def unpack(self, out: Dict[str, Any]) -> np.ndarray: """The readback's ``final_x0`` (``[R, 324]``) -> ``[321, 324]``; the rows past R (not computed: masked keys, never emitted) are zero.""" f = np.asarray(out["final_x0"], np.float32).reshape(-1, T.STATE_COLS) if f.shape[0] < T.AGENTS: f = np.concatenate([f, np.zeros((T.AGENTS - f.shape[0], T.STATE_COLS), np.float32)], 0) return f[:C.MAX_NUM_AGENTS] def forward(self, prepared: Any, *, ego_steps: bool = True, eager: bool = False, variant: Optional[str] = None) -> Dict[str, Any]: """One plan: upload, replay, read (``eager=True``: the same graph without the trace, for bring-up and replay-vs-eager checks). ``final_x0`` ``[321, 81, 4]`` (normalised, prefix-constrained), ``logit`` ``[5]``, ``denoising_steps``: the 11 iterates' ego rows as ``[1, 81, 4]`` arrays. ``variant``: force a plan variant (default :meth:`variant_for`).""" out = self._run(variant or self.variant_for(prepared), I.plan_inputs(prepared), eager) final = self.unpack(out) res = {"final_x0": final.reshape(C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM).astype(np.float32), "logit": out["logit"].reshape(-1)[:C.TURN_INDICATOR_OUTPUT_DIM].astype(np.float32)} if ego_steps: steps = out["ego_steps"].reshape(-1, T.STATE_COLS) res["denoising_steps"] = [s.reshape(1, C.OUTPUT_T + 1, C.POSE_DIM).astype(np.float32) for s in steps] return res def encoder_taps(self, prepared: Any, *, eager: bool = False) -> Dict[str, np.ndarray]: """Replay of ``encoder_taps``: ``{tap: array}`` with the reference's shapes (``[E, 256]`` category rows, ``[E, 64, 128]`` mixer taps, ``[564, 256]`` token / fusion / encoding rows).""" if not self.debug: raise RuntimeError("encoder_taps needs TtDiffusionPlanner(debug=True)") raw = self._run("encoder_taps", I.plan_inputs(prepared), eager) out = {} for name, a in raw.items(): a = np.asarray(a, np.float32) if name.endswith((".pre", ".mixer")): out[name] = a.reshape(-1, C.MIXER_TOKENS, C.MIXER_CHANNELS) elif name in ("enc.tokens", "enc.encoding") or name.startswith("enc.fusion."): out[name] = a.reshape(-1, C.HIDDEN_DIM)[:T.TOKENS_REAL] else: out[name] = a.reshape(-1, C.HIDDEN_DIM) return out def decode_once(self, prepared: Any, x: np.ndarray, t: float, *, encoding: np.ndarray, eager: bool = False) -> np.ndarray: """Replay of ``decode_once`` at evaluation time ``t`` (one of ``tables.eval_times``) on the teacher-forced ``x`` ``[321, 81, 4]`` (prefix-constrained) and ``encoding`` ``[564, 256]`` -> ``[321, 81, 4]`` (the t = 0 slot is 0: the masked projection).""" if not self.debug: raise RuntimeError("decode_once needs TtDiffusionPlanner(debug=True)") k = int(np.argmin([abs(float(t) - e) for e in self.tables.eval_times])) if abs(float(t) - self.tables.eval_times[k]) > 1e-6: raise ValueError(f"t={t} is not an evaluation time of the solver plan {self.tables.eval_times}") base = I.plan_inputs(prepared) enc = P.pad_rows(np.asarray(encoding, np.float32).reshape(T.TOKENS_REAL, C.HIDDEN_DIM), T.TOKENS) inputs = {"agent_key_row": base["agent_key_row"], "dbg_x": I.decoder_state(x, x[:, 0]), "dbg_enc": enc.reshape(1, 1, T.TOKENS, C.HIDDEN_DIM)} for i, blk in enumerate(self.tables.blocks[k]): for key in STEP_KEYS: inputs[f"dbg.{i}.{key}"] = np.asarray(blk[key], np.float32).reshape(1, 1, 1, -1) inputs["dbg.final.g"] = np.asarray(self.tables.final[k]["g"], np.float32).reshape(1, 1, 1, -1) inputs["dbg.final.b"] = np.asarray(self.tables.final[k]["b"], np.float32).reshape(1, 1, 1, -1) out = self._run("decode_once", inputs, eager) flat = np.asarray(out, np.float32).reshape(T.AGENTS, T.STATE_COLS)[:C.MAX_NUM_AGENTS] return flat.reshape(C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM) def trace_buffers_mb(self) -> Optional[float]: """DRAM held by the captured traces' command buffers: the allocated bytes of the device's TRACE region (``ttnn.get_memory_view``; all banks), None when the view is unavailable.""" import ttnn try: view = ttnn.get_memory_view(self.device, ttnn.BufferType.TRACE) return round(view.num_banks * view.total_bytes_allocated_per_bank / 2 ** 20, 2) except (AttributeError, RuntimeError, TypeError): return None def describe(self) -> Dict[str, Any]: return {"trace": self.runner.describe(), "trace_buffers_mb": self.trace_buffers_mb(), "precision": self.build.policy.describe(), "options": self.build.options(), "uploaded_mb": round(self.build.uploaded_bytes / 2 ** 20, 2), "build_ms": round(self.build_ms, 1), "tokens": T.TOKENS, "agents": T.AGENTS, "agent_buckets": list(self.buckets), "nfe": self.tables.nfe} def release(self) -> None: self.runner.release()