changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
15.8 kB
# 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.<i>.<key>``, ``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()