From 6ddf049867e5d5755ed861584681182da6b592b5 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Sat, 3 Oct 2026 20:34:15 -0400 Subject: [PATCH] [FIX] Refuse out-of-bounds accesses in the tracer instead of corrupting memory Triton's interpreter runs loads, stores and atomics on raw host memory. An unmasked out-of-bounds access in a traced kernel therefore wrote past the end of the tensor's allocation, corrupted glibc heap metadata, and the process aborted or segfaulted later at exit (rc 134/139). This showed up on CPU-only machines but happens with CUDA tensors too, since the interpreter overruns its host copies the same way. The tracer now records the storage ranges of its tensor arguments and, in its own before-callbacks, raises IndexError before a Triton load, store or atomic (or a Gluon gl.load/gl.store) whose pointer block touches a tensor argument's storage while an active lane falls outside every known storage. Accesses that never touch a known storage and programs that cast integers to pointers are not judged, so the guard is best-effort; the Sanitizer remains the complete check. Shared code only gains argument-normalising adapters for AtomicRMW and AtomicCas. --- README.md | 2 +- tests/end_to_end/test_gluon.py | 34 ++++++ tests/end_to_end/test_tracer.py | 177 ++++++++++++++++++++++++++++++ tests/unit/test_adapters.py | 20 ++++ tests/unit/test_tracer.py | 22 ++++ tilelens/clients/tracer/tracer.py | 131 ++++++++++++++++++++++ tilelens/core/frontend/triton.py | 4 + 7 files changed, 389 insertions(+), 1 deletion(-) 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,