Skip to content
Open
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
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
34 changes: 34 additions & 0 deletions tests/end_to_end/test_gluon.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import importlib.util
from pathlib import Path

import numpy as np
import pytest
import torch
from triton import knobs
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
177 changes: 177 additions & 0 deletions tests/end_to_end/test_tracer.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
import numpy as np
import pytest
import torch
import triton
import triton.language as tl
Expand Down Expand Up @@ -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))
20 changes: 20 additions & 0 deletions tests/unit/test_adapters.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@
from tilelens.core.client import Client, ClientManager
from tilelens.core.data import (
AddPtr,
AtomicCas,
AtomicRMW,
BinaryOp,
CastImpl,
Dot,
Expand Down Expand Up @@ -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]
Expand Down
22 changes: 22 additions & 0 deletions tests/unit/test_tracer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Loading
Loading