26b071c6e1
- _cudagraph_plugin: warm up autotune/JIT explicitly before graph capture (internal 5-iter warmup is too short), tag fallback markers with the failing phase; first-run latency no longer jitters run-to-run. - _device_guard_plugin (new): set GEMS_VENDOR via torch probe before importing flag_gems, avoiding its timeout-less nvidia-smi subprocess probe that can hang import in fork-broken environments. - _pretty_report_plugin (new) + _term_style (new): fold inputs identical across all result rows into a legend line, color status/plugin tags/markers on the live terminal; run.log is ANSI-stripped and keeps upstream SUCCESS/column wording for grep compatibility. - run_pytest.sh: make USE_FLAGTUNE overridable, group all knobs into a config section with Chinese comments, add start/end banners. - README: document the warmup semantics, new plugins/env vars, and the A/B rule that both sides must use the same USE_FLAGTUNE.
436 lines
17 KiB
Python
436 lines
17 KiB
Python
"""pytest plugin: annotate compiled kernels and dump their ttgir, grouped by
|
|
benchmark shape, at session end.
|
|
|
|
Pieces, all keyed by the kernel's Triton cache dir:
|
|
|
|
1. Compilation listener — writes a __constexprs.json sidecar (constexpr values
|
|
+ per-arg specializations like tt.divisibility) next to each compiled
|
|
kernel. The ttgir alone doesn't say which autotune variant it is.
|
|
2. Shape tracking — wraps every Benchmark subclass's get_input_iter (after the
|
|
shape-inject plugins have done their patching) to publish the benchmark
|
|
shape currently being fed, e.g. "16x256x7168x2048x8".
|
|
3. Launch tracking — kernel_load_end_hook maps a loaded function handle to its
|
|
cache dir (must be the END hook: at load-start the handle is still None);
|
|
launch_enter_hook counts launches per (kernel cache dir, current shape).
|
|
Launches made inside Autotuner._bench (the tuning sweep; LibTuner inherits
|
|
it) are counted separately, so "actually used" = launched at least once
|
|
OUTSIDE the sweep.
|
|
4. Dump at exit (only when FLAGGEMS_PERF_TTGIR_DUMP_DIR is set) — walks
|
|
TRITON_CACHE_DIR and copies each ttgir into
|
|
<dump>/<shape>/<kernel>/<distinguishing constexprs>_w{warps}s{stages}.ttgir
|
|
for every shape that actually used it. Only constexprs that VARY within a
|
|
(shape, kernel) group are named, abbreviated per-word (BLOCK_SIZE_M -> BSM;
|
|
legend in naming.md). Same-name collisions are disambiguated by signature
|
|
dtype (dtype sweeps, e.g. __fp16), then varying per-arg specializations
|
|
(e.g. __EMdiv16), then a cache-hash prefix. index.tsv lists every compiled
|
|
variant per shape with launch counts, including never-used sweep losers
|
|
(not copied). atexit (not sessionfinish) so a mid-run CUDA crash still
|
|
dumps whatever was compiled.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import atexit
|
|
import json
|
|
import os
|
|
import re
|
|
import sys
|
|
from collections import defaultdict
|
|
|
|
import pytest
|
|
|
|
_DUMP_DIR_ENV = "FLAGGEMS_PERF_TTGIR_DUMP_DIR"
|
|
|
|
_fn_to_cachedir: dict = {} # GPU function handle -> cache-dir basename
|
|
_launches: dict = {} # cache-dir basename -> {shape label: [real, sweep]}
|
|
_in_bench = 0 # >0 while inside Autotuner._bench (autotune sweep)
|
|
_current_shape: str | None = None
|
|
_dump_enabled = False
|
|
|
|
|
|
# --- compile-time sidecar -------------------------------------------------
|
|
|
|
def _plain(value):
|
|
value = getattr(value, "value", value) # unwrap tl.constexpr
|
|
if value is None or isinstance(value, (bool, int, float, str)):
|
|
return value
|
|
return repr(value)
|
|
|
|
|
|
def _compile_listener(*, src, metadata, metadata_group, times, cache_hit):
|
|
try:
|
|
fn = getattr(src, "fn", None)
|
|
arg_names = getattr(fn, "arg_names", None)
|
|
constants = getattr(src, "constants", None)
|
|
if not arg_names or constants is None:
|
|
return # IRSource or unexpected layout: nothing to record
|
|
paths = list(metadata_group.values())
|
|
if not paths:
|
|
return
|
|
path = os.path.join(os.path.dirname(paths[0]), "__constexprs.json")
|
|
if os.path.exists(path):
|
|
return # cache hit on a dir we already annotated
|
|
|
|
def arg_name(key):
|
|
# ASTSource keys constants/attrs by arg-index tuples; map back to names.
|
|
if isinstance(key, tuple):
|
|
return ".".join(
|
|
arg_names[i] if isinstance(i, int) and i < len(arg_names) else str(i)
|
|
for i in key
|
|
)
|
|
return str(key)
|
|
|
|
out = {arg_name(k): _plain(v) for k, v in constants.items()}
|
|
|
|
# Argument signature (dtypes): benchmarks sweep dtypes, and dtype is
|
|
# not a constexpr — without this, dtype variants collide into hash
|
|
# suffixes. Values look like "*fp16", "i32".
|
|
sig = {}
|
|
try:
|
|
sig = {str(k): str(v) for k, v in (getattr(src, "signature", None) or {}).items()
|
|
if str(v) != "constexpr"}
|
|
except Exception:
|
|
pass
|
|
|
|
# Per-arg specializations, e.g. ("tt.divisibility", 16) -> "div16".
|
|
specs = {}
|
|
for key, props in (getattr(src, "attrs", None) or {}).items():
|
|
encoded = []
|
|
for p in props or []:
|
|
try:
|
|
pname, pval = p[0], p[1]
|
|
except (TypeError, IndexError):
|
|
encoded.append(str(p))
|
|
continue
|
|
encoded.append(f"div{pval}" if "divisibility" in str(pname)
|
|
else f"{pname}={pval}")
|
|
if encoded:
|
|
specs[arg_name(key)] = sorted(encoded)
|
|
|
|
tmp = f"{path}.tmp.pid{os.getpid()}"
|
|
with open(tmp, "w") as f:
|
|
json.dump({"name": getattr(src, "name", "unknown"),
|
|
"constexprs": out, "attrs": specs, "signature": sig},
|
|
f, indent=1, sort_keys=True)
|
|
os.replace(tmp, path)
|
|
except Exception as exc: # never break compilation over a metadata dump
|
|
from _term_style import tag
|
|
print(f"{tag('[ir-meta-plugin]')} sidecar dump failed: {exc}", file=sys.stderr)
|
|
|
|
|
|
# --- shape + launch tracking ----------------------------------------------
|
|
|
|
def _load_hook(module, function, name, metadata_group, hash):
|
|
try:
|
|
paths = list(metadata_group.values())
|
|
if paths:
|
|
_fn_to_cachedir[function] = os.path.basename(os.path.dirname(paths[0]))
|
|
except Exception:
|
|
pass
|
|
|
|
|
|
def _launch_hook(md):
|
|
try:
|
|
cachedir = _fn_to_cachedir.get(md.data.get("function"))
|
|
if cachedir is None:
|
|
return
|
|
rec = _launches.setdefault(cachedir, {}).setdefault(
|
|
_current_shape or "shape_unknown", [0, 0])
|
|
rec[1 if _in_bench else 0] += 1
|
|
except Exception:
|
|
pass
|
|
|
|
|
|
def _wrap_bench(original):
|
|
def wrapped(self, *args, **kwargs):
|
|
global _in_bench
|
|
_in_bench += 1
|
|
try:
|
|
return original(self, *args, **kwargs)
|
|
finally:
|
|
_in_bench -= 1
|
|
return wrapped
|
|
|
|
|
|
def _shape_label(bench, idx):
|
|
shapes = getattr(bench, "shapes", None) or []
|
|
if idx < len(shapes):
|
|
s = shapes[idx]
|
|
if isinstance(s, (list, tuple)):
|
|
return "x".join(re.sub(r"\W", "", str(v)) for v in s)
|
|
return re.sub(r"\W", "", str(s))
|
|
return f"input{idx}"
|
|
|
|
|
|
def _wrap_input_iter(original):
|
|
# Publish the shape label BEFORE resuming the generator, so kernels
|
|
# launched while building the inputs attribute to the right shape too.
|
|
def patched(self, dtype):
|
|
global _current_shape
|
|
it = original(self, dtype)
|
|
idx = 0
|
|
while True:
|
|
_current_shape = _shape_label(self, idx)
|
|
try:
|
|
item = next(it)
|
|
except StopIteration:
|
|
_current_shape = None
|
|
return
|
|
yield item
|
|
idx += 1
|
|
return patched
|
|
|
|
|
|
@pytest.hookimpl(trylast=True)
|
|
def pytest_collection_finish(session):
|
|
# trylast: run after the shape-inject plugins have re-pointed
|
|
# get_input_iter, so we wrap the version that will actually execute.
|
|
if not _dump_enabled:
|
|
return
|
|
from benchmark import base as fg_base
|
|
|
|
def _subclasses(cls):
|
|
for sub in cls.__subclasses__():
|
|
yield sub
|
|
yield from _subclasses(sub)
|
|
|
|
seen, wrapped = set(), 0
|
|
for cls in (fg_base.Benchmark, *_subclasses(fg_base.Benchmark)):
|
|
if cls in seen:
|
|
continue
|
|
seen.add(cls)
|
|
own = cls.__dict__.get("get_input_iter")
|
|
if own is not None:
|
|
cls.get_input_iter = _wrap_input_iter(own)
|
|
wrapped += 1
|
|
from _term_style import tag
|
|
print(f"{tag('[ir-meta-plugin]')} shape tracking wrapped on {wrapped} Benchmark classes",
|
|
file=sys.stderr, flush=True)
|
|
|
|
|
|
# --- naming ----------------------------------------------------------------
|
|
|
|
_SAN = re.compile(r"[^A-Za-z0-9.-]+")
|
|
|
|
|
|
def _fmt_value(v):
|
|
# isinstance check first: 1 == True in Python, a plain dict lookup would
|
|
# render GROUP_SIZE_M=1 as "T".
|
|
if isinstance(v, bool):
|
|
return "T" if v else "F"
|
|
if v is None:
|
|
return "-"
|
|
s = str(v)
|
|
if "." in s and not re.fullmatch(r"-?\d+(\.\d+)?", s):
|
|
s = s.split(".")[-1] # dotted repr like triton.language.bfloat16
|
|
s = _SAN.sub("", s)
|
|
return s[:24] or "x"
|
|
|
|
|
|
def _abbrev(name):
|
|
words = [w for w in name.split("_") if w]
|
|
if len(name) <= 4 or len(words) < 2:
|
|
return name
|
|
out = []
|
|
for w in words:
|
|
m = re.match(r"^([A-Za-z])[A-Za-z]*?(\d*)$", w)
|
|
out.append((m.group(1) + m.group(2)) if m else w[0])
|
|
return "".join(out)
|
|
|
|
|
|
def _disambiguate(group):
|
|
"""Suffixes for variants whose constexprs + launch config coincide: name by
|
|
the signature dtypes that differ (dtype sweeps), else by differing per-arg
|
|
specializations, else fall back to a hash prefix."""
|
|
sig_keys = sorted({k for v in group for k in (v.get("signature") or {})})
|
|
sig_varying = [
|
|
k for k in sig_keys
|
|
if len({(v.get("signature") or {}).get(k) for v in group}) > 1
|
|
]
|
|
if sig_varying:
|
|
# One dtype sweep usually shifts every tensor arg together; the first
|
|
# varying arg's dtype identifies the variant. Sanitize "*fp16" -> fp16.
|
|
sufs = []
|
|
for v in group:
|
|
toks = [re.sub(r"\W", "", (v.get("signature") or {}).get(k) or "none")
|
|
for k in sig_varying]
|
|
uniq = sorted(set(toks))
|
|
sufs.append("__" + (uniq[0] if len(uniq) == 1 else "_".join(
|
|
f"{_abbrev(k)}{t}" for k, t in zip(sig_varying, toks))))
|
|
if len(set(sufs)) == len(group):
|
|
return sufs
|
|
attr_keys = sorted({k for v in group for k in (v.get("attrs") or {})})
|
|
varying = [
|
|
k for k in attr_keys
|
|
if len({tuple((v.get("attrs") or {}).get(k, [])) for v in group}) > 1
|
|
]
|
|
if varying:
|
|
sufs = []
|
|
for v in group:
|
|
toks = [
|
|
f"{_abbrev(k)}{'.'.join((v.get('attrs') or {}).get(k) or ['none'])}"
|
|
for k in varying
|
|
]
|
|
sufs.append("__" + "_".join(toks))
|
|
if len(set(sufs)) == len(group):
|
|
return sufs
|
|
return ["__" + v["hash"][:8] for v in group]
|
|
|
|
|
|
def _name_group(items):
|
|
"""Filenames for one (shape, kernel) group: only constexprs whose value
|
|
varies within the group, abbreviated. Returns ([(item, filename)], legend)."""
|
|
keys = sorted({k for it in items for k in it["constexprs"]})
|
|
varying = [
|
|
k for k in keys
|
|
if len({json.dumps(it["constexprs"].get(k), sort_keys=True) for it in items}) > 1
|
|
]
|
|
by_ab = defaultdict(list)
|
|
for k in varying:
|
|
by_ab[_abbrev(k)].append(k)
|
|
ab = {k: (a if len(ks) == 1 else k) for a, ks in by_ab.items() for k in ks}
|
|
|
|
named = defaultdict(list)
|
|
for it in items:
|
|
parts = [f"{ab[k]}{_fmt_value(it['constexprs'][k])}"
|
|
for k in varying if k in it["constexprs"]]
|
|
w, s = it.get("warps"), it.get("stages")
|
|
parts.append(f"w{w}s{s}" if w not in (None, "") else "cfg-unknown")
|
|
named["_".join(parts)].append(it)
|
|
|
|
results = []
|
|
for base, group in sorted(named.items()):
|
|
sufs = _disambiguate(group) if len(group) > 1 else [""]
|
|
for it, suf in zip(group, sufs):
|
|
results.append((it, f"{base}{suf}.ttgir"))
|
|
return results, {k: a for k, a in ab.items() if a != k}
|
|
|
|
|
|
# --- dump -------------------------------------------------------------------
|
|
|
|
def _dump_ttgir(cache_dir: str, dump_dir: str) -> None:
|
|
import pathlib
|
|
import shutil
|
|
|
|
cache, dest = pathlib.Path(cache_dir), pathlib.Path(dump_dir)
|
|
if not cache.is_dir():
|
|
from _term_style import tag
|
|
print(f"{tag('[ir-meta-plugin]')} no cache dir {cache}; nothing to dump",
|
|
file=sys.stderr)
|
|
return
|
|
dest.mkdir(parents=True, exist_ok=True)
|
|
|
|
def load(path):
|
|
try:
|
|
return json.loads(path.read_text())
|
|
except Exception:
|
|
return {}
|
|
|
|
infos = []
|
|
for ttgir in sorted(cache.rglob("*.ttgir")):
|
|
meta = load(ttgir.with_suffix(".json"))
|
|
side = load(ttgir.parent / "__constexprs.json")
|
|
infos.append({
|
|
"src": ttgir, "hash": ttgir.parent.name, "kernel": ttgir.stem,
|
|
"warps": meta.get("num_warps"), "stages": meta.get("num_stages"),
|
|
"constexprs": side.get("constexprs", {}), "attrs": side.get("attrs", {}),
|
|
"signature": side.get("signature", {}),
|
|
"per_shape": _launches.get(ttgir.parent.name, {}),
|
|
})
|
|
|
|
# (shape, kernel) -> variants really used there (launched outside the sweep)
|
|
groups = defaultdict(list)
|
|
for info in infos:
|
|
for shape, (real, _sweep) in info["per_shape"].items():
|
|
if real > 0:
|
|
groups[(shape, info["kernel"])].append(info)
|
|
|
|
rows, legends, copied = [], defaultdict(dict), 0
|
|
for (shape, kernel), items in sorted(groups.items()):
|
|
named, legend = _name_group(items)
|
|
legends[kernel].update(legend)
|
|
for it, fname in named:
|
|
target = dest / shape / kernel / fname
|
|
target.parent.mkdir(parents=True, exist_ok=True)
|
|
shutil.copyfile(it["src"], target)
|
|
copied += 1
|
|
real, sweep = it["per_shape"][shape]
|
|
rows.append((shape, kernel, fname, real, sweep,
|
|
it["warps"] or "", it["stages"] or "",
|
|
json.dumps(it["constexprs"], sort_keys=True), it["hash"]))
|
|
|
|
# Compiled but never really used anywhere (autotune losers): index-only.
|
|
for info in infos:
|
|
if not any(real > 0 for real, _ in info["per_shape"].values()):
|
|
sweep = sum(s for _, s in info["per_shape"].values())
|
|
rows.append(("-", info["kernel"], "(sweep loser, not dumped)", 0, sweep,
|
|
info["warps"] or "", info["stages"] or "",
|
|
json.dumps(info["constexprs"], sort_keys=True), info["hash"]))
|
|
|
|
rows.sort()
|
|
with open(dest / "index.tsv", "w") as f:
|
|
f.write("shape\tkernel\tfile\tlaunches\tsweep_launches\t"
|
|
"num_warps\tnum_stages\tconstexprs\tcache_hash\n")
|
|
for r in rows:
|
|
f.write("\t".join(str(x) for x in r) + "\n")
|
|
|
|
lines = ["# ttgir naming legend", "",
|
|
"Layout: `<shape>/<kernel>/<varying constexprs>_w{warps}s{stages}[__spec|__hash8].ttgir`",
|
|
"Shape dirs mirror the benchmark shape yaml. Only constexprs that vary",
|
|
"within a (shape, kernel) group appear; `index.tsv` has the full map.",
|
|
"Values: `T`/`F` = true/false, `-` = none/null."]
|
|
for kernel in sorted(legends):
|
|
if not legends[kernel]:
|
|
continue
|
|
lines += ["", f"## {kernel}"]
|
|
width = max(len(a) for a in legends[kernel].values())
|
|
for full, a in sorted(legends[kernel].items(), key=lambda kv: kv[1]):
|
|
lines.append(f"- `{a:<{width}}` = {full}")
|
|
(dest / "naming.md").write_text("\n".join(lines) + "\n")
|
|
|
|
shapes = sorted({r[0] for r in rows if r[0] != "-"})
|
|
losers = sum(1 for r in rows if r[0] == "-")
|
|
from _term_style import DIM, paint
|
|
print(f">>> [Dump] ttgir -> {dest} ({copied} files across {len(shapes)} shapes; "
|
|
f"{losers} unused variants index-only; legend: naming.md)", flush=True)
|
|
for s in shapes:
|
|
n = sum(1 for r in rows if r[0] == s)
|
|
print(paint(f">>> [Dump] {s}: {n}", DIM), flush=True)
|
|
|
|
|
|
# --- registration -----------------------------------------------------------
|
|
|
|
def pytest_configure(config):
|
|
global _dump_enabled
|
|
import triton
|
|
import triton.runtime.autotuner as _autotuner
|
|
|
|
prev = triton.knobs.compilation.listener
|
|
if prev is None:
|
|
triton.knobs.compilation.listener = _compile_listener
|
|
else:
|
|
def chained(**kwargs):
|
|
prev(**kwargs)
|
|
_compile_listener(**kwargs)
|
|
triton.knobs.compilation.listener = chained
|
|
|
|
dump_dir = os.environ.get(_DUMP_DIR_ENV, "").strip()
|
|
if dump_dir:
|
|
cache_dir = os.environ.get("TRITON_CACHE_DIR", "").strip()
|
|
if cache_dir:
|
|
_dump_enabled = True
|
|
triton.knobs.runtime.kernel_load_end_hook.add(_load_hook)
|
|
triton.knobs.runtime.launch_enter_hook.add(_launch_hook)
|
|
_autotuner.Autotuner._bench = _wrap_bench(_autotuner.Autotuner._bench)
|
|
atexit.register(_dump_ttgir, cache_dir, dump_dir)
|
|
else:
|
|
from _term_style import tag
|
|
print(f"{tag('[ir-meta-plugin]')} warning: dump dir set but TRITON_CACHE_DIR "
|
|
"is not; ttgir dump disabled", file=sys.stderr, flush=True)
|
|
|
|
from _term_style import tag
|
|
print(f"{tag('[ir-meta-plugin]')} compilation listener registered"
|
|
+ (f"; shape/launch tracking on, ttgir dump -> {dump_dir}"
|
|
if _dump_enabled else ""),
|
|
file=sys.stderr, flush=True)
|