Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
146 changes: 146 additions & 0 deletions tests/test_wrapper.py
Original file line number Diff line number Diff line change
@@ -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"
42 changes: 31 additions & 11 deletions triton_viz/wrapper.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,16 @@
import os
import shutil
import runpy
import sys
import pytest
import triton
import triton_viz
from triton_viz.clients import Sanitizer


# store the original triton.jit
_original_jit = triton.jit
_original_autotune = triton.autotune


def sanitizer_wrapper(kernel):
Expand All @@ -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.
Expand All @@ -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:
Expand All @@ -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:]
Expand All @@ -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():
Expand Down