diffusion-planner-p150 / code /scripts /profile_ops.py
changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
11.9 kB
#!/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 <ops_perf_results_*.csv> --start-signpost trace --end-signpost trace_end
python -m tt_diffusion_planner.ttaw.profiling <the same csv or its directory> --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:<stage>``: the stage of the plan: ``enc.<category>.pre`` / ``.mix`` / ``.head`` (the six mixer trunks, then
pool + entity head), ``enc.static``, ``enc.<goal|ego_shape|turn>``, ``enc.tokens`` (concat, validity, position
embedding), ``mask`` (key-bias expansion), ``enc.fusion<i>.attn`` / ``.mlp``, ``enc.final_ln``, ``dec.cross_kv``
(hoisted cross K / V), ``dec.e<k>.preproj``, ``dec.e<k>.b<i>.attn`` / ``.mlp1`` / ``.cross`` / ``.mlp2``,
``dec.e<k>.final`` (evaluation k = 0..10, DiT block i), ``dec.e<k>.solver`` (the DPM-Solver++(2M) update, the
prefix constraint and the ego-row slice of the next iterate), ``turn``, ``pack``;
- ``c:<kind>``: 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()