diff --git a/README.md b/README.md index 0efc0527..a0413df4 100644 --- a/README.md +++ b/README.md @@ -187,7 +187,7 @@ simulation runtimes live under `tilelens/core/simulation/`. Analyze kernels across visualization, profiling, and sanitization with a single line of code. -- Visualizer: currently supports load, store, and matmul operations for 1/2/3D tensors (more operations and dimensions coming soon). +- Visualizer: currently supports load, store, and matmul operations for 1/2/3D tensors (more operations and dimensions coming soon). Triton's interpreter runs loads, stores, and atomics on raw host memory, so the tracer raises `IndexError` (wrapped in an `InterpreterError` under Gluon) before a Triton load, store, or atomic, or a Gluon `gl.load`/`gl.store`, that runs off the storage of a tensor argument, instead of letting it corrupt memory. This guard is best-effort: accesses that never touch a tensor argument (for example a whole program past the end), programs that cast an integer to a pointer (`x.to(tl.pointer_type(...))`, as pointer tables do), and Gluon atomics and async copies are not checked, so use the Sanitizer for a complete out-of-bounds report. - Profiler: flags non-unrolled loops, inefficient mask usage, and missing buffer_load optimizations while tracking load/store byte counts with low-overhead sampling. - Sanitizer: symbolically checks tensor memory accesses for out-of-bounds errors and emits reports with tensor metadata, call stack, and expression trees; optional fake-memory storage avoids real reads. diff --git a/tests/end_to_end/test_gluon.py b/tests/end_to_end/test_gluon.py index d8882720..f7ba26fe 100644 --- a/tests/end_to_end/test_gluon.py +++ b/tests/end_to_end/test_gluon.py @@ -1,6 +1,7 @@ import importlib.util from pathlib import Path +import numpy as np import pytest import torch from triton import knobs @@ -105,6 +106,19 @@ def _range_memcpy_kernel(in_ptr, out_ptr, xnumel, BLOCK: gl.constexpr): gl.store(out_ptr + i, value) +@gluon.jit +def _unmasked_store_1d_memcpy_kernel( + in_ptr, + out_ptr, + xnumel, + BLOCK: gl.constexpr, + layout: gl.constexpr, +): + offsets = gl.program_id(0) * BLOCK + gl.arange(0, BLOCK, layout=layout) + value = gl.load(in_ptr + offsets, mask=offsets < xnumel, other=0.0) + gl.store(out_ptr + offsets, value) + + @gluon.jit def _masked_1d_memcpy_kernel( in_ptr, @@ -868,6 +882,26 @@ def test_gluon_core_ops_run_masked_1d_memcpy_on_cpu(): torch.testing.assert_close(out, inp, atol=0, rtol=0) +def test_gluon_trace_refuses_out_of_bounds_store(): + # out's storage is exactly 40 floats; the rest of buf is sentinel memory + buf = np.full(40 + 128, -7, dtype=np.float32) + out = torch.from_numpy(buf[64:104]) + inp = torch.arange(40, dtype=torch.float32) + layout = gl.BlockedLayout([1], [32], [1], [0]) + kernel = tilelens.trace("tracer", frontend="gluon")( + _unmasked_store_1d_memcpy_kernel + ) + + with pytest.raises(Exception) as exc_info: + kernel[(1,)](inp, out, inp.numel(), 64, layout, num_warps=1) + + # Gluon's interpreter wraps the tracer's IndexError in an InterpreterError + cause = exc_info.value.__cause__ + assert isinstance(cause, IndexError) + assert "out-of-bounds store" in str(cause) + assert (buf[:64] == -7).all() and (buf[104:] == -7).all() + + def test_gluon_core_ops_run_masked_2d_memcpy_on_cpu(): inp = torch.arange(24, dtype=torch.float32).reshape(4, 6) out = torch.full_like(inp, -1) diff --git a/tests/end_to_end/test_tracer.py b/tests/end_to_end/test_tracer.py index de75777c..60c6e0ff 100644 --- a/tests/end_to_end/test_tracer.py +++ b/tests/end_to_end/test_tracer.py @@ -1,3 +1,5 @@ +import numpy as np +import pytest import torch import triton import triton.language as tl @@ -223,3 +225,178 @@ def autotune_add_kernel_cache_on( assert ( bench_fn is not None and bench_fn.__name__ == "dummy_benchmarker" ), f"Expected dummy_benchmarker, got: {bench_fn}" + + +def _fenced(values, dtype=np.float32, pad=64): + """ + Return a tensor whose storage holds exactly `values`, plus the sentinel + memory on both sides of it. Out-of-bounds writes would land in the + sentinels instead of the heap, so tests can check them deterministically. + """ + buf = np.full(len(values) + 2 * pad, -7, dtype=dtype) + buf[pad:-pad] = values + return torch.from_numpy(buf[pad:-pad]), (buf[:pad], buf[-pad:]) + + +def _untouched(fences): + return all((fence == -7).all() for fence in fences) + + +@triton.jit +def unmasked_store_kernel(x_ptr, out_ptr, n_elements, BLOCK_SIZE: tl.constexpr): + offs = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + x = tl.load(x_ptr + offs, mask=offs < n_elements, other=0.0) + tl.store(out_ptr + offs, x) + + +@pytest.mark.parametrize("grid_idx", [None, 0]) +def test_tracer_refuses_out_of_bounds_load(grid_idx): + traced = tilelens.trace(client=Tracer(grid_idx=grid_idx))(copy_kernel) + x, x_fences = _fenced(np.arange(10)) + out, out_fences = _fenced(np.zeros(10)) + + with pytest.raises(IndexError, match=r"(?s)out-of-bounds load .*`x_ptr`"): + traced[(3,)](x, out, BLOCK_SIZE=4) + assert _untouched(x_fences) and _untouched(out_fences) + + +@pytest.mark.parametrize("num_sms", [1, 4]) +def test_tracer_refuses_out_of_bounds_store(monkeypatch, num_sms): + monkeypatch.setattr(tilelens.config, "num_sms", num_sms) + traced = tilelens.trace(client=Tracer())(unmasked_store_kernel) + x = torch.arange(10, dtype=torch.float32) + out, fences = _fenced(np.zeros(10)) + + with pytest.raises(IndexError, match=r"(?s)out-of-bounds store .*`out_ptr`"): + traced[(3,)](x, out, x.numel(), BLOCK_SIZE=4) + assert _untouched(fences) + + +@triton.jit +def unmasked_atomic_add_kernel(out_ptr, BLOCK_SIZE: tl.constexpr): + offs = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + tl.atomic_add(out_ptr + offs, 1) + + +@triton.jit +def unmasked_atomic_cas_kernel(out_ptr, BLOCK_SIZE: tl.constexpr): + offs = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + tl.atomic_cas(out_ptr + offs, tl.zeros_like(offs), tl.full(offs.shape, 1, tl.int32)) + + +@pytest.mark.parametrize( + "kernel, op_name", + [ + (unmasked_atomic_add_kernel, "atomic"), + (unmasked_atomic_cas_kernel, "atomic_cas"), + ], +) +def test_tracer_refuses_out_of_bounds_atomic(kernel, op_name): + traced = tilelens.trace(client=Tracer())(kernel) + out, fences = _fenced(np.zeros(10), dtype=np.int32) + + with pytest.raises(IndexError, match=f"out-of-bounds {op_name} "): + traced[(3,)](out, BLOCK_SIZE=4) + assert _untouched(fences) + + +@triton.jit +def unmasked_atomic_max_kernel(x_ptr, out_ptr, n_elements, BLOCK_SIZE: tl.constexpr): + offs = tl.arange(0, BLOCK_SIZE) + x = tl.load(x_ptr + offs, mask=offs < n_elements, other=-1.0) + tl.atomic_max(out_ptr + offs, x) + + +def test_tracer_refuses_sign_split_float_atomic_max(): + # Triton issues float atomic_max as two calls masked by sign; here the call + # for the negative lanes holds only the out-of-bounds ones + traced = tilelens.trace(client=Tracer())(unmasked_atomic_max_kernel) + x = torch.arange(1, 9, dtype=torch.float32) + out, fences = _fenced(np.zeros(8)) + + with pytest.raises(IndexError, match="out-of-bounds atomic "): + traced[(1,)](x, out, x.numel(), BLOCK_SIZE=16) + assert _untouched(fences) + + +@triton.jit +def packed_word_store_kernel(x_ptr, BLOCK_SIZE: tl.constexpr): + words = x_ptr.to(tl.pointer_type(tl.int32)) + offs = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + tl.store(words + offs, 1) + + +def test_tracer_refuses_word_overlapping_storage_end(): + # x holds 10 bytes, so the int32 word at byte 8 overlaps its end + traced = tilelens.trace(client=Tracer())(packed_word_store_kernel) + x, fences = _fenced(np.zeros(10), dtype=np.int8) + + with pytest.raises(IndexError, match=r"(?s)out-of-bounds store .*`x_ptr`"): + traced[(2,)](x, BLOCK_SIZE=2) + assert _untouched(fences) + + +def test_tracer_refuses_out_of_bounds_store_through_reinterpret(): + traced = tilelens.trace(client=Tracer())(unmasked_store_kernel) + x = torch.arange(10, dtype=torch.float16) + out, fences = _fenced(np.zeros(10), dtype=np.int16) + + with pytest.raises(IndexError, match=r"(?s)out-of-bounds store .*`out_ptr`"): + traced[(3,)](x, triton.reinterpret(out, tl.float16), x.numel(), BLOCK_SIZE=4) + assert _untouched(fences) + + +@triton.jit +def tuple_store_kernel(x_ptr, ptrs, BLOCK_SIZE: tl.constexpr): + offs = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + tl.store(ptrs[1] + offs, tl.load(x_ptr + offs)) + + +def test_tracer_refuses_out_of_bounds_store_through_tuple_arg(): + traced = tilelens.trace(client=Tracer())(tuple_store_kernel) + x = torch.zeros(16) + out, fences = _fenced(np.zeros(10)) + + with pytest.raises(IndexError, match=r"(?s)out-of-bounds store .*`ptrs\[1\]`"): + traced[(3,)](x, (x, out), BLOCK_SIZE=4) + assert _untouched(fences) + + +@triton.jit +def shifted_copy_kernel(x_ptr, out_ptr, SHIFT: tl.constexpr, BLOCK_SIZE: tl.constexpr): + offs = tl.arange(0, BLOCK_SIZE) + tl.store(out_ptr + offs, tl.load(x_ptr + offs - SHIFT)) + + +def test_tracer_allows_view_access_within_base_storage(): + base = torch.arange(16, dtype=torch.float32) + out = torch.empty(8) + traced = tilelens.trace(client=Tracer())(shifted_copy_kernel) + + # reads base[0:8] through a view that starts at base[4] + traced[(1,)](base[4:8], out, SHIFT=4, BLOCK_SIZE=8) + + torch.testing.assert_close(out, base[:8]) + + +@triton.jit +def pointer_table_kernel(table_ptr, arg_ptr, out_ptr, BLOCK_SIZE: tl.constexpr): + rows = tl.arange(0, 2) + cols = tl.arange(0, BLOCK_SIZE) + row_ptrs = tl.load(table_ptr + rows).to(tl.pointer_type(tl.float32)) + values = tl.load(row_ptrs[:, None] + cols[None, :]) + tl.store(out_ptr + rows[:, None] * BLOCK_SIZE + cols[None, :], values) + + +def test_tracer_allows_pointer_table_gather(): + # one row is a kernel arg and one is not, so a single gather touches known + # and unknown storage; pointers built from integers are not judged + arg_row = torch.arange(4, dtype=torch.float32) + other_row = torch.arange(4, 8, dtype=torch.float32) + table = torch.tensor([arg_row.data_ptr(), other_row.data_ptr()]) + out = torch.empty(8) + traced = tilelens.trace(client=Tracer())(pointer_table_kernel) + + traced[(1,)](table, arg_row, out, BLOCK_SIZE=4) + + torch.testing.assert_close(out, torch.arange(8, dtype=torch.float32)) diff --git a/tests/unit/test_adapters.py b/tests/unit/test_adapters.py index e0d0585a..176e64c9 100644 --- a/tests/unit/test_adapters.py +++ b/tests/unit/test_adapters.py @@ -7,6 +7,8 @@ from tilelens.core.client import Client, ClientManager from tilelens.core.data import ( AddPtr, + AtomicCas, + AtomicRMW, BinaryOp, CastImpl, Dot, @@ -172,6 +174,24 @@ def test_triton_addptr_adapter_orders_arguments(): assert result.kwargs == {} +def test_triton_atomic_rmw_adapter_returns_ptr_and_mask(): + """Adapter maps create_atomic_rmw(rmw_op, ptr, val, mask, sem, scope) to (ptr, mask).""" + adapter = TRITON_ADAPTERS[AtomicRMW] + rmw_op, ptr, val, mask, sem, scope = (object() for _ in range(6)) + result = adapter(rmw_op, ptr, val, mask, sem, scope) + assert result.args == (ptr, mask) + assert result.kwargs == {} + + +def test_triton_atomic_cas_adapter_returns_ptr_only(): + """Adapter maps create_atomic_cas(ptr, cmp, val, sem, scope) to (ptr,).""" + adapter = TRITON_ADAPTERS[AtomicCas] + ptr, cmp, val, sem, scope = (object() for _ in range(5)) + result = adapter(ptr, cmp, val, sem, scope) + assert result.args == (ptr,) + assert result.kwargs == {} + + def test_program_id_adapter_returns_axis_only(): """Adapter maps tl.program_id(axis) to (axis,).""" adapter = TRITON_ADAPTERS[ProgramId] diff --git a/tests/unit/test_tracer.py b/tests/unit/test_tracer.py index b0afbefa..4e842b84 100644 --- a/tests/unit/test_tracer.py +++ b/tests/unit/test_tracer.py @@ -2,6 +2,10 @@ from unittest.mock import MagicMock import numpy as np +import pytest +import torch +import triton.language as tl +from triton.runtime.interpreter import TensorHandle from tilelens.clients.tracer.tracer import Tracer, _convert_grid_idx from tilelens.core.data import Transfer @@ -158,3 +162,21 @@ def test_tracer_transfer_records_element_strides_as_bytes(): assert np.array_equal(record.src_offsets, linear * 4) assert np.array_equal(record.dst_offsets, linear * 2) assert record.bytes == dst_data.nbytes + + +# ======== Out-of-bounds Guard Tests =========== + + +def test_check_in_bounds_accepts_python_bool_mask(): + tracer = Tracer() + tensor = torch.zeros(4) + tracer.arg_callback("x_ptr", tensor, None) + base = tensor.data_ptr() + + def ptr(*offsets): + lanes = np.array([base + offset for offset in offsets], dtype=np.uint64) + return TensorHandle(lanes, tl.pointer_type(tl.float32)) + + tracer._check_in_bounds("load", ptr(0, 4, 8, 12), True) + with pytest.raises(IndexError, match="out-of-bounds load"): + tracer._check_in_bounds("load", ptr(8, 12, 16), True) diff --git a/tilelens/clients/tracer/tracer.py b/tilelens/clients/tracer/tracer.py index 8151ed95..c787bbba 100644 --- a/tilelens/clients/tracer/tracer.py +++ b/tilelens/clients/tracer/tracer.py @@ -11,6 +11,9 @@ Dot, Grid, Allocate, + AtomicCas, + AtomicRMW, + IntToPtr, ) from ...utils.traceback_utils import extract_user_frames from tilelens.core.masked_load_store import masked_load @@ -29,6 +32,21 @@ def _convert_grid_idx(grid_idx) -> tuple[int, int, int] | None: return grid_idx +def _storage_ranges(name, arg): + """Yield (name, tensor, start, end) for each torch storage behind a kernel arg.""" + if isinstance(arg, (tuple, list)): + for i, item in enumerate(arg): + yield from _storage_ranges(f"{name}[{i}]", item) + return + # triton.reinterpret wrappers and host TensorDescriptors keep the tensor in `base` + tensor = arg if hasattr(arg, "untyped_storage") else getattr(arg, "base", None) + if hasattr(tensor, "untyped_storage"): + storage = tensor.untyped_storage() + if storage.nbytes(): + start = storage.data_ptr() + yield name, tensor, start, start + storage.nbytes() + + class Tracer(Client): NAME = "tracer" @@ -42,6 +60,8 @@ def __init__( self.grid_idx = _convert_grid_idx(grid_idx) self.records: list = [] self.tensors: list = [] + # (arg name, tensor, start, end) of the storages behind the kernel args + self.storages: list = [] self.sample = True def _get_tensor(self, data_ptr): @@ -54,6 +74,95 @@ def _get_tensor(self, data_ptr): ret_idx = i return self.tensors[ret_idx] + def _check_in_bounds(self, op_name: str, ptr, mask) -> None: + """ + Raise IndexError before an access runs off the storage of a tensor arg. + + Triton's interpreter performs loads, stores and atomics on raw host + memory, so an active lane past the end of an allocation reads or + corrupts whatever follows it (glibc heap metadata, for instance, which + later crashes the process at exit). The tracer only judges memory it + was given: when any lane of the pointer block, masked-off lanes + included, touches a tensor arg's storage, every active lane must lie + fully inside some known storage. Accesses that never touch a known + storage, and every access of a program after it casts an integer to a + pointer (pointer tables), are left alone, so this is a best-effort + guard; the Sanitizer is the complete out-of-bounds check. + """ + element_ty = getattr(getattr(ptr, "dtype", None), "element_ty", None) + data = getattr(ptr, "data", None) + if element_ty is None or not isinstance(data, np.ndarray) or not self.storages: + return + if self._get_thread_local("int_to_ptr", False): + return # pointers built from integers cannot be attributed to an arg + + lanes = data.reshape(-1) + mask_data = True if mask is None else getattr(mask, "data", mask) + try: + active = np.broadcast_to( + np.asarray(mask_data, dtype=bool), data.shape + ).reshape(-1) + except ValueError: + return # the frontend adapter did not pair mask lanes with pointer lanes + active_lanes = lanes[active] + if active_lanes.size == 0: + return + # pointer elements (pointer tables) have no primitive_bitwidth + itemsize = max(1, getattr(element_ty, "primitive_bitwidth", 64) // 8) + lowest, highest = int(active_lanes.min()), int(active_lanes.max()) + if any( + start <= lowest and highest + itemsize <= end + for _, _, start, end in self.storages + ): + return # common case: every active lane is inside one storage + + inside = np.zeros(active_lanes.shape, dtype=bool) + for _, _, start, end in self.storages: + inside |= (active_lanes >= start) & (active_lanes <= end - itemsize) + if inside.all(): + return + # Masked-off lanes help attribute the access: Triton splits a float + # atomic_max/min into two calls by sign, and the call holding only the + # out-of-bounds lanes still points at the tensor through its other lanes. + touched = [ + storage + for storage in self.storages + if ((lanes > storage[2] - itemsize) & (lanes < storage[3])).any() + ] + if not touched: + return + + outside = active_lanes[~inside] + first_bad = int(outside[0]) + name, tensor, start, end = min( + touched, key=lambda s: max(s[2] - first_bad, first_bad - s[3] + 1, 0) + ) + base = tensor.data_ptr() + ranges = " and ".join( + f"[{int(part.min()) - base}, {int(part.max()) - base + itemsize})" + for part in (outside[outside < start], outside[outside >= start]) + if part.size + ) + frames = extract_user_frames(num_frames=1) + where = ( + f"{frames[0].filename}:{frames[0].lineno} in {frames[0].func_name}\n" + f" {frames[0].line_of_code.strip()}" + if frames + else "an unknown kernel location" + ) + raise IndexError( + f"Tracer refused an out-of-bounds {op_name} in program " + f"{self._get_thread_local('program_id')} at {where}\n" + f" {outside.size} of {active_lanes.size} active lanes fall outside " + f"`{name}` (dtype={tensor.dtype}, shape={tuple(tensor.shape)}): they " + f"touch bytes {ranges} from its data_ptr, but its storage spans " + f"[{start - base}, {end - base}).\n" + " Triton's interpreter would run this access on raw host memory, " + "reading or writing memory outside the tensor. Mask the out-of-range " + 'lanes, or run the kernel under the Sanitizer (tilelens.trace("sanitizer")) ' + "for a full out-of-bounds report." + ) + def pre_run_callback(self, fn: Callable) -> bool: return True @@ -69,8 +178,11 @@ def post_warmup_callback(self, jit_fn, ret) -> None: def arg_callback(self, name, arg, arg_cvt): if hasattr(arg, "data_ptr"): self.tensors.append(arg) + self.storages.extend(_storage_ranges(name, arg)) def grid_idx_callback(self, grid_idx: tuple[int, ...]): + self._set_thread_local("program_id", grid_idx) + self._set_thread_local("int_to_ptr", False) if self.grid_idx is not None and grid_idx != self.grid_idx: self.sample = False else: @@ -96,6 +208,8 @@ def _convert_keys_to_numpy(keys): @self.lock_fn def pre_load_callback(ptr, mask, keys): + if keys is None: + self._check_in_bounds("load", ptr, mask) if not self.sample: return @@ -114,6 +228,8 @@ def pre_load_callback(ptr, mask, keys): @self.lock_fn def pre_store_callback(ptr, mask, keys): + if keys is None: + self._check_in_bounds("store", ptr, mask) if not self.sample: return @@ -136,6 +252,17 @@ def pre_store_callback(ptr, mask, keys): rec.call_path = extract_user_frames(num_frames=1) self.records.append(rec) + @self.lock_fn + def pre_atomic_rmw_callback(ptr, mask): + self._check_in_bounds("atomic", ptr, mask) + + @self.lock_fn + def pre_atomic_cas_callback(ptr): + self._check_in_bounds("atomic_cas", ptr, None) + + def post_int_to_ptr_callback(ret, *args): + self._set_thread_local("int_to_ptr", True) + @self.lock_fn def pre_transfer_callback(src, dst, mem_src, mem_dst): # TODO: currently only works with NKI Beta 2. Make DSL-agnostic by @@ -210,6 +337,9 @@ def post_dot_callback(ret, input, other): Allocate: OpCallbacks(after_callback=post_allocate_callback), Load: OpCallbacks(before_callback=pre_load_callback), Store: OpCallbacks(before_callback=pre_store_callback), + AtomicRMW: OpCallbacks(before_callback=pre_atomic_rmw_callback), + AtomicCas: OpCallbacks(before_callback=pre_atomic_cas_callback), + IntToPtr: OpCallbacks(after_callback=post_int_to_ptr_callback), Transfer: OpCallbacks(before_callback=pre_transfer_callback), ReduceSum: OpCallbacks(after_callback=post_reduce_sum_callback), Dot: OpCallbacks(after_callback=post_dot_callback), @@ -230,4 +360,5 @@ def sample(self, value: bool) -> None: def finalize(self) -> list: with self._lock_context(): self.tensors.clear() + self.storages.clear() return self.records diff --git a/tilelens/core/frontend/triton.py b/tilelens/core/frontend/triton.py index 30b4846a..468b1917 100644 --- a/tilelens/core/frontend/triton.py +++ b/tilelens/core/frontend/triton.py @@ -193,6 +193,10 @@ mask, kwargs.get("keys"), ), + AtomicRMW: lambda _rmw_op, ptr, _val, mask, *_args, **_kwargs: AdapterResult( + ptr, mask + ), + AtomicCas: lambda ptr, *_args, **_kwargs: AdapterResult(ptr), Dot: lambda a, b, *_args, **_kwargs: AdapterResult(a, b), ReduceSum: lambda input_tensor, axis=None,