diff --git a/tests/test_wrapper.py b/tests/test_wrapper.py new file mode 100644 index 00000000..18849180 --- /dev/null +++ b/tests/test_wrapper.py @@ -0,0 +1,146 @@ +# tests/test_triton_sanitizer_trace_injection.py +import json +import os +import sys +import textwrap +import subprocess +from pathlib import Path +import shutil + + +def test_triton_sanitizer_injects_trace_outermost(tmp_path: Path, monkeypatch): + """ + Black-box verification: + - Pure @triton.jit kernels have exactly one @triton_viz.trace at the outermost layer; + - @triton.autotune + @triton.jit kernels: trace is only added at the autotune outer layer, inner jit should not be traced again. + Implementation approach: + Use sitecustomize to monkey-patch triton_viz.trace at subprocess startup, + record each trace decorator application target (function name/type), write to JSON on exit. + """ + + # 1) Write sitecustomize: takes effect at interpreter startup + trace_log = tmp_path / "trace_log.json" + sitecustomize = tmp_path / "sitecustomize.py" + sitecustomize.write_text( + textwrap.dedent( + f""" + import atexit, json, os + # Import and replace triton_viz.trace here + import triton_viz as _tv + + _orig_trace = _tv.trace + _records = [] + + def _spy_trace(*t_args, **t_kwargs): + # Call original trace, get the actual decorator + _decorator = _orig_trace(*t_args, **t_kwargs) + def _apply(target): + # Which object is being applied to (usually function name, or autotune/jit product) + name = getattr(target, "__name__", None) or getattr(target, "__qualname__", None) or repr(target) + typ = type(target).__name__ + mod = getattr(target, "__module__", None) + + _records.append({{ + "name": name, + "type": typ, + "module": mod, + }}) + + # Continue with original logic, keep functionality unchanged + return _decorator(target) + return _apply + + # Replace + _tv.trace = _spy_trace + + # Write records to disk on process exit + _LOG_PATH = {json.dumps(str(trace_log))} + @atexit.register + def _dump_records(): + try: + with open(_LOG_PATH, "w") as f: + json.dump(_records, f) + except Exception as e: + # Avoid affecting test process exit + pass + """ + ), + encoding="utf-8", + ) + + # 2) Write a minimal test script: two kernels + # - k1: pure jit + # - k2: autotune + jit + my_program = tmp_path / "my_program.py" + my_program.write_text( + textwrap.dedent( + """ + import triton + import triton.language as tl + + # --- Pure jit --- + @triton.jit + def k1(X, n: tl.constexpr): + return + + # --- autotune + jit --- + @triton.autotune(configs=[triton.Config({})], key=['n']) + @triton.jit + def k2(X, n: tl.constexpr): + return + + if __name__ == "__main__": + # No need to actually launch kernel; we just need the decorator chain to be built + pass + """ + ), + encoding="utf-8", + ) + + # 3) Assemble subprocess environment: ensure our sitecustomize is found first + env = os.environ.copy() + # Put tmp_path at the front, ensure sitecustomize.py is auto-imported + env["PYTHONPATH"] = str(tmp_path) + os.pathsep + env.get("PYTHONPATH", "") + + # 4) Find triton-sanitizer executable (also allow python -m entry as fallback) + exe = shutil.which("triton-sanitizer") + if exe: + cmd = [exe, str(my_program)] + else: + # If your entry point name is different, use the corresponding module (below is the wrapper.apply path) + cmd = [sys.executable, "-m", "triton_viz.wrapper", str(my_program)] + + # 5) Run subprocess + proc = subprocess.run( + cmd, + cwd=str(tmp_path), + env=env, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + ) + + # Optional debug output (easier to locate on failure) + if proc.returncode != 0: + print("STDOUT:\n", proc.stdout) + print("STDERR:\n", proc.stderr) + assert proc.returncode == 0, "triton-sanitizer run failed" + + # 6) Read trace records + assert ( + trace_log.exists() + ), "trace log not generated (sitecustomize may not have taken effect)" + records = json.loads(trace_log.read_text(encoding="utf-8")) + + # Build counter by function name + from collections import Counter + + counts = Counter(rec["name"] for rec in records) + + # Assertion: the 2 kernel names defined in our script should each appear exactly once + # - k1: pure jit → should have trace injected once at outer layer + # - k2: autotune + jit → should only appear once at autotune outer layer (should not trace inner jit again) + for expected in ["k1", "k2"]: + assert ( + counts[expected] == 1 + ), f"{expected} should be traced 1 time, actual count is {counts[expected]}; may have missing or duplicate injection" diff --git a/triton_viz/wrapper.py b/triton_viz/wrapper.py index 8d549f8c..3e1472bb 100644 --- a/triton_viz/wrapper.py +++ b/triton_viz/wrapper.py @@ -1,5 +1,8 @@ +import os +import shutil import runpy import sys +import pytest import triton import triton_viz from triton_viz.clients import Sanitizer @@ -7,6 +10,7 @@ # store the original triton.jit _original_jit = triton.jit +_original_autotune = triton.autotune def sanitizer_wrapper(kernel): @@ -28,6 +32,19 @@ def _decorator(f): return sanitizer_wrapper(k) +def _patched_autotune(fn=None, **autotune_kw): + if fn is None: + + def _decorator(f): + k = _original_autotune(**autotune_kw)(f) + return sanitizer_wrapper(k) + + return _decorator + else: + k = _original_autotune(fn) + return sanitizer_wrapper(k) + + def apply(): """ Apply the sanitizer wrapper to triton.jit and run the user script. @@ -39,6 +56,9 @@ def apply(): _interp.jit = _patched_jit + # patching triton.autotune + triton.autotune = _patched_autotune + # run user script # argv is like: ['triton-sanitizer', 'user_script.py', 'arg1', 'arg2', ...] if len(sys.argv) < 2: @@ -47,11 +67,6 @@ def apply(): script = sys.argv[1] - # Check if the first argument is a file that exists - import os - import subprocess - import shutil - if os.path.isfile(script): # It's a Python script file, run it directly sys.argv = sys.argv[1:] @@ -62,12 +77,17 @@ def apply(): # Check if it's an executable command if shutil.which(cmd[0]): - # Run the command with the sanitizer environment already set up - result = subprocess.run(cmd, env=os.environ.copy()) - sys.exit(result.returncode) - else: - print(f"Error: '{script}' is neither a valid file nor a command") - sys.exit(1) + if cmd[0] == "pytest": + sys.exit(pytest.main(cmd[1:])) + elif cmd[0] == "python": + sys.argv = cmd[1:] + try: + runpy.run_path(cmd[1], run_name="__main__") + except SystemExit as e: + sys.exit(e.code) + sys.exit(0) + print(f"Error: '{script}' is neither a valid file nor a command") + sys.exit(1) def enable_sanitizer():