Download code/tt_diffusion_planner/ttaw/profiling.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 18.6 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/ttaw/profiling.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/ttaw/profiling.py
-
curl -L -o profiling.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/ttaw/profiling.py
18.6 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """C07: measurement helpers -- the numbers OPT_BASELINE.md / OPT_REPORT.md / the model card quote. | |
| - :class:`StageBench` collects wall times per stage (``host_in``, ``h2d``, ``trace``, ``b2b``, ``d2h``, ``post``, | |
| ``e2e``) and reports p50 / p99 / mean / min / max in ms (RP section 3, rf-detr ``bench_breakdown``). | |
| - :func:`bench_trace_runner` measures those stages on a :class:`.trace.TraceRunner` variant. | |
| - :func:`signpost` / :func:`signposted` mark Tracy ranges (no-ops when the ``tracy`` module is absent), so | |
| ``tt-perf-report --start-signpost trace --end-signpost trace_end`` and :func:`summarize_ops` can cut the CSV. | |
| - :func:`summarize_ops` reads a device-profiler ``ops_perf_results*.csv`` (Tracy ``-r``) or the C++ | |
| ``cpp_device_perf_report.csv`` and returns op count, kernel sum, op-to-op gaps, span, per-op-code totals and the | |
| math-fidelity histogram. CLI: ``python -m <pkg>.ttaw.profiling <csv or dir> [--start trace] [--json out.json]``. | |
| - :class:`AiclkSampler` samples the chip clock (sysfs ``tt_aiclk``) and hwmon power / temperature in a background | |
| thread during a bench: AICLK sags from 1350 MHz under load (RP section 3), which explains span-vs-wall gaps. | |
| Profiling runs go through ``bin/devrun`` like any device job, with an absolute output directory under | |
| ``generated/profiler/<model>_<tag>`` (PLAN.md section 5.1). | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import contextlib | |
| import csv | |
| import glob | |
| import json | |
| import logging | |
| import os | |
| import statistics | |
| import sys | |
| import threading | |
| import time | |
| from collections import defaultdict | |
| from dataclasses import dataclass, field | |
| from pathlib import Path | |
| from typing import Any, Callable, Dict, Iterator, List, Mapping, Optional, Sequence, Union | |
| import numpy as np | |
| __all__ = [ | |
| "STAGES", | |
| "StageBench", | |
| "time_b2b", | |
| "bench_trace_runner", | |
| "signpost", | |
| "signposted", | |
| "read_device_profiler", | |
| "find_ops_csv", | |
| "OpsSummary", | |
| "summarize_ops", | |
| "AiclkSampler", | |
| ] | |
| log = logging.getLogger(__name__) | |
| STAGES = ("host_in", "h2d", "trace", "b2b", "d2h", "post", "e2e") | |
| DEFAULT_AICLK_MHZ = 1350.0 | |
| # ------------------------------------------------------------------------------------------- stage bench | |
| def _stats(samples: Sequence[float]) -> Dict[str, float]: | |
| a = np.asarray(samples, np.float64) | |
| return {"n": int(a.size), "p50": float(np.percentile(a, 50)), "p99": float(np.percentile(a, 99)), | |
| "mean": float(a.mean()), "min": float(a.min()), "max": float(a.max())} | |
| class StageBench: | |
| """Per-stage wall-time samples in milliseconds. | |
| :: | |
| bench = StageBench("centerpoint eth/1cq") | |
| for _ in range(100): | |
| with bench.stage("e2e"): | |
| with bench.stage("host_in"): host = prepare(points) | |
| ... | |
| print(bench.table()); bench.save_json("logs/centerpoint/bench_baseline.json", config=device_info) | |
| """ | |
| def __init__(self, name: str = ""): | |
| self.name = name | |
| self.samples: Dict[str, List[float]] = {} | |
| def stage(self, name: str) -> Iterator[None]: | |
| t0 = time.perf_counter() | |
| try: | |
| yield | |
| finally: | |
| self.add(name, (time.perf_counter() - t0) * 1e3) | |
| def add(self, name: str, ms: float) -> None: | |
| self.samples.setdefault(name, []).append(float(ms)) | |
| def summary(self) -> Dict[str, Dict[str, float]]: | |
| """``{stage: {n, p50, p99, mean, min, max}}`` in ms; known stages first, in pipeline order.""" | |
| order = [s for s in STAGES if s in self.samples] + [s for s in self.samples if s not in STAGES] | |
| return {s: _stats(self.samples[s]) for s in order if self.samples[s]} | |
| def table(self) -> str: | |
| rows = [f"| stage | n | p50 ms | p99 ms | mean ms | min ms |", "|---|---:|---:|---:|---:|---:|"] | |
| for s, st in self.summary().items(): | |
| rows.append(f"| {s} | {st['n']} | {st['p50']:.3f} | {st['p99']:.3f} | {st['mean']:.3f} | {st['min']:.3f} |") | |
| return (f"**{self.name}**\n\n" if self.name else "") + "\n".join(rows) | |
| def to_dict(self, **extra: Any) -> Dict[str, Any]: | |
| return {"name": self.name, "stages_ms": self.summary(), **extra} | |
| def save_json(self, path: Union[str, os.PathLike], **extra: Any) -> Path: | |
| path = Path(path) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| path.write_text(json.dumps(self.to_dict(**extra), indent=1, default=str) + "\n") | |
| return path | |
| def time_b2b(enqueue: Callable[[], Any], sync: Callable[[], Any], n: int = 100, warmup: int = 5) -> float: | |
| """Back-to-back device time per iteration (ms): ``warmup`` + ``n`` non-blocking ``enqueue()`` calls, one | |
| ``sync()`` each round. With a traced ``enqueue`` this is the device time per forward.""" | |
| for _ in range(warmup): | |
| enqueue() | |
| sync() | |
| t0 = time.perf_counter() | |
| for _ in range(n): | |
| enqueue() | |
| sync() | |
| return (time.perf_counter() - t0) * 1e3 / n | |
| def bench_trace_runner(runner, variant: Optional[str], inputs: Mapping[str, Any], *, | |
| params: Optional[Mapping[str, Any]] = None, iters: int = 100, warmup: int = 10, | |
| post: Optional[Callable[[Any], Any]] = None, b2b_iters: Optional[int] = None, | |
| name: str = "") -> StageBench: | |
| """Stage breakdown of one :class:`.trace.TraceRunner` variant (medians are the headline numbers): | |
| ``host_in`` host tensor conversion, ``h2d`` upload + sync, ``trace`` one replay + sync, ``d2h`` read, | |
| ``post`` the optional ``post(outputs)`` callback, ``e2e`` ``runner(variant, inputs)`` + ``post``, and ``b2b`` | |
| back-to-back replays (device time per forward).""" | |
| import ttnn | |
| from .tensors import to_host_tensor | |
| dev = runner.device | |
| bench = StageBench(name or f"{runner.name}/{variant}") | |
| slots = {k: runner._slot(k, "input") for k in inputs} | |
| for _ in range(warmup): | |
| out = runner(variant, inputs, params) | |
| if post is not None: | |
| post(out) | |
| for _ in range(iters): | |
| t0 = time.perf_counter() | |
| host = {k: to_host_tensor(v, slots[k].dtype, slots[k].layout, shape=slots[k].shape) for k, v in inputs.items()} | |
| t1 = time.perf_counter() | |
| runner.upload(host, params) | |
| ttnn.synchronize_device(dev) | |
| t2 = time.perf_counter() | |
| runner.replay(variant) | |
| ttnn.synchronize_device(dev) | |
| t3 = time.perf_counter() | |
| out = runner.read(variant) | |
| t4 = time.perf_counter() | |
| if post is not None: | |
| post(out) | |
| t5 = time.perf_counter() | |
| for stage, a, b in (("host_in", t0, t1), ("h2d", t1, t2), ("trace", t2, t3), ("d2h", t3, t4)): | |
| bench.add(stage, (b - a) * 1e3) | |
| if post is not None: | |
| bench.add("post", (t5 - t4) * 1e3) | |
| for _ in range(iters): | |
| with bench.stage("e2e"): | |
| out = runner(variant, inputs, params) | |
| if post is not None: | |
| post(out) | |
| bench.add("b2b", time_b2b(lambda: runner.replay(variant), lambda: ttnn.synchronize_device(dev), | |
| n=b2b_iters or iters)) | |
| return bench | |
| # --------------------------------------------------------------------------------------------- signposts | |
| def signpost(name: str, message: Optional[str] = None) -> bool: | |
| """Emit a Tracy signpost (a row in the ops CSV); returns False (no-op) when ``tracy`` is not importable.""" | |
| try: | |
| from tracy import signpost as _signpost | |
| except ImportError: | |
| return False | |
| _signpost(header=name, message=message) | |
| return True | |
| def signposted(name: str) -> Iterator[None]: | |
| """``signpost(name)`` ... ``signpost(name + "_end")`` around the block.""" | |
| signpost(name) | |
| try: | |
| yield | |
| finally: | |
| signpost(f"{name}_end") | |
| def read_device_profiler(device) -> bool: | |
| """Flush the device profiler buffer (``ttnn.ReadDeviceProfiler``); call it per trace segment, the buffer holds | |
| about 1000 ops. Returns False when the API is absent.""" | |
| import ttnn | |
| fn = getattr(ttnn, "ReadDeviceProfiler", None) | |
| if fn is None: | |
| return False | |
| fn(device) | |
| return True | |
| # ------------------------------------------------------------------------------------- ops CSV summary | |
| def find_ops_csv(path: Union[str, os.PathLike]) -> Path: | |
| """A CSV file as-is, or the newest ``ops_perf_results*.csv`` (else ``cpp_device_perf_report.csv``) below a dir.""" | |
| p = Path(path) | |
| if p.is_file(): | |
| return p | |
| for pattern in ("ops_perf_results*.csv", "cpp_device_perf_report.csv"): | |
| found = sorted(glob.glob(str(p / "**" / pattern), recursive=True), key=os.path.getmtime) | |
| if found: | |
| return Path(found[-1]) | |
| raise FileNotFoundError(f"no ops_perf_results*.csv or cpp_device_perf_report.csv under {p}") | |
| def _num(row: Mapping[str, str], key: str) -> float: | |
| value = (row.get(key) or "").strip() | |
| try: | |
| return float(value) if value else 0.0 | |
| except ValueError: | |
| return 0.0 | |
| def _chip_freq_mhz(csv_path: Path) -> Optional[float]: | |
| """CHIP_FREQ[MHz] from a ``profile_log_device.csv`` header near the ops CSV, if any.""" | |
| for cand in [csv_path.parent / "profile_log_device.csv", csv_path.parent.parent / "profile_log_device.csv", | |
| csv_path.parent / ".logs" / "profile_log_device.csv"]: | |
| try: | |
| head = cand.read_text(errors="ignore").splitlines()[0] | |
| except (OSError, IndexError): | |
| continue | |
| for part in head.split(","): | |
| if "CHIP_FREQ" in part and ":" in part: | |
| try: | |
| return float(part.split(":")[1].strip()) | |
| except ValueError: | |
| pass | |
| return None | |
| class OpsSummary: | |
| """Summary of one profiled section (times in microseconds).""" | |
| source: str | |
| section: str | |
| ops: int | |
| kernel_sum_us: float | |
| fw_sum_us: float | |
| op2op_sum_us: float | |
| span_us: Optional[float] | |
| freq_mhz: float | |
| by_op: List[Dict[str, Any]] = field(default_factory=list) | |
| fidelity: Dict[str, int] = field(default_factory=dict) | |
| top_gaps: List[Dict[str, Any]] = field(default_factory=list) | |
| def to_dict(self) -> Dict[str, Any]: | |
| return {k: getattr(self, k) for k in self.__dataclass_fields__} | |
| def table(self, top: int = 20) -> str: | |
| span = "n/a" if self.span_us is None else f"{self.span_us:.1f}" | |
| lines = [f"{self.source} [{self.section}]", | |
| f"ops {self.ops} kernel_sum {self.kernel_sum_us:.1f} us op2op_sum {self.op2op_sum_us:.1f} us " | |
| f"span {span} us (@{self.freq_mhz:.0f} MHz) fidelity {self.fidelity}", | |
| f"{'op code':55s} {'count':>6s} {'kernel us':>11s} {'share':>7s}"] | |
| for row in self.by_op[:top]: | |
| lines.append(f"{row['op'][:55]:55s} {row['count']:6d} {row['kernel_us']:11.1f} {row['share']:6.1%}") | |
| return "\n".join(lines) | |
| def summarize_ops(source: Union[str, os.PathLike, Sequence[Mapping[str, str]]], *, start: Optional[str] = "trace", | |
| end: Optional[str] = None, last_replay_session: Optional[bool] = None, | |
| freq_mhz: Optional[float] = None, top_gaps: int = 10) -> OpsSummary: | |
| """Summarize device ops between signposts ``start`` and ``end`` (default: the next signpost). | |
| ``source``: an ``ops_perf_results*.csv`` / ``cpp_device_perf_report.csv`` path, a directory holding one, or | |
| already-parsed rows. ``start=None`` takes every op. ``last_replay_session`` (default: on when the rows carry | |
| ``METAL TRACE REPLAY SESSION ID`` values) keeps only the last trace replay session, the way ``prof_cpp.py`` | |
| did. Span uses the FW start/end cycles at ``freq_mhz`` (default: the CHIP_FREQ of ``profile_log_device.csv``, | |
| else 1350 MHz; AICLK sags under load, so treat span as approximate).""" | |
| if isinstance(source, (str, os.PathLike)): | |
| path = find_ops_csv(source) | |
| with open(path, newline="") as f: | |
| rows = list(csv.DictReader(f)) | |
| label = str(path) | |
| freq = freq_mhz or _chip_freq_mhz(path) or DEFAULT_AICLK_MHZ | |
| else: | |
| rows, label, freq = list(source), "<rows>", freq_mhz or DEFAULT_AICLK_MHZ | |
| name_key = "OP CODE" if rows and "OP CODE" in rows[0] else "OP NAME" | |
| section = "all" | |
| if start is not None and any((r.get("OP TYPE") or "") == "signpost" for r in rows): | |
| selected, on = [], False | |
| for r in rows: | |
| if (r.get("OP TYPE") or "") == "signpost": | |
| code = r.get(name_key, "") | |
| if on and (end is None or code == end): | |
| break | |
| on = on or code == start | |
| continue | |
| if on: | |
| selected.append(r) | |
| rows, section = selected, f"{start}..{end or 'next signpost'}" | |
| else: | |
| rows = [r for r in rows if (r.get("OP TYPE") or "") != "signpost"] | |
| sessions = [r.get("METAL TRACE REPLAY SESSION ID", "") for r in rows] | |
| if last_replay_session is None: | |
| last_replay_session = any(s.strip() for s in sessions) | |
| if last_replay_session and rows: | |
| traced = [r for r in rows if (r.get("METAL TRACE ID") or "").strip()] | |
| if traced: | |
| last = traced[-1].get("METAL TRACE REPLAY SESSION ID") | |
| rows = [r for r in traced if r.get("METAL TRACE REPLAY SESSION ID") == last] | |
| section += f" (replay session {last})" | |
| kernel = [_num(r, "DEVICE KERNEL DURATION [ns]") / 1e3 for r in rows] | |
| fw = [_num(r, "DEVICE FW DURATION [ns]") / 1e3 for r in rows] | |
| gaps = [_num(r, "OP TO OP LATENCY [ns]") / 1e3 for r in rows] | |
| starts = [_num(r, "DEVICE FW START CYCLE") for r in rows] | |
| ends = [_num(r, "DEVICE FW END CYCLE") for r in rows] | |
| span = (max(ends) - min(s for s in starts if s > 0)) / freq if rows and any(starts) and any(ends) else None | |
| agg: Dict[str, List[float]] = defaultdict(lambda: [0, 0.0]) | |
| fidelity: Dict[str, int] = defaultdict(int) | |
| for r, k in zip(rows, kernel): | |
| agg[r.get(name_key, "?")][0] += 1 | |
| agg[r.get(name_key, "?")][1] += k | |
| fid = (r.get("MATH FIDELITY") or "").strip() | |
| if fid: | |
| fidelity[fid] += 1 | |
| total = sum(kernel) or 1.0 | |
| by_op = [{"op": op, "count": int(c), "kernel_us": t, "share": t / total} | |
| for op, (c, t) in sorted(agg.items(), key=lambda kv: -kv[1][1])] | |
| order = sorted(range(1, len(rows)), key=lambda i: -gaps[i])[:top_gaps] | |
| worst = [{"index": i, "op": rows[i].get(name_key, "?"), "gap_us": gaps[i], "after": rows[i - 1].get(name_key, "?")} | |
| for i in order] | |
| return OpsSummary(label, section, len(rows), sum(kernel), sum(fw), sum(gaps[1:]), span, freq, by_op, | |
| dict(fidelity), worst) | |
| # ------------------------------------------------------------------------------------------------ AICLK | |
| class AiclkSampler: | |
| """Background sampling of ``tt_aiclk`` (MHz) and hwmon power (W) / temperature (C) of one chip. | |
| :: | |
| with AiclkSampler(chip=0, interval_s=0.05) as clk: | |
| run_bench() | |
| print(clk.summary()) # {"aiclk_mhz": {"median": 1302, "min": 1206, ...}, "power_w": {...}, ...} | |
| Reads sysfs only (read-only, works while another process holds the device); ``available`` is False on hosts | |
| without the Tenstorrent KMD.""" | |
| def __init__(self, chip: int = 0, interval_s: float = 0.05, *, root: str = "/sys/class/tenstorrent"): | |
| self.interval_s = float(interval_s) | |
| base = Path(root) / f"tenstorrent!{chip}" | |
| self._aiclk = base / "tt_aiclk" | |
| hwmon = sorted(glob.glob(str(base / "device" / "hwmon" / "hwmon*"))) | |
| self._hwmon = Path(hwmon[0]) if hwmon else None | |
| self.available = self._aiclk.is_file() | |
| self.samples: Dict[str, List[float]] = {"aiclk_mhz": [], "power_w": [], "temp_c": []} | |
| self._stop = threading.Event() | |
| self._thread: Optional[threading.Thread] = None | |
| def _read(path: Optional[Path], scale: float) -> Optional[float]: | |
| if path is None: | |
| return None | |
| try: | |
| return float(path.read_text().split()[0]) * scale | |
| except (OSError, ValueError, IndexError): | |
| return None | |
| def sample_once(self) -> Dict[str, Optional[float]]: | |
| hw = self._hwmon | |
| values = {"aiclk_mhz": self._read(self._aiclk, 1.0), | |
| "power_w": self._read(hw / "power1_input" if hw else None, 1e-6), | |
| "temp_c": self._read(hw / "temp1_input" if hw else None, 1e-3)} | |
| for k, v in values.items(): | |
| if v is not None: | |
| self.samples[k].append(v) | |
| return values | |
| def _loop(self) -> None: | |
| while not self._stop.is_set(): | |
| self.sample_once() | |
| self._stop.wait(self.interval_s) | |
| def start(self) -> "AiclkSampler": | |
| if self.available and self._thread is None: | |
| self._stop.clear() | |
| self._thread = threading.Thread(target=self._loop, name="aiclk-sampler", daemon=True) | |
| self._thread.start() | |
| return self | |
| def stop(self) -> None: | |
| if self._thread is not None: | |
| self._stop.set() | |
| self._thread.join(timeout=5) | |
| self._thread = None | |
| def __enter__(self) -> "AiclkSampler": | |
| return self.start() | |
| def __exit__(self, *exc) -> None: | |
| self.stop() | |
| def summary(self) -> Dict[str, Any]: | |
| out: Dict[str, Any] = {"available": self.available} | |
| for k, v in self.samples.items(): | |
| if v: | |
| out[k] = {"median": statistics.median(v), "min": min(v), "max": max(v), "n": len(v)} | |
| return out | |
| def main(argv: Optional[Sequence[str]] = None) -> int: | |
| """``python -m <pkg>.ttaw.profiling <csv|dir> [--start trace] [--end trace_end] [--json out] [--top 25]``.""" | |
| ap = argparse.ArgumentParser(description="Summarize a device-profiler ops CSV") | |
| ap.add_argument("source") | |
| ap.add_argument("--start", default="trace", help="signpost that opens the section ('' = all ops)") | |
| ap.add_argument("--end", default=None) | |
| ap.add_argument("--freq-mhz", type=float, default=None) | |
| ap.add_argument("--top", type=int, default=25) | |
| ap.add_argument("--json", default=None) | |
| a = ap.parse_args(argv) | |
| summary = summarize_ops(a.source, start=a.start or None, end=a.end, freq_mhz=a.freq_mhz) | |
| print(summary.table(a.top)) | |
| if a.json: | |
| Path(a.json).write_text(json.dumps(summary.to_dict(), indent=1) + "\n") | |
| return 0 | |
| if __name__ == "__main__": | |
| sys.exit(main()) | |