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