From 02cabf83e7ec582814dc4f9b5415f4b49aacd5d0 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 12 Sep 2025 08:02:38 -0400 Subject: [PATCH 1/3] [FIX] Add __call__ method to Trace class for nested JIT function calls MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit When a traced JIT function is called from within another JIT function, the Trace wrapper needs to properly delegate to the underlying function. This fix ensures compatibility with nested Triton JIT function calls. 🤖 Generated with [Claude Code](https://claude.ai/code) Co-Authored-By: Claude --- triton_viz/core/trace.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/triton_viz/core/trace.py b/triton_viz/core/trace.py index 88078eecb..b4115f2ff 100644 --- a/triton_viz/core/trace.py +++ b/triton_viz/core/trace.py @@ -69,6 +69,11 @@ def run(self, *args, **kwargs): self.finalize() return ret + def __call__(self, *args, **kwargs): + # When a traced JIT function is called from within another JIT function, + # we need to execute the underlying function directly + return self.base_fn(*args, **kwargs) + def warmup(self, *args, **kwargs): raise NotImplementedError From d10d6d13a880d50f79c6917f0b51102d517b900d Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 19 Sep 2025 08:54:04 -0400 Subject: [PATCH 2/3] add unittest --- tests/test_core.py | 27 +++++++++++++++++++++++++++ 1 file changed, 27 insertions(+) diff --git a/tests/test_core.py b/tests/test_core.py index 5ae38bb18..375bb1cf3 100644 --- a/tests/test_core.py +++ b/tests/test_core.py @@ -59,6 +59,33 @@ def my_kernel(x_ptr, y_ptr, out_ptr, BLOCK_SIZE: tl.constexpr): assert sum(c == "tracer" for c in clients) == 1 +# ======== Nested JIT Call Tests ========= +@triton_viz.trace(clients=Sanitizer(abort_on_error=True)) +@triton.jit +def trace_nested_inner_kernel(x): + return x * 2 + +def test_trace_nested_jit_calls(): + """ + Test that Trace class properly handles nested JIT function calls via __call__ method. + + When a traced JIT function is called from within another JIT function, + the Trace wrapper needs to properly delegate to the underlying function. + This test ensures compatibility with nested Triton JIT function calls. + """ + + @triton_viz.trace(clients=Sanitizer(abort_on_error=True)) + @triton.jit + def trace_nested_call_kernel(ptr, n: tl.constexpr): + x = tl.load(ptr + tl.arange(0, n)) + y = trace_nested_inner_kernel(x) # This nested call requires __call__ method when wrapped by trace + tl.store(ptr + tl.arange(0, n), y) + + # Test execution + data = torch.ones(8) + trace_nested_call_kernel[(1,)](data, 8) + + # ======== Wrapper Tests ========= # Test that triton_viz.wrapper works correctly: # 1. It should patch triton.jit / triton.language.jit / From 396c5e85bf803ffe3b9891c463067aca77baf3fb Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 19 Sep 2025 09:03:27 -0400 Subject: [PATCH 3/3] pre-commit formatting --- tests/test_core.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/tests/test_core.py b/tests/test_core.py index 4ecff3697..7b19013f3 100644 --- a/tests/test_core.py +++ b/tests/test_core.py @@ -58,6 +58,7 @@ def my_kernel(x_ptr, y_ptr, out_ptr, BLOCK_SIZE: tl.constexpr): def trace_nested_inner_kernel(x): return x * 2 + def test_trace_nested_jit_calls(): """ Test that Trace class properly handles nested JIT function calls via __call__ method. @@ -71,7 +72,9 @@ def test_trace_nested_jit_calls(): @triton.jit def trace_nested_call_kernel(ptr, n: tl.constexpr): x = tl.load(ptr + tl.arange(0, n)) - y = trace_nested_inner_kernel(x) # This nested call requires __call__ method when wrapped by trace + y = trace_nested_inner_kernel( + x + ) # This nested call requires __call__ method when wrapped by trace tl.store(ptr + tl.arange(0, n), y) # Test execution