#!/usr/bin/env python3 # SPDX-License-Identifier: Apache-2.0 """Device profile of diffusion-planner-p150: one eager plan and one traced replay between Tracy signposts. ROOT=/home/ubuntu/experiments/tt-models $ROOT/bin/devrun -t 3600 -- python -m tracy -r -p -v --op-support-count 16000 --no-web-server \\ -o $ROOT/generated/profiler/diffusion-planner_baseline code/scripts/profile_ops.py tt-perf-report --start-signpost trace --end-signpost trace_end python -m tt_diffusion_planner.ttaw.profiling --start trace --end trace_end Sections (signposts): ``eager`` .. ``eager_end``, one eager run of the ``plan`` variant after the warm-up (``TraceRunner.run_eager``: the same graph and programs as the trace, nothing compiled), and ``trace`` .. ``trace_end``, ``--replays`` replays of the captured ``plan`` trace (the served device path) on the uploaded sample. The device profiler buffer is flushed (``ttnn.ReadDeviceProfiler``) before and after each section. Inside the eager section two levels of signposts attribute every op (and, through the identical op order, every replayed op): - ``m:``: the stage of the plan: ``enc..pre`` / ``.mix`` / ``.head`` (the six mixer trunks, then pool + entity head), ``enc.static``, ``enc.``, ``enc.tokens`` (concat, validity, position embedding), ``mask`` (key-bias expansion), ``enc.fusion.attn`` / ``.mlp``, ``enc.final_ln``, ``dec.cross_kv`` (hoisted cross K / V), ``dec.e.preproj``, ``dec.e.b.attn`` / ``.mlp1`` / ``.cross`` / ``.mlp2``, ``dec.e.final`` (evaluation k = 0..10, DiT block i), ``dec.e.solver`` (the DPM-Solver++(2M) update, the prefix constraint and the ego-row slice of the next iterate), ``turn``, ``pack``; - ``c:``: the layer kind: ``split`` (split hi / lo matmul, ``SplitLinear``), ``linear`` (one ``ttnn.linear``), ``ln32`` (fp32 LayerNorm decomposition), ``ln`` (``ttnn.layer_norm``), ``attn_mm`` (fp32 matmul attention), ``sdpa``, ``heads`` (head split / merge), ``mask``, and ``glue`` for every op outside a layer (residual adds, adaLN gates, mixer transposes, solver updates, slices, concats). ``--no-layer-signposts`` turns both levels off. One plan issues several thousand programs, more than the profiler's default 1000-program buffer, so run it with ``--op-support-count`` above the program count of the largest section (the warm-up before the first flush included). The output folder must be absolute and outside the bundle (PLAN.md 5.1). """ from __future__ import annotations import argparse import functools import json import time from pathlib import Path from typing import Any, Callable, Dict, List import numpy as np from tt_diffusion_planner import DiffusionPlanner from tt_diffusion_planner.ttaw.profiling import read_device_profiler, signpost, signposted SAMPLE = Path(__file__).resolve().parents[1] / "tt_diffusion_planner" / "samples" / "kashiwanoha_dense.npz" class _Proxy: """A callable stand-in for a layer object that emits a signpost (``label()``), then calls the layer.""" def __init__(self, inner: Any, label: Callable[[], str]): self.inner, self._label = inner, label def __call__(self, *args, **kwargs): signpost(self._label()) return self.inner(*args, **kwargs) def __getattr__(self, name: str): return getattr(self.inner, name) def install_signposts(tt) -> Callable[[], None]: """Signposts on a ``tt.model.TtDiffusionPlanner`` (instance attributes, layer classes and the attention helpers); returns the function that removes them again.""" from tt_diffusion_planner.tt import layers as L from tt_diffusion_planner.tt import model as M from tt_diffusion_planner.ttaw.ops import attention as A undo: List[Callable[[], None]] = [] state = {"k": -1, "ln": {}} stack: List[str] = [] def set_attr(obj, name, value): had = name in vars(obj) old = vars(obj).get(name) setattr(obj, name, value) undo.append(lambda: setattr(obj, name, old) if had else delattr(obj, name)) def set_item(d, key, value): old = d[key] d[key] = value undo.append(lambda: d.__setitem__(key, old)) # ---- layer kinds (class level, with a stack so nested layers restore the outer kind) --------------------- def kind_wrap(owner, name, kind_of): fn = getattr(owner, name) @functools.wraps(fn) def inner(*args, **kwargs): kind = kind_of(*args) stack.append(kind) signpost(f"c:{kind}") try: return fn(*args, **kwargs) finally: stack.pop() signpost(f"c:{stack[-1] if stack else 'glue'}") setattr(owner, name, inner) undo.append(lambda: setattr(owner, name, fn)) kind_wrap(L.Linear, "__call__", lambda *a: "linear") kind_wrap(L.SplitLinear, "__call__", lambda *a: "split") kind_wrap(L.LayerNorm, "__call__", lambda self, *a: "ln32" if self.mode == "fp32" else "ln") kind_wrap(A, "attention_matmul", lambda *a: "attn_mm") kind_wrap(A, "sdpa", lambda *a: "sdpa") for name in ("split_qkv", "split_q_kv", "split_heads", "merge_heads"): kind_wrap(A, name, lambda *a: "heads") kind_wrap(A, "expand_key_bias", lambda *a: "mask") # ---- stages ------------------------------------------------------------------------------------------------ def method_wrap(obj, name, before=None, after=None): fn = getattr(obj, name) @functools.wraps(fn) def inner(*args, **kwargs): if before: signpost(before()) out = fn(*args, **kwargs) if after: signpost(after()) return out set_attr(obj, name, inner) enc, dec = tt.encoder, tt.decoder for cat, trunk in enc.trunks.items(): method_wrap(trunk, "pre", before=lambda c=cat: f"m:enc.{c}.pre") method_wrap(trunk, "mix", before=lambda c=cat: f"m:enc.{c}.mix") method_wrap(trunk, "pool", before=lambda c=cat: f"m:enc.{c}.head") set_attr(enc, "static1", _Proxy(enc.static1, lambda: "m:enc.static")) for cat in list(enc.small): set_item(enc.small, cat, _Proxy(enc.small[cat], lambda c=cat: f"m:enc.{c}")) set_attr(enc, "pad_tokens", _Proxy(enc.pad_tokens, lambda: "m:enc.tokens")) for i, blk in enumerate(enc.blocks): set_item(blk, "kv", _Proxy(blk["kv"], lambda i=i: f"m:enc.fusion{i}.attn")) set_item(blk, "n2", _Proxy(blk["n2"], lambda i=i: f"m:enc.fusion{i}.mlp")) set_attr(enc, "final_norm", _Proxy(enc.final_norm, lambda: "m:enc.final_ln")) method_wrap(dec, "cross_kv", before=lambda: "m:dec.cross_kv") method_wrap(dec, "solve", before=lambda: "m:dec.solve") def next_eval(): state["k"] += 1 state["ln"] = {} return f"m:dec.e{state['k']}.preproj" method_wrap(dec, "evaluate", before=next_eval, after=lambda: f"m:dec.e{state['k']}.solver") for i, blk in enumerate(dec.blocks): def ln_label(i=i): n = state["ln"][i] = state["ln"].get(i, 0) + 1 return f"m:dec.e{state['k']}.b{i}.{'attn' if n % 2 else 'mlp1'}" set_item(blk, "ln", _Proxy(blk["ln"], ln_label)) set_item(blk, "n3", _Proxy(blk["n3"], lambda i=i: f"m:dec.e{state['k']}.b{i}.cross")) set_item(blk, "n4", _Proxy(blk["n4"], lambda i=i: f"m:dec.e{state['k']}.b{i}.mlp2")) set_attr(dec, "fin_ln", _Proxy(dec.fin_ln, lambda: f"m:dec.e{state['k']}.final")) set_attr(tt, "turn", _Proxy(tt.turn, lambda: "m:turn")) pack = M.pack_outputs def pack_wrap(*args, **kwargs): signpost("m:pack") return pack(*args, **kwargs) M.pack_outputs = pack_wrap undo.append(lambda: setattr(M, "pack_outputs", pack)) # the mask expansion is attributed to its own stage (encoder fusion mask, decoder self-attention mask) expand = A.expand_key_bias def expand_wrap(*args, **kwargs): signpost("m:mask") return expand(*args, **kwargs) A.expand_key_bias = expand_wrap undo.append(lambda: setattr(A, "expand_key_bias", expand)) def remove() -> None: for fn in reversed(undo): fn() return remove def main() -> None: ap = argparse.ArgumentParser(description=__doc__.split("\n")[0]) ap.add_argument("--input", default=str(SAMPLE)) ap.add_argument("--dispatch", default=None, choices=["eth", "worker"]) ap.add_argument("--num-cqs", type=int, default=None, choices=[1, 2]) ap.add_argument("--replays", type=int, default=1, help="replays inside the trace section") ap.add_argument("--no-eager", dest="eager", action="store_false") ap.add_argument("--no-layer-signposts", dest="layers", action="store_false") ap.add_argument("--json", default=None, help="write the run description (device, trace, timings) here") a = ap.parse_args() import ttnn from tt_diffusion_planner.host import pipeline as hp from tt_diffusion_planner.reference import config as C from tt_diffusion_planner.reference.weights import find_weights_dir from tt_diffusion_planner.tt import inputs as I from tt_diffusion_planner.ttaw.io import load_named_arrays wd = find_weights_dir() t0 = time.perf_counter() with DiffusionPlanner.from_pretrained(dispatch=a.dispatch, num_command_queues=a.num_cqs, weights_dir=str(wd) if wd else None) as model: tt, runner, dev = model.tt, model.runner, model.device print("loaded in %.1f s:" % (time.perf_counter() - t0), json.dumps(model.device_info), flush=True) raw = load_named_arrays(a.input, C.INPUT_SCHEMA) prep = hp.prepare(raw, model.normalization.observation) inputs = I.plan_inputs(prep) variant = tt.variant_for(prep) # COMPACT: the served bucket trace print("variant", variant, flush=True) served = model(inputs=raw) # one served plan: upload + replay + readback ttnn.synchronize_device(dev) read_device_profiler(dev) # warm-up / capture / first plan out of the buffer timings: Dict[str, float] = {} if a.eager: remove = install_signposts(tt) if a.layers else (lambda: None) try: t1 = time.perf_counter() with signposted("eager"): eager = runner.run_eager(variant, inputs=inputs) ttnn.synchronize_device(dev) timings["eager_ms"] = (time.perf_counter() - t1) * 1e3 finally: remove() read_device_profiler(dev) final = np.asarray(eager["final_x0"], np.float32) print("eager final_x0 finite:", bool(np.isfinite(final).all()), flush=True) runner.upload(inputs) ttnn.synchronize_device(dev) read_device_profiler(dev) t1 = time.perf_counter() with signposted("trace"): runner.replay(variant, n=a.replays) ttnn.synchronize_device(dev) timings["trace_ms"] = (time.perf_counter() - t1) * 1e3 / a.replays read_device_profiler(dev) out = runner.read(variant) same = bool(np.array_equal(out["final_x0"], eager["final_x0"])) if a.eager else None desc = {"device": model.device_info, "timings_ms": timings, "replays": a.replays, "replay_equals_eager": same, "turn_command": int(served.turn_indicator["command"]), "options": tt.build.options(), "trace": runner.describe()} print("profile run:", json.dumps(desc, default=str), flush=True) if a.json: Path(a.json).write_text(json.dumps(desc, indent=1, default=str) + "\n") if __name__ == "__main__": main()