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