"""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 ///_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 print(f"[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 print(f"[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(): print(f"[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: `//_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] == "-") 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(f">>> [Dump] {s}: {n}", 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: print("[ir-meta-plugin] warning: dump dir set but TRITON_CACHE_DIR " "is not; ttgir dump disabled", file=sys.stderr, flush=True) print("[ir-meta-plugin] compilation listener registered" + (f"; shape/launch tracking on, ttgir dump -> {dump_dir}" if _dump_enabled else ""), file=sys.stderr, flush=True)