diff --git a/README.md b/README.md index 0efc0527e..9aff6f85d 100644 --- a/README.md +++ b/README.md @@ -124,6 +124,8 @@ uv sync --extra test # tests but no NKI support * To run core TileLens tests, run `pytest tests/`. * (if NKI installed) To run NKI-specific tests, run `pytest tests/ -m nki`. * To run all tests (Triton + NKI), run `pytest tests/ -m ""`. +* The IR-mode tests (the compiled sanitizer, below, and its IR layer) need + Triton 3.8, the one release IR mode runs on; on another release they fail. * To run visualizer web UI tests, run `npm run test:frontend`. ## Working with Examples @@ -191,6 +193,95 @@ Analyze kernels across visualization, profiling, and sanitization with a single - 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. +### Compiled sanitizer + +`Sanitizer(compile=True)` checks each launch against the kernel Triton compiles +for it (its TTIR) instead of interpreting it: out-of-bounds accesses (against the +tensor's view, strides included), integer-width overflows in address, mask and +branch arithmetic, and divisions by zero, each with a witness (program ids, lanes, +loop iteration). The kernel is compiled on the host, exactly as the JIT would +compile the launch but only through the TTIR stage, so **no GPU is needed**: CPU +tensors work, and so does a machine without a driver. From the CLI, give the +flag before the script name (the legacy `triton-sanitizer` alias takes it too): + +```sh +tile-sanitizer --compile my_script.py --my-script-flag +``` + +```py +from tilelens.clients import Sanitizer + + +@tilelens.trace(Sanitizer(compile=True)) +@triton.jit +def kernel(x_ptr, n, BLOCK: tl.constexpr): + ... +``` + +- Kernels are compiled and checked, **not run**: their outputs are never written, + so a script that checks its own results fails under `--compile`, and a launch + whose arguments the script computes from an earlier kernel's output is checked + with those unwritten values. +- Kernels are compiled for a fixed target, `cuda:89` (sm89, e.g. RTX 4090) by + default, so a verdict does not depend on the machine it was computed on. + Choose another with `Sanitizer(compile=True, target="cuda:90")` (or + `"hip:gfx942"`, or a Triton `GPUTarget`), or for every client that names none + with the environment variable `TILELENS_IR_TARGET=cuda:90` (e.g. for + `tile-sanitizer --compile`). The kernel's + own target queries (`tl.target_info.is_cuda()`, `cuda_capability_geq()`, + `is_hip()`) answer for that target too, and `TRITON_OVERRIDE_ARCH` does not + apply: the target is the one named. The TTIR can differ between targets (e.g. + tensor descriptors, target-dependent branches), and a verdict holds for the + target it was checked for. +- Targets also differ in what compiles at all: `fp8e4nv` (`torch.float8_e4m3fn`, + `tl.float8e4nv`) needs `cuda:89` or later, `num_ctas > 1` and 16-bit tensor + descriptor atomic min/max need `cuda:90`. A kernel or autotune config that fails + to compile for the target never stops the script (it was not going to run + anyway): it is reported `unsupported`, kind `compile-failed`, naming the target + and how to choose another, since it may run on a GPU of another kind unchecked. + Name a target it compiles for to check it. Only a failure no target compiles + past (a failing `tl.static_assert`, or a Python construct Triton never compiles, + unless the kernel asked Triton's driver anything first, e.g. through + `tl.target_info` or a device query it catches, or an earlier compile of the + kernel for the target did) is just a note in the verdict: that config never + launches (the autotuner skips it too). A target answer the kernel's own code + keeps from outside the check (another trace or target, the untraced program) + and never asks for again cannot be seen. A launch none of whose configs + compiled is `unsupported` (`compile-failed`), and its notes are printed with + it. Each report names where the kernel failed (`file:line`) and the innermost + error in one line. +- A call that does not match the kernel's signature (a missing, extra or + misnamed argument), or that Triton cannot key (e.g. an unhashable constexpr + value), is a bug in the call, not a compile failure: it raises the very + `TypeError` the untraced call raises, on any GPU, and the script stops there + (under `tile-sanitizer --compile` with a traceback and exit status 1). So does + a trace that also interprets the kernel (e.g. with the `Tracer`), before the + interpreter runs, which would fail on the call too. An autotuned call that + passes an autotuned meta-parameter itself raises the autotuner's own + `Conflicting meta-parameters` `ValueError`. A keyword that names no parameter + is a compile option, which the target may not know (e.g. `waves_per_eu`, a + HIP option, under a CUDA target): that is the target's `compile-failed`, and + the report names the keyword (misspelled, it fails on every GPU untraced too). + All of this holds on Triton 3.8 (below): on another release nothing is + compiled, so nothing binds the call, and the launch is `unsupported` + (`host-compile-unavailable`), a call that does not bind included (a trace + that also interprets the kernel raises the interpreter's own error for it). +- Each launch gets an `IRVerdict` in its records: `ok` is a proof for that + launch's scalar arguments, grid and tensors (`scope="launch"`); `violations` + comes with the findings; `unsupported` names what was not checked (e.g. a + data-dependent address, a construct the TTIR reader does not model, a Z3 + query that timed out, which can depend on the machine's load, a config that + failed to compile for the target, `compile-failed`, or a compile the host + cannot run, `host-compile-unavailable`, e.g. under `TRITON_KERNEL_OVERRIDE`, + `USE_IR_LOC` or a custom pipeline, which change the TTIR through stages the + host compile does not run). An autotuned + launch checks every config, each with its own arguments and grid and with a + `ConfigVerdict` of its own, also when several configs compile to one kernel. +- With the default `abort_on_error=True` the findings are printed and the + process exits with status 1; `ENABLE_SANITIZER=0` leaves kernels untraced. +- Requires Triton 3.8; on another release every launch is `unsupported` + (`host-compile-unavailable`), nothing compiled, and the program goes on. + ### Save and load traces ```py diff --git a/tests/end_to_end/test_compiled_sanitizer.py b/tests/end_to_end/test_compiled_sanitizer.py new file mode 100644 index 000000000..4869443ef --- /dev/null +++ b/tests/end_to_end/test_compiled_sanitizer.py @@ -0,0 +1,1875 @@ +"""Sanitizer(compile=True) on real kernels: each launch is compiled on the host +for the client's target, checked against its TTIR and not run, on +CPU tensors and with Triton's driver unreachable: no GPU is needed. Ports +#361's tests/end_to_end/test_compiled_sanitizer.py onto the IR lifecycle. +Counterparts on fake launches live in +tests/unit/sanitizer_compiled/test_client.py. +""" + +from __future__ import annotations + +import importlib +import inspect +import os +import subprocess +import sys +import threading +from pathlib import Path + +import pytest +import torch +import triton +import triton.language as tl +from triton.backends.compiler import GPUTarget + +import tilelens +from tilelens.clients import Sanitizer, Tracer +from tilelens.clients.sanitizer.compiled import CompiledSanitizer +from tilelens.clients.sanitizer.data import CompiledSanitizerRecord +from tilelens.core.config import DEFAULT_IR_TARGET, Config +from tilelens.core.data import Load, Store +from tilelens.ir import IRVerdict +from tilelens.ir.verdict import SourceLocation + +trace_module = importlib.import_module("tilelens.core.trace") +config_module = importlib.import_module("tilelens.core.config") +REPO = Path(__file__).resolve().parents[2] + + +def _real_compiles_available() -> bool: + # Triton imported under TRITON_INTERPRET=1 builds its own standard library + # as InterpretedFunctions, so nothing can compile for real in-process. No + # GPU is needed: IR mode compiles on the host. + import triton.language.standard as tl_standard + from triton.runtime.jit import JITFunction + + return isinstance(tl_standard.cdiv, JITFunction) + + +pytestmark = pytest.mark.skipif( + not _real_compiles_available(), + reason="Triton was imported under TRITON_INTERPRET=1: nothing compiles in-process", +) + + +@pytest.fixture(autouse=True) +def _no_driver(unreachable_driver): + """IR mode needs no GPU: Triton's driver is unreachable here, as on + a machine without one (where it raises "0 active drivers").""" + unreachable_driver("IR mode queried Triton's driver") + + +@pytest.fixture(autouse=True) +def _default_ir_target(monkeypatch): + """The default IR target, whatever TILELENS_IR_TARGET the caller + set: in the process config, and in any Config read from the environment + (configured_target below sets its own).""" + for name in ("TILELENS_IR_TARGET", "TRITON_VIZ_IR_TARGET"): + monkeypatch.delenv(name, raising=False) + monkeypatch.setattr(config_module.config, "ir_target", DEFAULT_IR_TARGET) + + +@pytest.fixture(autouse=True) +def _real_jit(monkeypatch): + # tests/unit/test_multithreading.py sets TRITON_INTERPRET=1 at import time, + # and a traced launch's patch scope restores knobs.runtime.interpret as an + # explicit override. These tests need @triton.jit to build real + # JITFunctions, so pin the knob off and put back exactly what was there. + from triton import knobs + + monkeypatch.delenv("TRITON_INTERPRET", raising=False) + missing = object() + previous = knobs.runtime.__dict__.get("interpret", missing) + knobs.runtime.__dict__["interpret"] = False + yield + if previous is missing: + knobs.runtime.__dict__.pop("interpret", None) + else: + knobs.runtime.__dict__["interpret"] = previous + + +def _sanitizer() -> CompiledSanitizer: + det = Sanitizer(compile=True, abort_on_error=False) + assert isinstance(det, CompiledSanitizer) + return det + + +def _line(kernel, needle: str) -> int: + """The source line of ``kernel`` (a JITFunction) holding ``needle``.""" + lines, start = inspect.getsourcelines(kernel.fn) + (index,) = [i for i, line in enumerate(lines) if needle in line] + return start + index + + +def _lane(record: CompiledSanitizerRecord) -> int: + (lane,) = [v for k, v in record.witness.items() if k.startswith("arange_")] + return lane + + +def _make_add(): + @triton.jit + def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr): + pid = tl.program_id(0) + offs = pid * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + x = tl.load(x_ptr + offs, mask=mask) + y = tl.load(y_ptr + offs, mask=mask) + tl.store(out_ptr + offs, x + y, mask=mask) + + return add_kernel + + +def _make_add_nomask(): + @triton.jit + def add_nomask(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + pid = tl.program_id(0) + offs = pid * BLOCK + tl.arange(0, BLOCK) + x = tl.load(x_ptr + offs) # no mask: OOB on a ragged tail + tl.store(out_ptr + offs, x) + + return add_nomask + + +# ======== proofs and findings ========= + + +def test_an_in_bounds_launch_is_proved_and_nothing_runs(): + det = _sanitizer() + traced = tilelens.trace(det)(_make_add()) + n = 3000 # a ragged tail, masked + x, y = torch.randn(n), torch.randn(n) + out = torch.zeros(n) + + kernel = traced[(triton.cdiv(n, 1024),)](x, y, out, n, BLOCK=1024) + + assert det.last_status == "ok" and det.records == [] + verdict = det.last_verdict + assert trace_module.launches[-1].records == [verdict] + (config,) = verdict.per_config + assert (config.specialization, config.status) == (kernel.hash, "ok") + assert verdict.refusal is None and verdict.notes == () + # LAUNCH="skip": compiled and checked, never run. + assert torch.count_nonzero(out) == 0 + + +def test_an_unmasked_tail_is_reported_at_its_line_and_device_address(): + det = _sanitizer() + add_nomask = _make_add_nomask() + traced = tilelens.trace(det)(add_nomask) + n = 3000 + x, out = torch.randn(n), torch.zeros(n) + + traced[(triton.cdiv(n, 1024),)](x, out, n, BLOCK=1024) + + assert det.last_status == "violations" + load, store = det.records + assert trace_module.launches[-1].records == [load, store, det.last_verdict] + assert [(r.kind, r.op_type, r.tensor_name) for r in (load, store)] == [ + ("out-of-bounds", Load, "x_ptr"), + ("out-of-bounds", Store, "out_ptr"), + ] + for record, tensor, needle in ((load, x, "tl.load"), (store, out, "tl.store")): + offset = record.violation_offset + assert n <= offset < 3 * 1024 + assert record.witness["pid_0"] == 2 + assert 2 * 1024 + _lane(record) == offset + # the address of the offending element + assert record.violation_address == tensor.data_ptr() + offset * 4 + assert record.tensor_facts.data_ptr == tensor.data_ptr() + (tb,) = record.user_code_tracebacks + assert (tb.filename, tb.lineno, tb.func_name) == ( + __file__, + _line(add_nomask, needle), + "add_nomask", + ) + assert needle in tb.line_of_code + assert torch.count_nonzero(out) == 0 + + +def test_each_launch_is_checked_against_its_own_arguments(): + det = _sanitizer() + traced = tilelens.trace(det)(_make_add_nomask()) + + # An exact multiple of BLOCK: in bounds. + x, out = torch.randn(4096), torch.zeros(4096) + traced[(4,)](x, out, 4096, BLOCK=1024) + assert (det.last_status, det.records) == ("ok", []) + + # A ragged tail with the same specialization: out of bounds. + x, out = torch.randn(3000), torch.zeros(3000) + traced[(3,)](x, out, 3000, BLOCK=1024) + assert det.last_status == "violations" and len(det.records) == 2 + + +def test_a_loop_advancing_a_pointer_is_checked_per_iteration(): + @triton.jit + def loop_sum(x_ptr, out_ptr, n_iters, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + ptrs = x_ptr + offs + acc = tl.zeros((BLOCK,), tl.float32) + for _ in range(n_iters): + acc += tl.load(ptrs) + ptrs += BLOCK + tl.store(out_ptr + offs, acc) + + det = _sanitizer() + traced = tilelens.trace(det)(loop_sum) + x, out = torch.randn(64), torch.zeros(16) + + traced[(1,)](x, out, 4, BLOCK=16) # 4 * 16 == 64 + assert (det.last_status, det.records) == ("ok", []) + traced[(1,)](x, out, 0, BLOCK=16) # a zero-trip loop reads nothing + assert (det.last_status, det.records) == ("ok", []) + + traced[(1,)](x, out, 5, BLOCK=16) + (record,) = det.records + assert (record.kind, record.op_type, record.tensor_name) == ( + "out-of-bounds", + Load, + "x_ptr", + ) + assert record.witness["iter_loop"] == 4 + assert record.violation_offset == 4 * 16 + _lane(record) + assert record.user_code_tracebacks[0].lineno == _line(loop_sum, "tl.load") + + +def test_a_store_loop_without_an_accumulator_is_checked(): + @triton.jit + def store_loop(out_ptr, iters, BLOCK: tl.constexpr): + for i in range(0, iters): + offs = i * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.full((BLOCK,), 1.0, tl.float32)) + + det = _sanitizer() + traced = tilelens.trace(det)(store_loop) + out = torch.zeros(16) + traced[(1,)](out, 4, BLOCK=4) # 4 * 4 == 16, exactly fits + assert (det.last_status, det.records) == ("ok", []) + traced[(1,)](out, 6, BLOCK=4) # 6 * 4 == 24 > 16 + assert det.last_status == "violations" + assert {r.op_type for r in det.records} == {Store} + assert torch.count_nonzero(out) == 0 + + +def test_a_strided_view_is_checked_against_its_elements(): + @triton.jit + def strided(x_ptr, out_ptr, n, stride, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs * stride, mask=mask), mask=mask) + + @triton.jit + def ignores_stride(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask), mask=mask) + + x = torch.randn(64, 2)[:, 0] # stride 2: every other float + out = torch.zeros(64) + assert not x.is_contiguous() + + det = _sanitizer() + tilelens.trace(det)(strided)[(1,)](x, out, 64, x.stride(0), BLOCK=64) + assert (det.last_status, det.records) == ("ok", []) + + # Offsets 0..63 land in the view's gaps, which are out of bounds. + tilelens.trace(det)(ignores_stride)[(1,)](x, out, 64, BLOCK=64) + (record,) = det.records + assert (record.kind, record.op_type) == ("out-of-bounds", Load) + assert record.violation_offset % 2 == 1 + assert (record.tensor_facts.strides, record.tensor_facts.contiguous) == ( + (2,), + False, + ) + + +def test_an_i32_wrap_is_an_integer_overflow(): + @triton.jit + def i32_wrap(x_ptr, S): + pid = tl.program_id(0) + off = (pid * S) * S # wraps in i32 for S = 65536 + tl.store(x_ptr + off, 1) + + det = _sanitizer() + traced = tilelens.trace(det)(i32_wrap) + x = torch.zeros(16, dtype=torch.int32) + + traced[(4,)](x, 2) + assert (det.last_status, det.records) == ("ok", []) + + traced[(4,)](x, 65536) + (record,) = det.records + assert (record.kind, record.op_type, record.tensor_name) == ( + "integer-overflow", + Store, + "x_ptr", + ) + assert (record.violation_offset, record.violation_address) == (None, None) + assert not -(1 << 31) <= record.witness["value"] < 1 << 31 + assert record.user_code_tracebacks[0].lineno == _line(i32_wrap, "off = ") + assert torch.count_nonzero(x) == 0 + + +def test_a_division_by_a_zero_argument_is_reported(): + @triton.jit + def divide(x_ptr, d): + pid = tl.program_id(0) + q = pid // d + tl.store(x_ptr + q, 1) + + det = _sanitizer() + traced = tilelens.trace(det)(divide) + x = torch.zeros(2, dtype=torch.int32) + + traced[(4,)](x, 2) + assert (det.last_status, det.records) == ("ok", []) + + traced[(4,)](x, 0) + (record,) = det.records + assert (record.kind, record.op_type) == ("division-by-zero", Store) + (tb,) = record.user_code_tracebacks + assert (tb.lineno, tb.line_of_code.strip()) == ( + _line(divide, "q = "), + "q = pid // d", + ) + + +BIG32 = (1 << 31) - 1 + + +def _make_circular(): + """At pid 1, t = pid + BIG wraps in i32, so the kernel divides by 3 (-3) + where the unbounded reading divides by 0.""" + + @triton.jit + def circ_offset(x_ptr, BIG): + pid = tl.program_id(0) + t = pid + BIG + d = tl.where(t > BIG, 0, 3) + off = (t // d).to(tl.int64) - BIG // 3 + tl.store(x_ptr + off, 1.0) # pid 1: a wild store + + @triton.jit + def circ_loop(x_ptr, BIG): + pid = tl.program_id(0) + t = pid + BIG + d = tl.where(t > BIG, 0, -3) + hi = tl.minimum((t // d) - BIG // -3 + 1, 4) + for i in range(0, hi): + tl.store(x_ptr + i, 1.0) # pid 1: x[1..3] + + return circ_offset, circ_loop + + +@pytest.mark.parametrize("which", [0, 1], ids=["offset", "loop"]) +def test_a_wrap_that_decides_its_own_divisor_is_never_a_proof(which): + kernel = _make_circular()[which] + det = _sanitizer() + traced = tilelens.trace(det)(kernel) + x = torch.zeros(1) + + traced[(1,)](x, BIG32) # pid 0 alone: offset 0, one iteration + assert (det.last_status, det.records) == ("ok", []) + + traced[(2,)](x, BIG32) + assert det.last_status == "violations" + (record,) = det.records + assert (record.kind, record.witness["pid_0"]) == ("integer-overflow", 1) + assert record.user_code_tracebacks[0].lineno == _line(kernel, "t = pid + BIG") + + +def test_a_wrap_a_where_discards_is_no_finding(): + """i * S wraps on the lanes the where discards.""" + + @triton.jit + def guarded_scale(x_ptr, n, S, BLOCK: tl.constexpr): + i = tl.arange(0, BLOCK) + off = tl.where(i < n, (i * S) // S, 0) + tl.store(x_ptr + off, 1.0) + + det = _sanitizer() + traced = tilelens.trace(det)(guarded_scale) + x = torch.zeros(256) + traced[(1,)](x, 100, 1 << 24, BLOCK=256) + assert (det.last_status, det.records) == ("ok", []) + traced[(1,)](x, 200, 1 << 24, BLOCK=256) # lanes 128..199 read i * S: wraps + assert [r.kind for r in det.records] == ["integer-overflow"] + + +def test_launches_on_two_host_threads_check_at_the_same_time(): + """Two traced kernels, each with its own sanitizer, launched from two + host threads at once: their Z3 checks overlap (one shared Z3 context + segfaulted here).""" + kernels = [_make_add_nomask(), _make_add_nomask()] + dets = [_sanitizer(), _sanitizer()] + traced = [tilelens.trace(d)(k) for d, k in zip(dets, kernels)] + tensors = [(torch.randn(4096), torch.zeros(4096)) for _ in kernels] + for t, (x, out) in zip(traced, tensors): # compile first, one at a time + t[(3,)](x, out, 3000, BLOCK=1024) + barrier = threading.Barrier(len(kernels)) + errors: list[BaseException] = [] + statuses: list[list] = [[] for _ in kernels] + + def launch(i): + try: + barrier.wait() + for k in range(6): + # another n each time: a check of its own, not a remembered one + traced[i][(5,)](*tensors[i], 4097 + k, BLOCK=1024) + statuses[i].append(dets[i].last_status) + except BaseException as exc: # noqa: BLE001 - reported below + errors.append(exc) + + threads = [threading.Thread(target=launch, args=(i,)) for i in range(len(kernels))] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + assert errors == [] + assert statuses == [["violations"] * 6] * len(kernels) + + +def test_a_grouped_swizzle_matmul_is_modeled(): + """The tutorial-03 grouped swizzle (``//``, ``%`` and ``min`` over launch + quantities) with the ``% M`` / ``% N`` row clamps removed (TritonBench's + matmul_triton2): a proof when M, N, K cover the blocks, the real OOB + otherwise.""" + + @triton.jit + def swizzle_matmul( + a_ptr, b_ptr, c_ptr, M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn, + BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, + GROUP_M: tl.constexpr, + ): # fmt: skip + pid = tl.program_id(0) + num_pid_m = tl.cdiv(M, BLOCK_M) + num_pid_n = tl.cdiv(N, BLOCK_N) + num_pid_in_group = GROUP_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_M) + pid_m = first_pid_m + (pid % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + offs_am = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) # no `% M` clamp + offs_bn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) # no `% N` clamp + offs_k = tl.arange(0, BLOCK_K) + a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak + b_ptrs = b_ptr + offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn + acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) + for k in range(0, tl.cdiv(K, BLOCK_K)): + a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_K, other=0.0) + b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k * BLOCK_K, other=0.0) + acc += tl.dot(a, b) + a_ptrs += BLOCK_K * stride_ak + b_ptrs += BLOCK_K * stride_bk + offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N) + tl.store(c_ptrs, acc, mask=c_mask) + + def launch(m, n, k): + det = _sanitizer() + a = torch.randn(m, k) + b = torch.randn(k, n) + c = torch.empty(m, n) + grid = (triton.cdiv(m, 32) * triton.cdiv(n, 32),) + tilelens.trace(det)(swizzle_matmul)[grid]( + a, b, c, m, n, k, + a.stride(0), a.stride(1), b.stride(0), b.stride(1), + c.stride(0), c.stride(1), + BLOCK_M=32, BLOCK_N=32, BLOCK_K=32, GROUP_M=8, + ) # fmt: skip + return det + + clean = launch(64, 64, 64) + assert (clean.last_status, clean.records) == ("ok", []), clean.last_verdict + # M = N = K = 16 < BLOCK: the K-only masks leave rows 16..31 of A (and + # cols 16..31 of B) unguarded. + buggy = launch(16, 16, 16) + assert buggy.last_status == "violations", buggy.last_verdict + assert {r.tensor_name for r in buggy.records} >= {"a_ptr", "b_ptr"} + + +# ======== branches: path conditions and abstentions ========= + + +def test_a_modeled_branch_condition_leaves_no_false_witness(): + """``if t > 0: load(p + t * n_cols + offs - n_cols)`` never reads offset + -n_cols: the t == 0 iteration takes the other branch.""" + + @triton.jit + def guarded_scan(x_ptr, out_ptr, n_steps, n_cols, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + mask = offs < n_cols + acc = tl.zeros((BLOCK,), tl.float32) + for i in range(n_steps): + t = n_steps - 1 - i + if t > 0: + prev = tl.load(x_ptr + t * n_cols + offs - n_cols, mask=mask, other=0) + else: + prev = tl.zeros((BLOCK,), tl.float32) + acc += prev + tl.store(out_ptr + offs, acc, mask=mask) + + det = _sanitizer() + x, out = torch.randn(5 * 8), torch.zeros(8) + tilelens.trace(det)(guarded_scan)[(1,)](x, out, 5, 8, BLOCK=8) + assert (det.last_status, det.records) == ("ok", []), det.last_verdict + + +def test_a_data_dependent_branch_abstains(): + """A possible OOB under a branch on loaded data is unsupported: never a + witness from a branch that may not run, never "ok".""" + + @triton.jit + def flag_gated(flag_ptr, x_ptr, out_ptr, n_cols, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + mask = offs < n_cols + flag = tl.load(flag_ptr) + acc = tl.zeros((BLOCK,), tl.float32) + if flag > 0: + acc = tl.load(x_ptr + offs - n_cols, mask=mask, other=0) + tl.store(out_ptr + offs, acc, mask=mask) + + det = _sanitizer() + flag = torch.zeros(1, dtype=torch.int32) + x, out = torch.randn(8), torch.zeros(8) + tilelens.trace(det)(flag_gated)[(1,)](flag, x, out, 8, BLOCK=8) + assert (det.last_status, det.records) == ("unsupported", []) + refusal = det.last_verdict.refusal + assert refusal.kind == "unmodelable-condition" + assert refusal.loc.line == _line(flag_gated, "offs - n_cols") + + +def test_an_unguarded_oob_is_reported_beside_a_branch(): + @triton.jit + def mixed(x_ptr, out_ptr, n, flag, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + x = tl.load(x_ptr + offs) # unguarded: OOB when numel < BLOCK + if flag > 0: + x += tl.load(x_ptr + offs, mask=offs < n, other=0) # guarded, safe + tl.store(out_ptr + offs, x, mask=offs < n) + + det = _sanitizer() + x, out = torch.randn(8), torch.zeros(8) + tilelens.trace(det)(mixed)[(1,)](x, out, 8, 1, BLOCK=16) + assert det.last_status == "violations", det.last_verdict + (record,) = det.records + assert record.op_type is Load + assert record.user_code_tracebacks[0].lineno == _line(mixed, "unguarded") + + +# ======== what cannot be checked ========= + + +def test_a_gather_is_unsupported_at_its_line(): + @triton.jit + def gather(idx_ptr, src_ptr, out_ptr, n, BLOCK: tl.constexpr): + pid = tl.program_id(0) + offs = pid * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + idx = tl.load(idx_ptr + offs, mask=mask) + vals = tl.load(src_ptr + idx, mask=mask) # data-dependent address + tl.store(out_ptr + offs, vals, mask=mask) + + det = Sanitizer(compile=True) # abort_on_error: unsupported never exits + n = 1024 + idx = torch.zeros(n, dtype=torch.int32) + src, out = torch.randn(n), torch.zeros(n) + tilelens.trace(det)(gather)[(4,)](idx, src, out, n, BLOCK=256) + + assert (det.last_status, det.records) == ("unsupported", []) + refusal = det.last_verdict.refusal + assert refusal.kind == "indirect-address" + assert (refusal.loc.file, refusal.loc.line) == ( + __file__, + _line(gather, "src_ptr + idx"), + ) + assert refusal.message.startswith(f"{__file__}:{refusal.loc.line}: ") + assert det.last_verdict.per_config[0].refusal == refusal + + +def test_nested_loops_are_unsupported(): + @triton.jit + def nested(in_ptr, out_ptr, M, N, BLOCK: tl.constexpr): + for i in range(0, M): + for j in range(0, N): + offs = (i * N + j) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(in_ptr + offs)) + + det = _sanitizer() + inp, out = torch.randn(20), torch.zeros(20) + tilelens.trace(det)(nested)[(1,)](inp, out, 8, 2, BLOCK=4) + assert (det.last_status, det.records) == ("unsupported", []) + assert det.last_verdict.refusal.kind == "nested-loop" + + +# How a kernel taking a tuple, or a host TensorDescriptor, is refused: Triton +# names each TTIR argument the parameter flattens to by its path in the tuple +# (``ptrs.0``), which reads, and binds to no launch argument (a host +# descriptor's leaves still repeat a name, ``d.shape.0``). Never checked, +# never a finding either way. +_AGGREGATE_REFUSALS = { + "tuple of pointers": "missing-binding", + "tuple of ints": "missing-binding", + "descriptor": "other", +} + + +def _tuple_of_pointers(): + @triton.jit + def tuple_of_pointers(ptrs, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + vals = tl.load(ptrs[0] + offs) # unmasked: out of bounds of 32 elements + tl.store(ptrs[1] + offs, vals, mask=offs < n) + + return tuple_of_pointers, ((torch.zeros(32), torch.zeros(64)), 64), {"BLOCK": 64} + + +def _tuple_of_ints(): + @triton.jit + def tuple_of_ints(x_ptr, bounds, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + bounds[0] + tl.store(x_ptr + offs, 1.0, mask=offs < bounds[1]) # out of bounds from 8 + + return tuple_of_ints, (torch.zeros(64), (8, 72)), {"BLOCK": 64} + + +def _host_descriptor(): + from triton.tools.tensor_descriptor import TensorDescriptor + + @triton.jit + def descriptor_copy(desc, out_ptr, BM: tl.constexpr, BN: tl.constexpr): + offs = tl.arange(0, BM)[:, None] * BN + tl.arange(0, BN)[None, :] + tl.store(out_ptr + offs, desc.load([0, 0])) + + desc = TensorDescriptor.from_tensor(torch.zeros(64, 64), [32, 32]) + return descriptor_copy, (desc, torch.zeros(16 * 16)), {"BM": 32, "BN": 32} + + +@pytest.mark.parametrize( + "case, make", + [ + ("tuple of pointers", _tuple_of_pointers), + ("tuple of ints", _tuple_of_ints), + ("descriptor", _host_descriptor), + ], +) +def test_a_tuple_or_host_descriptor_parameter_is_unsupported(case, make): + kernel, args, kwargs = make() + det = _sanitizer() + tilelens.trace(det)(kernel)[(1,)](*args, **kwargs) + assert (det.last_status, det.records) == ("unsupported", []), det.last_verdict + refusal = det.last_verdict.refusal + assert refusal.kind == _AGGREGATE_REFUSALS[case], refusal + + +# ======== configs ========= + + +def test_autotune_reports_the_one_config_that_goes_out_of_bounds(): + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16}, num_warps=1), + triton.Config({"BLOCK": 128}, num_warps=1), + ], + key=["n"], + ) + @triton.jit + def copy(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) # no mask + + det = _sanitizer() + x, out = torch.randn(64), torch.zeros(64) + ret = tilelens.trace(det)(copy)[lambda meta: (triton.cdiv(64, meta["BLOCK"]),)]( + x, out, 64 + ) + + assert ret is None # a skipped autotuned launch picks no config + verdict = det.last_verdict + assert verdict.status == "violations" + assert [(c.config["BLOCK"], c.status) for c in verdict.per_config] == [ + (16, "ok"), + (128, "violations"), + ] + assert {r.config["BLOCK"] for r in det.records} == {128} + assert verdict.per_config[1].n_reports == len(det.records) == 2 + assert torch.count_nonzero(out) == 0 + + +def test_configs_compiling_to_one_kernel_are_each_checked(): + """S is a runtime int that 2 and 3 + specialize alike, so both configs compile to one kernel; only S=3, on + its larger grid, goes out of bounds. Deduplicating events by kernel + alone would check S=2's binding only and prove the launch.""" + + @triton.autotune( + configs=[ + triton.Config({"S": 2}, num_warps=1), + triton.Config({"S": 3}, num_warps=1), + ], + key=["n"], + ) + @triton.jit + def strided_copy(x_ptr, out_ptr, n, S, BLOCK: tl.constexpr): + # Program p covers [p * BLOCK * S, p * BLOCK * S + BLOCK). + offs = tl.program_id(0) * BLOCK * S + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) + + det = _sanitizer() + x, out = torch.randn(64), torch.zeros(64) + # S=2: 2 programs, up to element 47; S=3: 4 programs, up to 159. + tilelens.trace(det)(strided_copy)[lambda meta: (4 if meta["S"] == 3 else 2,)]( + x, out, 64, BLOCK=16 + ) + + verdict = det.last_verdict + assert verdict.status == "violations" + assert [(c.config["S"], c.status) for c in verdict.per_config] == [ + (2, "ok"), + (3, "violations"), + ] + assert len({c.specialization for c in verdict.per_config}) == 1 + assert {(r.tensor_name, r.config["S"]) for r in det.records} == { + ("x_ptr", 3), + ("out_ptr", 3), + } + assert torch.count_nonzero(out) == 0 + + +def test_a_config_that_fails_to_compile_is_a_note(): + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16}, num_warps=1), + triton.Config({"BLOCK": 64}, num_warps=1), + ], + key=["n"], + ) + @triton.jit + def add_one(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + tl.static_assert(BLOCK <= 32) + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask) + 1, mask=mask) + + det = _sanitizer() + x, out = torch.randn(64), torch.zeros(64) + tilelens.trace(det)(add_one)[lambda meta: (triton.cdiv(64, meta["BLOCK"]),)]( + x, out, 64 + ) + + verdict = det.last_verdict + assert verdict.status == "ok" + (config,) = verdict.per_config + assert config.config["BLOCK"] == 16 + (note,) = verdict.notes + assert "'BLOCK': 64" in note and "CompileTimeAssertionFailure" in note + + +# ======== composition, persistence and the process exit ========= + + +def test_a_mixed_trace_with_the_eager_tracer(): + """The compiled sanitizer checks the compiled kernel; the tracer + gets the interpreted run, which writes the outputs.""" + det, tracer = _sanitizer(), Tracer() + traced = tilelens.trace(tracer)(tilelens.trace(det)(_make_add())) + n = 100 + x, y = torch.randn(n), torch.randn(n) + out = torch.zeros(n) + + traced[(1,)](x, y, out, n, BLOCK=128) + + assert det.last_status == "ok" + records = trace_module.launches[-1].records + assert det.last_verdict in records + assert {type(r) for r in records if not isinstance(r, IRVerdict)} >= {Load, Store} + torch.testing.assert_close(out, x + y) + + +def test_a_real_launch_round_trips_through_a_saved_trace(tmp_path): + det = _sanitizer() + traced = tilelens.trace(det)(_make_add_nomask()) + x, out = torch.randn(3000), torch.zeros(3000) + traced[(3,)](x, out, 3000, BLOCK=1024) + launch = trace_module.launches[-1] + assert launch.records[:-1] == det.records and len(det.records) == 2 + + saved = list(trace_module.launches) + trace_module.launches[:] = [launch] + try: + path = tilelens.save(tmp_path / "trace.zip") + (loaded,) = tilelens.load(path) + finally: + trace_module.launches[:] = saved + assert loaded.records == launch.records + assert loaded.records[-1] == det.last_verdict + + +# Triton's on-disk cache keys a kernel by its source and first line, not its +# file, and a cached TTIR keeps the locs of the file first compiled: the +# comment makes each script's kernel its own (the TTIR-only host compile +# never reaches that cache; the JIT's own compile does). +_OOB_SCRIPT = """\ +import torch, triton, triton.language as tl +{prelude} + +{decorator} +@triton.jit +def add_nomask(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + # {tag} + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) + +n = {n} +x = torch.randn(n) +out = torch.zeros(n) +add_nomask[(triton.cdiv(n, 1024),)](x, out, n, BLOCK=1024) +print("launch returned") +""" + + +def _run(script: Path, *argv: str, **env_vars: str) -> subprocess.CompletedProcess: + # No GPU in the child either, and no IR target but env_vars'. This + # checkout first on the path, then whatever the caller put there (e.g. + # another Triton release). + unset = ("TRITON_INTERPRET", "TILELENS_IR_TARGET", "TRITON_VIZ_IR_TARGET") + env = {k: v for k, v in os.environ.items() if k not in unset} + path = os.pathsep.join([str(REPO), *filter(None, [os.environ.get("PYTHONPATH")])]) + env.update(PYTHONPATH=path, CUDA_VISIBLE_DEVICES="", **env_vars) + return subprocess.run( + [sys.executable, *argv], capture_output=True, text=True, env=env, cwd=REPO + ) + + +@pytest.mark.parametrize("n", [3000, 4096]) +def test_abort_on_error_exits_after_reporting(tmp_path, n): + script = tmp_path / "oob.py" + script.write_text( + _OOB_SCRIPT.format( + prelude="import tilelens\nfrom tilelens.clients import Sanitizer", + decorator="@tilelens.trace(Sanitizer(compile=True))", + n=n, + tag=script, + ) + ) + proc = _run(script, str(script)) + if n == 4096: # in bounds + assert proc.returncode == 0, proc.stderr + assert "launch returned" in proc.stdout + return + assert proc.returncode == 1, proc.stderr + assert proc.stdout.count("Out-Of-Bounds Access Detected") == 2 + lines = script.read_text().splitlines() + (line,) = [i for i, text in enumerate(lines, 1) if "tl.store(" in text] + assert f"File: {script}, Line: {line}, in add_nomask" in proc.stdout + assert "launch returned" not in proc.stdout + + +@pytest.mark.parametrize("command", ["tile-sanitizer", "triton-sanitizer"]) +def test_the_cli_compile_flag_runs_the_compiled_sanitizer(tmp_path, command): + script = tmp_path / "oob.py" + script.write_text(_OOB_SCRIPT.format(prelude="", decorator="", n=3000, tag=script)) + cli = ( + f"import sys; sys.argv = [{command!r}, '--compile', {str(script)!r}]; " + "from tilelens.wrapper import apply_sanitizer; apply_sanitizer()" + ) + proc = _run(script, "-c", cli) + assert proc.returncode == 1, proc.stderr + assert proc.stdout.count("(compiled sanitizer)") == 2 + assert f"File: {script}, " in proc.stdout + assert "launch returned" not in proc.stdout + + +# ======== the target ========= + + +def _make_two_cta_configs(): + # num_ctas=2 is an sm90+ option: a compile for the default cuda:89 (or + # any target below sm90) rejects the config. + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16}, num_ctas=1), + triton.Config({"BLOCK": 16}, num_ctas=2), + ], + key=["n"], + ) + @triton.jit + def copy_ctas(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) + + return copy_ctas + + +def test_the_default_target_is_cuda89_whatever_the_machine(): + det = _sanitizer() + assert det.ir_target is None # the configured default + x, out = torch.zeros(4096), torch.zeros(4096) + kernel = tilelens.trace(det)(_make_add_nomask())[(4,)](x, out, 4096, BLOCK=1024) + assert kernel.target == GPUTarget("cuda", 89, 32) + assert list(kernel.asm) == ["ttir"] # only what the sanitizer reads + assert det.last_status == "ok" + + +def _assert_compile_failed(refusal, target, error): + """A compile-failed refusal: the config failed to compile for ``target``, with + ``error``, and the refusal says how to check it for another target. It + is one line, starting where in the kernel's source file it failed.""" + assert refusal.kind == "compile-failed" + assert f"failed to compile for {target} (" in refusal.message + assert error in refusal.message + assert "Sanitizer(compile=True, target=...)" in refusal.message + assert "TILELENS_IR_TARGET" in refusal.message + assert "\n" not in refusal.message + loc = refusal.loc + assert loc is not None and refusal.message.startswith(f"{loc.file}:{loc.line}: ") + + +def test_a_target_passed_to_the_sanitizer_is_the_one_checked(): + x, out = torch.zeros(64), torch.zeros(64) + # For the default cuda:89 the sm90 config fails to compile: unchecked, + # and it may run on an sm90 GPU, so the launch is unsupported. + det = _sanitizer() + tilelens.trace(det)(_make_two_cta_configs())[(4,)](x, out, 64) + assert det.last_status == "unsupported" + one, two = det.last_verdict.per_config + assert (one.config["num_ctas"], one.status) == (1, "ok") + assert (two.config["num_ctas"], two.status) == (2, "unsupported") + + for target in ("cuda:90", GPUTarget("cuda", 90, 32)): + det = Sanitizer(compile=True, abort_on_error=False, target=target) + assert det.ir_target == GPUTarget("cuda", 90, 32) + tilelens.trace(det)(_make_two_cta_configs())[(4,)](x, out, 64) + assert det.last_status == "ok", det.last_verdict + per_config = det.last_verdict.per_config + assert [c.config["num_ctas"] for c in per_config] == [1, 2] + + +# ======== a kernel that fails to compile for the target ========= + + +def test_a_config_that_needs_two_ctas_is_unsupported_and_the_program_goes_on(capsys): + """An autotuned kernel one config of which needs num_ctas=2 (sm90+): that + config is unsupported, kind compile-failed, naming the target and how to + name another; the other is checked; the launch returns (None, as for any + skipped autotuned launch) even with abort_on_error, and the program + goes on to its next launch.""" + det = Sanitizer(compile=True) # abort_on_error=True + kernel = _make_two_cta_configs() + traced = tilelens.trace(det)(kernel) + x, out = torch.zeros(64), torch.zeros(64) + + assert traced[(4,)](x, out, 64) is None + + verdict = det.last_verdict + assert (verdict.status, verdict.scope, verdict.notes) == ("unsupported", None, ()) + ok, failed = verdict.per_config + assert (ok.config["num_ctas"], ok.status) == (1, "ok") + assert (failed.config["num_ctas"], failed.status) == (2, "unsupported") + assert failed.specialization is None and verdict.refusal == failed.refusal + _assert_compile_failed( + failed.refusal, "cuda:89", "num_ctas > 1 requires NVIDIA SM90" + ) + assert det.records == [] + # No source position for an option error: the kernel's def line. + assert failed.refusal.loc == SourceLocation( + kernel.fn.fn.__code__.co_filename, _line(kernel.fn, "def copy_ctas(") + ) + (line,) = capsys.readouterr().out.splitlines() + assert line.startswith("[CompiledSanitizer] not checked (config {") + assert ( + f": compile-failed: {failed.refusal.loc.file}:{failed.refusal.loc.line}: " + "it failed to compile for cuda:89 (ValueError: num_ctas > 1 requires" + ) in line + + # The program goes on: its next launch is checked like any other. + traced[(4,)](x, out, 64) + assert det.last_status == "unsupported" + assert [c.status for c in det.last_verdict.per_config] == ["ok", "unsupported"] + + +def test_a_plain_kernel_launched_with_two_ctas_is_unsupported(): + """A plain kernel whose only config fails to compile for the target: the + launch is unsupported compile-failed and returns None (no kernel to + return), and the next launch compiles as usual.""" + det = Sanitizer(compile=True) # abort_on_error=True + traced = tilelens.trace(det)(_make_add_nomask()) + x, out = torch.zeros(4096), torch.zeros(4096) + + assert traced[(4,)](x, out, 4096, BLOCK=1024, num_ctas=2) is None + + verdict = det.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + (failed,) = verdict.per_config + assert (failed.specialization, failed.config, failed.status) == ( + None, + {}, + "unsupported", + ) + assert verdict.refusal == failed.refusal + _assert_compile_failed( + failed.refusal, "cuda:89", "num_ctas > 1 requires NVIDIA SM90" + ) + + kernel = traced[(4,)](x, out, 4096, BLOCK=1024) + assert kernel.target == GPUTarget("cuda", 89, 32) + assert det.last_status == "ok" + + +def _make_block_asserting_configs(*blocks): + @triton.autotune( + configs=[triton.Config({"BLOCK": block}, num_warps=1) for block in blocks], + key=["n"], + ) + @triton.jit + def add_one(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + tl.static_assert(BLOCK <= 32) + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask) + 1, mask=mask) + + return add_one + + +def test_a_launch_none_of_whose_configs_compiled_is_unsupported(): + """Every config failed: unsupported compile-failed, whether each could + compile for another target (num_ctas=2 here: its refusal) or for none + (a failing tl.static_assert: noted), and the launch returns.""" + x, out = torch.zeros(64), torch.zeros(64) + + def grid(meta): + return (triton.cdiv(64, meta["BLOCK"]),) + + det = Sanitizer(compile=True) + kernel = _make_block_asserting_configs(64, 128) + assert tilelens.trace(det)(kernel)[grid](x, out, 64) is None + verdict = det.last_verdict + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "compile-failed") + assert verdict.per_config == () + first, second = verdict.notes + assert first.startswith("config {'BLOCK': 64, ") + assert second.startswith("config {'BLOCK': 128, ") + # Where the assertion is, and the error in one line. + path = kernel.fn.fn.__code__.co_filename + at = f"{path}:{_line(kernel.fn, 'tl.static_assert(')}: " + assert all( + f"was not checked: {at}it failed to compile for cuda:89 " + "(CompileTimeAssertionFailure), " in note + for note in verdict.notes + ) + assert verdict.refusal.message == ( + f"{path}:{_line(kernel.fn, 'def add_one(')}: no config of the launch " + "compiled for cuda:89, so nothing was checked: each failed with an " + "error of its own code whatever the target (see the notes)" + ) + + @triton.autotune(configs=[triton.Config({"BLOCK": 16}, num_ctas=2)], key=["n"]) + @triton.jit + def copy_two_ctas(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) + + det = Sanitizer(compile=True) + assert tilelens.trace(det)(copy_two_ctas)[(4,)](x, out, 64) is None + verdict = det.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + (failed,) = verdict.per_config + assert verdict.refusal == failed.refusal + _assert_compile_failed(failed.refusal, "cuda:89", "num_ctas > 1") + + +def test_a_static_assert_on_the_target_is_the_targets(): + """A tl.static_assert that asks for the target (tl.target_info) fails + for this target only: the config may run on the user's GPU, so it is + unsupported, not a note; a target it compiles for checks it.""" + + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16, "HOPPER": False}), + triton.Config({"BLOCK": 16, "HOPPER": True}), + ], + key=["n"], + ) + @triton.jit + def maybe_hopper(x_ptr, n, BLOCK: tl.constexpr, HOPPER: tl.constexpr): + tl.static_assert(not HOPPER or tl.target_info.cuda_capability_geq(9, 0)) + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + x = torch.zeros(64) + det = _sanitizer() + tilelens.trace(det)(maybe_hopper)[(4,)](x, 64) + verdict = det.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + ok, failed = verdict.per_config + assert (ok.config["HOPPER"], ok.status) == (False, "ok") + assert (failed.config["HOPPER"], failed.status) == (True, "unsupported") + _assert_compile_failed(failed.refusal, "cuda:89", "CompileTimeAssertionFailure") + + det = Sanitizer(compile=True, abort_on_error=False, target="cuda:90") + tilelens.trace(det)(maybe_hopper)[(4,)](x, 64) + assert det.last_status == "ok" + assert [c.status for c in det.last_verdict.per_config] == ["ok", "ok"] + + +def _make_device_asserting_configs(asks): + # Two configs, both calling ``asks()`` (the first compiles first); the + # second, which goes out of bounds (64 lanes, no mask, over 16 + # elements), holds only where ``asks()`` answers yes. + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16, "BIG": False}), + triton.Config({"BLOCK": 64, "BIG": True}), + ], + key=["n"], + ) + @triton.jit + def big_only(x_ptr, n, BLOCK: tl.constexpr, BIG: tl.constexpr): + yes: tl.constexpr = asks() + tl.static_assert(not BIG or yes) + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 1.0) + + return big_only + + +def test_a_static_assert_on_a_caught_device_query_is_the_devices(): + """The host compile refuses a device query; a kernel that catches that + and falls back to an answer of its own ("no big shared memory") fails + a static_assert a GPU with more might pass (an H100 would run the out + of bounds config): that config was not checked, so the launch is never + "ok", and it is no note.""" + from triton.runtime.jit import constexpr_function + + @constexpr_function + def has_big_smem(): + from triton.runtime import driver + + try: + properties = driver.active.utils.get_device_properties(0) + except Exception: + return False + return properties["max_shared_mem"] >= 200_000 + + kernel = _make_device_asserting_configs(has_big_smem) + det = _sanitizer() + tilelens.trace(det)(kernel)[(1,)](torch.zeros(16), 16) + verdict = det.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + ok, failed = verdict.per_config + assert (ok.config["BIG"], ok.status) == (False, "ok") + assert (failed.config["BIG"], failed.status) == (True, "unsupported") + _assert_compile_failed(failed.refusal, "cuda:89", "CompileTimeAssertionFailure") + assert failed.refusal.loc.line == _line(kernel.fn, "tl.static_assert(") + + +def test_a_static_assert_on_a_kept_target_answer_is_the_targets(): + """A kernel's code may keep the target answer after its first compile + asked (a memo): a later config failing a static_assert on the kept + answer, without asking again, is the target's all the same. At a + target it holds for, the config is checked (and goes out of bounds).""" + from triton.runtime.jit import constexpr_function + + @constexpr_function + def is_hopper(_memo={}): # noqa: B006 the kernel's own memo + if "arch" not in _memo: + from triton.runtime import driver + + _memo["arch"] = driver.active.get_current_target().arch + return _memo["arch"] >= 90 + + kernel = _make_device_asserting_configs(is_hopper) + det = _sanitizer() + tilelens.trace(det)(kernel)[(1,)](torch.zeros(16), 16) + assert is_hopper.fn.__defaults__[0] == {"arch": 89} + verdict = det.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + ok, failed = verdict.per_config + assert (ok.config["BIG"], ok.status) == (False, "ok") + assert (failed.config["BIG"], failed.status) == (True, "unsupported") + _assert_compile_failed(failed.refusal, "cuda:89", "CompileTimeAssertionFailure") + + is_hopper.fn.__defaults__[0].clear() + det = Sanitizer(compile=True, abort_on_error=False, target="cuda:90") + tilelens.trace(det)(kernel)[(1,)](torch.zeros(16), 16) + assert det.last_status == "violations" + assert {r.config["BIG"] for r in det.records} == {True} + + +# ======== calls that do not bind the kernel's parameters ========= + + +def _make_two_args(): + @triton.jit + def two_args(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + return two_args + + +# Calls of two_args(x_ptr, n, BLOCK) that do not bind: the arguments after +# x_ptr, and the keyword arguments. +_UNBOUND_CALLS = { + "missing-argument": ((), {"BLOCK": 64}), + "extra-argument": ((64, 64, 7), {}), + "misnamed-keyword": ((), {"N": 64, "BLOCK": 64}), + "repeated-keyword": ((64,), {"n": 64, "BLOCK": 64}), + # binds, but the JIT cannot key the call (compute_cache_key), on any GPU + "unhashable-constexpr": ((64,), {"BLOCK": [64]}), +} + + +class _StandInDriver: + """What JITFunction.run asks Triton's driver for before it binds a call, + on a machine with a GPU of the default IR target.""" + + def get_current_device(self): + return 0 + + def get_current_stream(self, device=None): + return 0 + + def get_current_target(self): + return GPUTarget("cuda", 89, 32) + + +def _untraced_error(kernel, args, kwargs) -> BaseException: + """What the untraced JIT raises for ``kernel[grid](*args, **kwargs)``, + for a call that does not bind: JITFunction.run's own binder, on a + stand-in driver (no GPU: the call fails before anything is compiled or + launched).""" + from triton.runtime.driver import driver + + owner = type(driver) + saved = owner.__dict__["active"] + stand_in = _StandInDriver() + owner.active = property(lambda self: stand_in) + try: + kernel.warmup(*args, grid=(1,), **kwargs) + except Exception as exc: + return exc + finally: + owner.active = saved + raise AssertionError("the untraced call bound") + + +@pytest.mark.parametrize("case", list(_UNBOUND_CALLS)) +def test_a_call_that_does_not_bind_raises_as_untraced(case, capsys): + """A call that does not match the kernel's signature is a bug in + the call, whatever the target: the launch raises the untraced JIT's own + error (the same type and message) instead of reporting the launch as + not checked. Nothing is printed or recorded, and the trace goes on.""" + kernel = _make_two_args() + x = torch.zeros(64) + rest, kwargs = _UNBOUND_CALLS[case] + expected = _untraced_error(kernel, (x, *rest), kwargs) + assert type(expected) is TypeError + det = Sanitizer(compile=True) # abort_on_error=True + traced = tilelens.trace(det)(kernel) + launches = len(trace_module.launches) + + with pytest.raises(TypeError) as raised: + traced[(1,)](x, *rest, **kwargs) + + assert (type(raised.value), str(raised.value)) == (type(expected), str(expected)) + assert det.last_verdict is None and det.records == [] + assert len(trace_module.launches) == launches + assert capsys.readouterr().out == "" + # The trace is not left mid-launch: a call that binds is checked. + traced[(1,)](x, 64, BLOCK=64) + assert det.last_status == "ok" + + +def test_an_autotuned_call_that_does_not_bind_raises_as_untraced(): + """Through the autotuner too: the first config's compile raises.""" + + @triton.autotune( + configs=[triton.Config({"BLOCK": 16}), triton.Config({"BLOCK": 32})], + key=["n"], + ) + @triton.jit + def tuned(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + x = torch.zeros(64) + # The JIT's binder, as the autotuner calls it for its first config. + expected = _untraced_error(tuned.fn, (x,), {"BLOCK": 16}) + det = _sanitizer() + + with pytest.raises(TypeError) as raised: + tilelens.trace(det)(tuned)[(1,)](x) + + assert str(raised.value) == str(expected) + assert det.last_verdict is None + + +def _make_tuned(): + @triton.autotune( + configs=[triton.Config({"BLOCK": 16}), triton.Config({"BLOCK": 32})], + key=["n"], + ) + @triton.jit + def tuned(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + return tuned + + +def test_an_autotuned_call_passing_a_tuned_parameter_raises_as_untraced(): + """A call that passes an autotuned meta-parameter itself: the untraced + launch's autotuning refuses it before benchmarking (Autotuner._bench's + ValueError; no device involved), and so does a traced launch, IR-only + or mixed, not a TypeError naming a tilelens internal; the trace goes + on.""" + x = torch.zeros(64) + with pytest.raises(ValueError) as untraced: + _make_tuned()[(1,)](x, 64, BLOCK=64) + assert "Conflicting meta-parameters: BLOCK" in str(untraced.value) + det = _sanitizer() + traced = tilelens.trace(det)(_make_tuned()) + + with pytest.raises(ValueError) as raised: + traced[(1,)](x, 64, BLOCK=64) + + assert str(raised.value) == str(untraced.value) + assert det.last_verdict is None + traced[(1,)](x, 64) + assert det.last_status == "ok" + mixed = tilelens.trace(Tracer())(tilelens.trace(_sanitizer())(_make_tuned())) + with pytest.raises(ValueError) as raised: + mixed[(1,)](x, 64, BLOCK=64) + assert str(raised.value) == str(untraced.value) + assert torch.equal(x, torch.zeros(64)) + + +def test_a_mixed_trace_raises_the_untraced_error_before_interpreting(): + """In a mixed trace, the call raises the untraced JIT's error + from the compile pass, before the interpreter runs, which would fail on + the same call too (with Python's own TypeError for the kernel's + function), so no output is written either way.""" + kernel = _make_two_args() + x = torch.zeros(64) + expected = _untraced_error(kernel, (x,), {"BLOCK": 64}) + traced = tilelens.trace(Tracer())(tilelens.trace(_sanitizer())(kernel)) + + with pytest.raises(TypeError) as raised: + traced[(1,)](x, BLOCK=64) + + assert str(raised.value) == str(expected) + assert torch.equal(x, torch.zeros(64)) + with pytest.raises(TypeError, match="'n'"): + tilelens.trace(Tracer())(_make_two_args())[(1,)](x, BLOCK=64) + assert torch.equal(x, torch.zeros(64)) + + +def test_an_option_the_target_does_not_know_is_a_compile_failure(): + """A keyword naming no parameter is a compile option, which the + target's backend may not know (waves_per_eu is a HIP option): no bind + failure, but the target's compile failure, so the launch is + unsupported, the program goes on, and another target checks it.""" + x = torch.zeros(64) + det = _sanitizer() + call = dict(BLOCK=64, waves_per_eu=2) + + assert tilelens.trace(det)(_make_two_args())[(1,)](x, 64, **call) is None + + refusal = det.last_verdict.refusal + assert (det.last_status, refusal.kind) == ("unsupported", "compile-failed") + assert "failed to compile for cuda:89 (KeyError:" in refusal.message + assert "waves_per_eu" in refusal.message and "TILELENS_IR_TARGET" in refusal.message + hip = Sanitizer(compile=True, abort_on_error=False, target="hip:gfx942") + tilelens.trace(hip)(_make_two_args())[(1,)](x, 64, **call) + assert hip.last_status == "ok" + + +@pytest.mark.parametrize("target", ["cuda:89", "hip:gfx942"]) +def test_an_unknown_option_is_named_not_blamed_on_the_kernel(target): + """A keyword no backend knows (a misspelled option) fails for every + target: the refusal says the call passes an option the target lacks, + that a misspelled one fails on every GPU, and that only a target of a + backend with that option can check it; it does not claim the kernel may + compile for another target.""" + det = Sanitizer(compile=True, abort_on_error=False, target=target) + + assert ( + tilelens.trace(det)(_make_two_args())[(1,)]( + torch.zeros(64), 64, BLOCK=64, bogus=1 + ) + is None + ) + + refusal = det.last_verdict.refusal + assert (det.last_status, refusal.kind) == ("unsupported", "compile-failed") + assert ( + f"the call passes 'bogus', neither a parameter of the kernel nor a " + f"compile option for {target}, so it was not checked" in refusal.message + ) + assert "a misspelled option: on every GPU" in refusal.message + assert "name a target of that backend" in refusal.message + assert "a kernel can compile for one target and fail for another" not in ( + refusal.message + ) + + +_UNBOUND_SCRIPT = """\ +import torch, triton, triton.language as tl + + +@triton.jit +def two_args(x_ptr, n, BLOCK: tl.constexpr): + # {tag} + offs = tl.arange(0, BLOCK) + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + +x = torch.zeros(64) +print("before the launch") +two_args[(1,)]({call}) # the launch +print("launch returned") +""" + + +@pytest.mark.parametrize( + "case", ["missing-argument", "extra-argument", "misnamed-keyword"] +) +def test_the_cli_raises_a_call_that_does_not_bind(tmp_path, case): + """tile-sanitizer --compile: the script fails at the call as it does + untraced, with a non-zero exit status and the traceback, whose last + line is the untraced JIT's error.""" + rest, kwargs = _UNBOUND_CALLS[case] + expected = _untraced_error(_make_two_args(), (torch.zeros(64), *rest), kwargs) + call = ", ".join( + ["x", *map(repr, rest), *(f"{k}={v!r}" for k, v in kwargs.items())] + ) + script = tmp_path / "unbound.py" + script.write_text(_UNBOUND_SCRIPT.format(tag=script, call=call)) + cli = ( + f"import sys; sys.argv = ['tile-sanitizer', '--compile', {str(script)!r}]; " + "from tilelens.wrapper import apply_sanitizer; apply_sanitizer()" + ) + + proc = _run(script, "-c", cli) + + assert proc.returncode == 1, proc.stderr + assert proc.stdout.splitlines() == ["before the launch"] + lines = script.read_text().splitlines() + (line,) = [i for i, text in enumerate(lines, 1) if "# the launch" in text] + assert "Traceback (most recent call last):" in proc.stderr + assert f'File "{script}", line {line}, in ' in proc.stderr + assert proc.stderr.rstrip().splitlines()[-1] == f"TypeError: {expected}" + + +def test_a_mixed_trace_still_interprets_after_a_compile_failure(): + """The interpreted peer runs (and writes the outputs) while + the compiled sanitizer reports the config it could not compile.""" + det, tracer = _sanitizer(), Tracer() + traced = tilelens.trace(tracer)(tilelens.trace(det)(_make_add_nomask())) + x, out = torch.randn(64), torch.zeros(64) + + traced[(1,)](x, out, 64, BLOCK=64, num_ctas=2) + + torch.testing.assert_close(out, x) + verdict = det.last_verdict + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "compile-failed") + assert any(isinstance(r, Store) for r in trace_module.launches[-1].records) + + +def test_a_hip_target_compiles_with_its_own_backend(): + det = Sanitizer(compile=True, abort_on_error=False, target="hip:gfx942") + x, out = torch.zeros(3000), torch.zeros(3000) + tilelens.trace(det)(_make_add_nomask())[(3,)](x, out, 3000, BLOCK=1024) + assert det.last_status == "violations" + assert {r.tensor_name for r in det.records} == {"x_ptr", "out_ptr"} + + +# The target-dependent branches of a kernel are the IR target's. + + +def _make_unmasked_on_cuda(): + @triton.jit + def unmasked_on_cuda(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + if tl.target_info.is_cuda(): + tl.store(x_ptr + offs, 5.0) # what every CUDA device runs + else: + tl.store(x_ptr + offs, 6.0, mask=offs < n) + + return unmasked_on_cuda + + +def _make_unmasked_from_sm89(): + @triton.jit + def unmasked_from_sm89(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + if tl.target_info.cuda_capability_geq(8, 9): + tl.store(x_ptr + offs, 1.0) + else: + tl.store(x_ptr + offs, 2.0, mask=offs < n) + + return unmasked_from_sm89 + + +def _make_unmasked_on_hip(): + @triton.jit + def unmasked_on_hip(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + if tl.target_info.is_hip(): + tl.store(x_ptr + offs, 3.0) + else: + tl.store(x_ptr + offs, 4.0, mask=offs < n) + + return unmasked_on_hip + + +# (kernel, target, status): 128 lanes over 64 elements, so the unmasked +# branch goes out of bounds. Without a GPU, tl.target_info used to read +# "no target": the default target's CUDA branch was never checked (a false +# "ok"). +TARGET_BRANCHES = [ + (_make_unmasked_on_cuda, None, "violations"), + (_make_unmasked_on_cuda, "hip:gfx942", "ok"), + (_make_unmasked_from_sm89, None, "violations"), + (_make_unmasked_from_sm89, "cuda:80", "ok"), + (_make_unmasked_from_sm89, "cuda:89", "violations"), + (_make_unmasked_from_sm89, "cuda:90", "violations"), + (_make_unmasked_on_hip, None, "ok"), + (_make_unmasked_on_hip, "hip:gfx942", "violations"), +] + + +def _check_branches(make, target): + det = Sanitizer(compile=True, abort_on_error=False, target=target) + tilelens.trace(det)(make())[(8,)](torch.zeros(64), 64, BLOCK=16) + return det.last_status, sorted( + (r.kind, r.tensor_name, r.violation_offset, tuple(sorted(r.witness.items()))) + for r in det.records + ) + + +@pytest.mark.parametrize( + "make, target, status", + TARGET_BRANCHES, + ids=[f"{m.__name__[6:]}-{t or 'default'}" for m, t, _ in TARGET_BRANCHES], +) +def test_target_dependent_branches_are_the_ir_targets(make, target, status): + assert _check_branches(make, target)[0] == status + + +class _Machine: + """A stand-in for Triton's active driver on a machine with a GPU of + ``target``; it must never be asked during IR mode's compile.""" + + def __init__(self, target): + self.target = target + + def get_current_target(self): + raise AssertionError("IR mode asked the machine's driver for its target") + + +@pytest.mark.parametrize( + "machine", + [ + GPUTarget("cuda", 89, 32), + GPUTarget("cuda", 90, 32), + GPUTarget("hip", "gfx942", 64), + ], + ids=["sm89", "sm90", "gfx942"], +) +def test_a_verdict_does_not_depend_on_the_machine(monkeypatch, machine): + """Every verdict and finding above is the same on a GPU of any kind as + without one (this module's default: the driver is unreachable).""" + without_gpu = [_check_branches(make, target) for make, target, _ in TARGET_BRANCHES] + from triton.runtime.driver import driver + + stand_in = _Machine(machine) + monkeypatch.setattr(type(driver), "active", property(lambda self: stand_in)) + on_gpu = [_check_branches(make, target) for make, target, _ in TARGET_BRANCHES] + assert on_gpu == without_gpu + assert [status for status, _ in on_gpu] == [s for _, _, s in TARGET_BRANCHES] + + +def _make_to_fp8e4nv(): + @triton.jit + def to_fp8e4nv(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + x = tl.load(x_ptr + offs, mask=mask) + tl.store(out_ptr + offs, x.to(tl.float8e4nv).to(tl.float32), mask=mask) + + return to_fp8e4nv + + +def _make_load_fp8e4nv(): + @triton.jit + def load_fp8e4nv(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store( + out_ptr + offs, tl.load(x_ptr + offs, mask=mask).to(tl.float32), mask=mask + ) + + return load_fp8e4nv + + +FP8E4NV_KERNELS = pytest.mark.parametrize( + "make, dtype", + [(_make_to_fp8e4nv, torch.float32), (_make_load_fp8e4nv, torch.float8_e4m3fn)], + ids=["cast", "tensor"], +) + + +@FP8E4NV_KERNELS +def test_fp8e4nv_compiles_under_the_default_target(make, dtype): + """The default target, cuda:89, is the first with fp8e4nv, + so such kernels are checked without naming a target.""" + x, out = torch.zeros(64).to(dtype), torch.zeros(64) + det = _sanitizer() + kernel = tilelens.trace(det)(make())[(4,)](x, out, 64, BLOCK=16) + assert kernel.target == GPUTarget("cuda", 89, 32) + assert det.last_status == "ok", det.last_verdict + + +@FP8E4NV_KERNELS +def test_an_explicit_cuda80_refuses_fp8e4nv_as_unsupported(make, dtype): + """cuda:80 has no fp8e4nv: the launch fails to compile there (whatever + this machine's GPU), which is unsupported compile-failed naming cuda:80 + and how to name another target, never an exception.""" + x, out = torch.zeros(64).to(dtype), torch.zeros(64) + det = Sanitizer(compile=True, target="cuda:80") # abort_on_error=True + assert tilelens.trace(det)(make())[(4,)](x, out, 64, BLOCK=16) is None + verdict = det.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + (failed,) = verdict.per_config + assert verdict.refusal == failed.refusal + _assert_compile_failed(failed.refusal, "cuda:80", "fp8e4nv not supported") + + +def _device_is_zero(): + # A host function a kernel's constexpr function calls (marked like + # tl.target_info.current_target), asking Triton's driver for a device. + return triton.runtime.driver.active.get_current_device() == 0 + + +_device_is_zero.__triton_builtin__ = True # type: ignore[attr-defined] + + +def test_a_host_compile_that_cannot_run_is_unsupported_not_raised(): + """A compile that asks for a device cannot run on the host: no error of + the kernel's (a GPU would compile it), so the launch is unsupported and + goes on, never failed.""" + from triton.runtime.jit import constexpr_function + + @constexpr_function + def on_device_zero(): + return _device_is_zero() + + @triton.jit + def device_dependent(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + if on_device_zero(): + tl.store(x_ptr + offs, 1.0, mask=offs < n) + + det = _sanitizer() + tilelens.trace(det)(device_dependent)[(4,)](torch.zeros(64), 64, BLOCK=16) + verdict = det.last_verdict + assert (verdict.status, verdict.refusal.kind) == ( + "unsupported", + "host-compile-unavailable", + ) + assert "'get_current_device'" in verdict.refusal.message + assert verdict.notes == () + + +def test_another_triton_release_is_unsupported_and_the_program_goes_on(monkeypatch): + """IR mode runs on Triton 3.8 only: on any other release the host + compile is unavailable, so every launch is unsupported, nothing is + compiled or run, and the program goes on. A call that does not bind is + one more such launch: the JIT's binder is never reached.""" + from tilelens.core.host_compile import triton_api + + triton_api.cache_clear() # a refusal is not cached + monkeypatch.setattr(triton, "__version__", "3.6.0") + det = _sanitizer() + traced = tilelens.trace(det)(_make_add()) + x, y, out = torch.ones(64), torch.ones(64), torch.zeros(64) + + assert traced[(4,)](x, y, out, 64, BLOCK=16) is None + + verdict = det.last_verdict + assert (verdict.status, verdict.refusal.kind) == ( + "unsupported", + "host-compile-unavailable", + ) + assert verdict.refusal.message.endswith( + "IR mode supports Triton 3.8.x only; the installed Triton is 3.6.0" + ) + assert det.records == [] and torch.equal(out, torch.zeros(64)) + assert traced[(4,)](x, y, out, BLOCK=16) is None # n is missing + assert det.last_verdict.refusal.kind == "host-compile-unavailable" + + +@pytest.fixture +def configured_target(monkeypatch): + """Set TILELENS_IR_TARGET and reload the process config from it.""" + + def set_target(value): + monkeypatch.setenv("TILELENS_IR_TARGET", value) + monkeypatch.setattr(config_module, "config", Config()) + + return set_target + + +def test_the_environment_sets_the_default_target(configured_target): + configured_target("cuda:90") + x, out = torch.zeros(64), torch.zeros(64) + det = _sanitizer() + tilelens.trace(det)(_make_two_cta_configs())[(4,)](x, out, 64) + assert [c.config["num_ctas"] for c in det.last_verdict.per_config] == [1, 2] + + # A target the client names wins over the environment's. + det = Sanitizer(compile=True, abort_on_error=False, target="cuda:80") + tilelens.trace(det)(_make_two_cta_configs())[(4,)](x, out, 64) + assert det.last_status == "unsupported" + _assert_compile_failed(det.last_verdict.refusal, "cuda:80", "num_ctas > 1") + + +@pytest.mark.parametrize("spec", ["sm90", "cuda:", "hip:942", "cuda:90:x", 90]) +def test_a_spec_that_names_no_target_is_refused_up_front(spec): + with pytest.raises(ValueError, match=f"invalid IR target {spec!r}"): + Sanitizer(compile=True, target=spec) + + +def test_a_configured_spec_that_names_no_target_fails_the_launch(configured_target): + configured_target("gfx942") + det = _sanitizer() + x, out = torch.zeros(64), torch.zeros(64) + with pytest.raises(ValueError, match=r"TILELENS_IR_TARGET\) is 'gfx942'"): + tilelens.trace(det)(_make_add_nomask())[(1,)](x, out, 64, BLOCK=64) + assert det.last_verdict is None + + +_TARGET_SCRIPT = """\ +import torch, triton, triton.language as tl + +@triton.autotune( + configs=[ + triton.Config({{"BLOCK": 16}}, num_ctas=1), + triton.Config({{"BLOCK": 16}}, num_ctas=2), + ], + key=["n"], +) +@triton.jit +def copy_ctas(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + # {tag} + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs)) + +x, out = torch.zeros(64), torch.zeros(64) +copy_ctas[(4,)](x, out, 64) +print("launch returned") +""" + + +@pytest.mark.parametrize( + "target, returncode, output", + [ + ("cuda:90", 0, "launch returned"), + # The sm90 config is reported as not checked, and the script + # goes on. + ( + None, + 0, + ": it failed to compile for cuda:89 (ValueError: " + "num_ctas > 1 requires NVIDIA SM90", + ), + ("sm90", 1, "TILELENS_IR_TARGET) is 'sm90'"), + ], +) +def test_the_cli_reads_the_target_from_the_environment( + tmp_path, target, returncode, output +): + script = tmp_path / "ctas.py" + script.write_text(_TARGET_SCRIPT.format(tag=script)) + cli = ( + f"import sys; sys.argv = ['tile-sanitizer', '--compile', {str(script)!r}]; " + "from tilelens.wrapper import apply_sanitizer; apply_sanitizer()" + ) + env = {} if target is None else {"TILELENS_IR_TARGET": target} + proc = _run(script, "-c", cli, **env) + assert proc.returncode == returncode, proc.stderr + assert output in proc.stdout + proc.stderr + assert ("launch returned" in proc.stdout) == (returncode == 0) + + +_UNCOMPILABLE_SCRIPT = """\ +import torch, triton, triton.language as tl + +@triton.jit +def copy(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + # {tag} + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=offs < n), mask=offs < n) + +@triton.autotune(configs=[triton.Config({{"BLOCK": 64}})], key=["n"]) +@triton.jit +def bounded(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + tl.static_assert(BLOCK <= 32) + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=offs < n), mask=offs < n) + +x, out = torch.zeros(64), torch.zeros(64) +copy[(1,)](x, out, 64, BLOCK=64, num_ctas=2) +print("first launch returned") +bounded[(1,)](x, out, 64) +print("second launch returned") +""" + + +def test_the_cli_goes_on_past_kernels_that_fail_to_compile(tmp_path): + """tile-sanitizer --compile on kernels no config of which compiles for + the target (num_ctas=2 below sm90; a failing tl.static_assert): no + finding, so exit status 0, each launch reported as not checked, in one + line naming where the kernel failed. A launch whose configs all failed + as notes prints the notes too: they say why.""" + script = tmp_path / "uncompilable.py" + script.write_text(_UNCOMPILABLE_SCRIPT.format(tag=script)) + cli = ( + f"import sys; sys.argv = ['tile-sanitizer', '--compile', {str(script)!r}]; " + "from tilelens.wrapper import apply_sanitizer; apply_sanitizer()" + ) + proc = _run(script, "-c", cli) + assert proc.returncode == 0, proc.stderr + lines = proc.stdout.splitlines() + assert "first launch returned" in lines and "second launch returned" in lines + source = script.read_text().splitlines() + + def at(needle): + (line,) = [i for i, text in enumerate(source, 1) if needle in text] + return f"{script}:{line}: " + + printed = [line for line in lines if line.startswith("[CompiledSanitizer]")] + first, *unchecked = printed + assert first.startswith( + "[CompiledSanitizer] not checked: compile-failed: " + f"{at('def copy(')}it failed to compile for cuda:89 (ValueError: num_ctas " + "> 1 requires NVIDIA SM90" + ) + assert unchecked == [ + "[CompiledSanitizer] not checked: compile-failed: " + f"{at('def bounded(')}no config of the launch compiled for cuda:89, so " + "nothing was checked: each failed with an error of its own code whatever " + "the target (see the notes)", + "[CompiledSanitizer] note: config {'BLOCK': 64, 'num_warps': 4, " + "'num_ctas': 1, 'num_stages': 3} was not checked: " + f"{at('tl.static_assert(')}it failed to compile for cuda:89 " + "(CompileTimeAssertionFailure), an error of its own code whatever the " + "target, so it never launches", + ] + assert "Traceback" not in proc.stderr diff --git a/tests/end_to_end/test_host_compile.py b/tests/end_to_end/test_host_compile.py new file mode 100644 index 000000000..0658a78f1 --- /dev/null +++ b/tests/end_to_end/test_host_compile.py @@ -0,0 +1,318 @@ +"""Host compile vs the JIT's own compile, for one target, on the CPU. + +IR mode compiles on the host for its target instead of through +``JITFunction.run``. The JIT compiles a launch for whatever Triton's active +driver says the device is, and a stand-in driver (a device id, a stream and +a target, nothing more) lets it compile without a GPU, for any target: the +oracle for what HostCompiler lifts from ``JITFunction.run`` and +``triton.compile`` (the binder call, the options the JIT adds, the +truncated compile's hash key), wherever these tests run. + +For every launch below and each of cuda:80, cuda:90 and hip:gfx942, the +JIT's kernel and the host's must be one kernel: the same hash (Triton's +name for the specialization), the same TTIR text, and the same +compiled-sanitizer verdict, down to each finding's witness. The launches cover the JIT's integer typing at the +i32/i64 boundary (2**31 - 1, 2**31, -2**31, 2**32), the equal-to-1 +specialization, bool and float arguments, a tensor descriptor and the +front end's target queries (tl.target_info). +""" + +from __future__ import annotations + +import pytest +import torch +import triton +import triton.language as tl +from triton.backends.compiler import GPUTarget + +import tilelens +from tilelens.clients import Sanitizer +from tilelens.core.client import ClientManager, LaunchCall +from tilelens.core.host_compile import HostCompiler, parse_ir_target +from tilelens.ir.ttir_reader import UnsupportedTTIR, parse_ttir + + +def _real_compiles_available() -> bool: + # Triton imported under TRITON_INTERPRET=1 builds its own standard library + # as InterpretedFunctions, so nothing can compile for real in-process. + import triton.language.standard as tl_standard + from triton.runtime.jit import JITFunction + + return isinstance(tl_standard.cdiv, JITFunction) + + +pytestmark = pytest.mark.skipif( + not _real_compiles_available(), + reason="Triton was imported under TRITON_INTERPRET=1: nothing compiles in-process", +) + +CUDA80 = GPUTarget("cuda", 80, 32) +# IR target -> the stand-in's device id: JITFunction.device_caches keeps a +# target per device, so every target gets a device of its own. +TARGETS = {"cuda:80": 0, "cuda:90": 1, "hip:gfx942": 2} + + +@pytest.fixture(autouse=True) +def _real_jit(monkeypatch, tmp_path_factory): + # tests/unit/test_multithreading.py sets TRITON_INTERPRET=1 at import + # time; pin the knob off so @triton.jit builds real JITFunctions. The + # JIT's compiles go to a cache of this module's own: its kernels are + # compiled here, never read back from another checkout's cache. + from triton import knobs + + monkeypatch.delenv("TRITON_INTERPRET", raising=False) + monkeypatch.setenv( + "TRITON_CACHE_DIR", str(tmp_path_factory.getbasetemp() / "jit-cache") + ) + missing = object() + previous = knobs.runtime.__dict__.get("interpret", missing) + knobs.runtime.__dict__["interpret"] = False + yield + if previous is missing: + knobs.runtime.__dict__.pop("interpret", None) + else: + knobs.runtime.__dict__["interpret"] = previous + + +class _StandInDriver: + """What JITFunction.run asks Triton's active driver for, on a machine + whose device ``device`` is a GPU of ``target``.""" + + def __init__(self, target, device): + self.target, self.device = target, device + + def get_current_device(self): + return self.device + + def get_current_stream(self, device=None): + return 0 + + def get_current_target(self): + return self.target + + +@pytest.fixture(params=list(TARGETS)) +def target(request, monkeypatch): + """The IR target, and a stand-in driver for a device of it.""" + from triton.runtime.driver import driver + + target = parse_ir_target(request.param) + stand_in = _StandInDriver(target, TARGETS[request.param]) + monkeypatch.setattr(type(driver), "active", property(lambda self: stand_in)) + return target + + +def _masked_copy(): + @triton.jit + def masked_copy(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask), mask=mask) + + return masked_copy + + +def _strided_store(): + @triton.jit + def strided_store(x_ptr, S, BLOCK: tl.constexpr): + off = tl.program_id(0) * S # an i32 product for an i32 S + tl.store(x_ptr + off + tl.arange(0, BLOCK), 1.0) + + return strided_store + + +def _flagged_scale(): + @triton.jit + def flagged_scale(x_ptr, out_ptr, flag, scale, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + if flag: + tl.store(out_ptr + offs, tl.load(x_ptr + offs) * scale) # unmasked + + return flagged_scale + + +def _target_branches(): + @triton.jit + def target_branches(x_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + if tl.target_info.cuda_capability_geq(9, 0): + tl.store(x_ptr + offs, 1.0) # unmasked from sm90 + elif tl.target_info.is_hip(): + tl.store(x_ptr + offs, 2.0) # unmasked on HIP + else: + tl.store(x_ptr + offs, 3.0, mask=offs < n) + + return target_branches + + +def _descriptor_bump(): + @triton.jit + def descriptor_bump(desc, BLOCK: tl.constexpr): + desc.store([0, 0], desc.load([0, 0]) + 1) + + return descriptor_bump + + +def _descriptor(): + from triton.tools.tensor_descriptor import TensorDescriptor + + return (TensorDescriptor.from_tensor(torch.zeros(64, 64), [16, 16]),) + + +BOUNDARY = [2**31 - 1, 2**31, -(2**31), 2**32, 1, 0, 64] + + +def _cases(): + """(id, kernel factory, args builder, kwargs, grid).""" + cases = [] + for n in BOUNDARY: + cases.append( + ( + f"masked_copy-n={n}", + _masked_copy, + lambda n=n: (torch.zeros(64), torch.zeros(64), n), + {"BLOCK": 16}, + (8,), # 128 lanes over 64 elements: in bounds iff n <= 64 + ) + ) + cases.append( + ( + f"strided_store-S={n}", + _strided_store, + lambda n=n: (torch.zeros(64), n), + {"BLOCK": 16}, + (3,), + ) + ) + for flag in (True, False): + for scale in (1.5, -0.0): + cases.append( + ( + f"flagged_scale-flag={flag}-scale={scale}", + _flagged_scale, + lambda flag=flag, scale=scale: ( + torch.zeros(64), + torch.zeros(64), + flag, + scale, + ), + {"BLOCK": 16}, + (8,), # 128 lanes over 64 elements: out of bounds if flag + ) + ) + cases.append( + ( + "target_branches", + _target_branches, + lambda: (torch.zeros(64), 64), + {"BLOCK": 16}, + (8,), # out of bounds on sm90+ and HIP only + ) + ) + cases.append( + ("descriptor_bump", _descriptor_bump, _descriptor, {"BLOCK": 16}, (1,)) + ) + return cases + + +CASES = _cases() + + +def _reading(text: str): + try: + return parse_ttir(text) + except UnsupportedTTIR as e: + return ("refused", e.kind, str(e.loc)) + + +def _verdict_from(kernel, jit_fn, args, kwargs, grid): + """The compiled sanitizer's verdict and findings for one launch whose + compiled kernel is ``kernel``, delivered as the core delivers it.""" + san = Sanitizer(compile=True, abort_on_error=False) + manager = ClientManager([san]) + manager.begin_launch( + LaunchCall(jit_fn=jit_fn, args=args, kwargs=kwargs, grid=grid, capture=True) + ) + event = ClientManager._launch_event( + jit_fn, args, kwargs, grid, kernel, target=kernel.metadata.target + ) + manager._dispatch_ir("before_launch", event, manager.ir_clients()) + manager.finalize() + return san.last_verdict, san.records + + +def _summary(verdict, records): + return ( + verdict.status, + verdict.scope, + None if verdict.refusal is None else verdict.refusal.kind, + tuple((c.specialization, c.status, c.n_reports) for c in verdict.per_config), + sorted( + ( + r.kind, + r.op_type.__name__, + r.tensor_name, + r.violation_offset, + tuple(sorted(r.witness.items())), + ) + for r in records + ), + ) + + +@pytest.mark.parametrize( + "make, build, kwargs, grid", [c[1:] for c in CASES], ids=[c[0] for c in CASES] +) +def test_the_jit_and_the_host_compile_the_same_kernel( + target, make, build, kwargs, grid +): + kernel = make() + args = build() + jit = kernel.warmup(*args, grid=grid, **kwargs) + host = HostCompiler().compile(kernel, args, kwargs, target=target) + + assert jit.metadata.target == host.target == target + assert host.hash == jit.hash + assert host.asm["ttir"] == jit.asm["ttir"] + from_host = _verdict_from(host, kernel, args, kwargs, grid) + assert _summary(*from_host) == _summary( + *_verdict_from(jit, kernel, args, kwargs, grid) + ) + + # A traced launch (the host path end to end, which never asks the + # driver) reaches that verdict too. + san = Sanitizer(compile=True, abort_on_error=False, target=target) + tilelens.trace(san)(kernel)[grid](*args, **kwargs) + assert _summary(san.last_verdict, san.records) == _summary(*from_host) + + +def test_the_cases_are_not_vacuous(): + """The launches exercise both verdicts, every finding kind they can, + both integer widths, and branches the target decides.""" + statuses, kinds, widths = set(), set(), set() + for _, make, build, kwargs, grid in CASES: + kernel = make() + args = build() + host = HostCompiler().compile(kernel, args, kwargs, target=CUDA80) + graph = _reading(host.asm["ttir"]) + if not isinstance(graph, tuple): + widths |= {a.int_bits for a in graph.func_args if a.int_bits} + verdict, records = _verdict_from(host, kernel, args, kwargs, grid) + statuses.add(verdict.status) + kinds |= {r.kind for r in records} + assert {"ok", "violations"} <= statuses + assert kinds == {"out-of-bounds", "integer-overflow"} + assert {32, 64} <= widths + + by_target = { + spec: HostCompiler() + .compile( + _target_branches(), + (torch.zeros(64), 64), + {"BLOCK": 16}, + target=parse_ir_target(spec), + ) + .asm["ttir"] + for spec in TARGETS + } + assert len(set(by_target.values())) == len(TARGETS) diff --git a/tests/end_to_end/test_ir_lifecycle_compiled.py b/tests/end_to_end/test_ir_lifecycle_compiled.py new file mode 100644 index 000000000..21da22c07 --- /dev/null +++ b/tests/end_to_end/test_ir_lifecycle_compiled.py @@ -0,0 +1,836 @@ +"""End-to-end tests of the core IR lifecycle on real kernels: IR clients receive +kernels compiled on the host through ClientManager.ir_capture, without the +real launch, across plain, @heuristics and @autotune∘@heuristics kernels. +Counterparts on a fake compile live in tests/unit/test_ir_lifecycle.py. + +Only what needs a device runs on one: an untraced or interpreted real +launch, device-memory accounting, and an interpreting client's voted warmup +(the JIT's own compile). Everything else runs on CPU tensors with Triton's +driver unreachable (``_no_driver``), as on a machine without a GPU. +""" + +import importlib +import types + +import pytest +import torch +import triton +import triton.language as tl +from triton.backends.compiler import GPUTarget +from triton.compiler.errors import CompilationError, CompileTimeAssertionFailure + +import tilelens +from tilelens.core.callbacks import ForLoopCallbacks, OpCallbacks +from tilelens.core.client import Client +from tilelens.core.config import DEFAULT_IR_TARGET +from tilelens.core.data import Store +from tilelens.core.host_compile import HostCompiler, HostKernel + +# `tilelens.core.trace` the attribute is the trace() decorator; the module +# holds the `launches` list. +trace_module = importlib.import_module("tilelens.core.trace") +config_module = importlib.import_module("tilelens.core.config") + + +def _real_compiles_available() -> bool: + # Triton imported under TRITON_INTERPRET=1 builds its own standard library + # as InterpretedFunctions, so nothing can compile for real in-process. + import triton.language.standard as tl_standard + from triton.runtime.jit import JITFunction + + return isinstance(tl_standard.cdiv, JITFunction) + + +pytestmark = pytest.mark.skipif( + not _real_compiles_available(), + reason="Triton was imported under TRITON_INTERPRET=1: nothing compiles in-process", +) + +GPU_REASON = "a real launch (or the JIT's own compile) needs a CUDA GPU" +# Marks what needs a device; everything else runs without Triton's driver. +needs_gpu = pytest.mark.skipif(not torch.cuda.is_available(), reason=GPU_REASON) +# The default IR target (sm89: the first that compiles fp8e4nv), and one +# passed explicitly. +CUDA89 = GPUTarget("cuda", 89, 32) +CUDA80 = GPUTarget("cuda", 80, 32) + + +@pytest.fixture(autouse=True) +def _no_driver(request, unreachable_driver): + """Unless the test is marked needs_gpu: Triton's driver is unreachable, + as on a machine without a GPU, so any driver query on the IR path fails + the test.""" + if any( + mark.kwargs.get("reason") == GPU_REASON + for mark in request.node.iter_markers("skipif") + ): + return + unreachable_driver("the IR path queried Triton's driver") + + +@pytest.fixture(autouse=True) +def _default_ir_target(monkeypatch): + """The default IR target, whatever TILELENS_IR_TARGET the caller + set: in the process config, and in any Config read from the environment.""" + for name in ("TILELENS_IR_TARGET", "TRITON_VIZ_IR_TARGET"): + monkeypatch.delenv(name, raising=False) + monkeypatch.setattr(config_module.config, "ir_target", DEFAULT_IR_TARGET) + + +@pytest.fixture(autouse=True) +def _real_jit(monkeypatch): + # tests/unit/test_multithreading.py sets TRITON_INTERPRET=1 at import time, + # and a traced launch's patch scope restores knobs.runtime.interpret as an + # explicit override. These tests need @triton.jit to build real + # JITFunctions, so pin the knob off and put back exactly what was there. + from triton import knobs + + monkeypatch.delenv("TRITON_INTERPRET", raising=False) + missing = object() + previous = knobs.runtime.__dict__.get("interpret", missing) + knobs.runtime.__dict__["interpret"] = False + yield + if previous is missing: + knobs.runtime.__dict__.pop("interpret", None) + else: + knobs.runtime.__dict__["interpret"] = previous + + +class _IRClient(Client): + NEEDS_INTERPRETER = False + IR_STAGES = frozenset({"ttir"}) + + def __init__(self): + super().__init__() + self.log: list = [] + self.events: list = [] + self.failures: list = [] + self.finalized: list = [] + + def begin_launch(self, call): + self.log.append("begin") + self.events = [] + self.failures = [] + + def abort_launch(self, exc): + self.log.append("abort") + + def before_launch(self, event): + self.log.append("before") + self.events.append(event) + + def after_launch(self, event): + self.log.append("after") + + def compile_failed(self, event): + self.log.append("compile_failed") + self.failures.append(event) + + def finalize(self): + self.log.append("finalize") + self.finalized.append(list(self.events)) + return [] + + def pre_warmup_callback(self, jit_fn, *args, **kwargs): + return False + + def post_warmup_callback(self, jit_fn, ret): + pass + + def _unreachable(self, *args, **kwargs): + raise AssertionError(f"interpreter hook reached IR client {self.NAME}") + + pre_run_callback = _unreachable + post_run_callback = _unreachable + arg_callback = _unreachable + grid_callback = _unreachable + grid_idx_callback = _unreachable + register_op_callback = _unreachable + register_for_loop_callback = _unreachable + + +class _SkipIRClient(_IRClient): + NAME = "ir_skip" + LAUNCH = "skip" + + +class _EagerCounter(Client): + """Interpreting client counting stores; optionally votes for a warmup.""" + + NAME = "eager_counter" + + def __init__(self, warmup_vote=False): + super().__init__() + self.stores = 0 + self.warmup_vote = warmup_vote + self.warmups: list = [] + + def _on_store(self, *args, **kwargs): + self.stores += 1 + + def pre_run_callback(self, fn): + return True + + def post_run_callback(self, fn): + return True + + def arg_callback(self, name, arg, arg_cvt): + pass + + def grid_callback(self, grid): + pass + + def grid_idx_callback(self, grid_idx): + pass + + def register_op_callback(self, op_type, *args, **kwargs): + if op_type is Store: + return OpCallbacks(before_callback=self._on_store) + return OpCallbacks() + + def register_for_loop_callback(self): + return ForLoopCallbacks() + + def finalize(self): + return [] + + def pre_warmup_callback(self, jit_fn, *args, **kwargs): + return self.warmup_vote + + def post_warmup_callback(self, jit_fn, ret): + self.warmups.append(ret) + + +def _make_add_one(): + @triton.jit + def add_one(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask) + 1, mask=mask) + + return add_one + + +def _make_autotuned(**autotune_kwargs): + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16}, num_warps=1), + triton.Config({"BLOCK": 32}, num_warps=2), + ], + key=["n"], + **autotune_kwargs, + ) + @triton.heuristics({"EVEN": lambda args: args["n"] % args["BLOCK"] == 0}) + @triton.jit + def add_one_tuned(x_ptr, out_ptr, n, BLOCK: tl.constexpr, EVEN: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + if EVEN: + tl.store(out_ptr + offs, tl.load(x_ptr + offs) + 1) + else: + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask) + 1, mask=mask) + + return add_one_tuned + + +def _grid(meta): + return (triton.cdiv(meta["n"], meta["BLOCK"]),) + + +def _grid64(meta): + # The interpreter hands grid callables tensor-converted runtime args, so + # interpreted launches may only read constexprs here. + return (triton.cdiv(64, meta["BLOCK"]),) + + +def _inputs(n=64, device="cpu"): + x = torch.arange(n, dtype=torch.float32, device=device) + return x, torch.zeros_like(x) + + +def test_ir_only_skip_compiles_on_the_host_without_launching(): + ir = _SkipIRClient() + kernel = _make_add_one() + hooks = [] + kernel.add_pre_run_hook(lambda *args, **kwargs: hooks.append(1)) + traced = tilelens.trace(ir)(kernel) + x, out = _inputs() + + ret = traced[(4,)](x, out, 64, BLOCK=16) + + assert torch.equal(out, torch.zeros_like(x)) + assert ir.log == ["begin", "before", "after", "finalize"] + (events,) = ir.finalized + (event,) = events + # Compiled on the host for the default target, through TTIR only. + assert isinstance(event.kernel, HostKernel) + assert event.target == event.kernel.target == CUDA89 + assert list(event.kernel.asm) == ["ttir"] + assert "tt.func" in event.kernel.asm["ttir"] + assert event.resolved_grid == (4, 1, 1) + assert event.specialization == event.kernel.hash + assert "run" not in vars(traced.jit_fn) + # JITFunction.run was never entered: no pre_run_hook fired. + assert hooks == [] + # One config: the launch returns its kernel, as the untraced launch + # would, and the Launch carries the grid, but no tensor (none may + # outlive the launch). + assert ret is event.kernel + launch = trace_module.launches[-1] + assert launch.grid == (4, 1, 1) + assert not launch.tensors + + +def test_a_relaunch_reuses_the_traces_host_compile(monkeypatch): + compiles = [] + compile = HostCompiler._compile_source + + def counting(*args, **kwargs): + compiles.append(1) + return compile(*args, **kwargs) + + monkeypatch.setattr(HostCompiler, "_compile_source", staticmethod(counting)) + ir = _SkipIRClient() + traced = tilelens.trace(ir)(_make_add_one()) + x, out = _inputs() + + first = traced[(4,)](x, out, 64, BLOCK=16) + again = traced[(4,)](torch.zeros(64), torch.zeros(64), 64, BLOCK=16) + other = traced[(4,)](x, out, 63, BLOCK=16) # 63: not divisible by 16 + + assert again is first and other is not first + assert len(compiles) == 2 + assert [len(events) for events in ir.finalized] == [1, 1, 1] + assert ir.finalized[0][0].specialization == ir.finalized[1][0].specialization + + +def test_autotune_over_heuristics_reports_every_config(): + user = _make_autotuned() + ir = _SkipIRClient() + traced = tilelens.trace(ir)(user) + x, out = _inputs() + + rets = [traced[_grid](x, out, 64), traced[_grid](x, out, 64)] + grids = [launch.grid for launch in trace_module.launches[-2:]] + + # Every launch reports every config, whatever the autotune cache holds. + for events in ir.finalized: + assert [e.kwargs["BLOCK"] for e in events] == [16, 32] + assert {e.kwargs["EVEN"] for e in events} == {True} + assert len({e.specialization for e in events}) == 2 + assert torch.equal(out, torch.zeros_like(x)) + # No config was picked: nothing to return, and the configs' grids + # differ, so the Launch has none. + assert rets == [None, None] and grids == [None, None] + assert user.cache == {} + + +def _make_runtime_stride_configs(): + # S is a runtime int that 2 and 3 specialize alike: both configs compile + # to one kernel, yet cover other elements (and here grids). + @triton.autotune( + configs=[ + triton.Config({"S": 2}, num_warps=1), + triton.Config({"S": 3}, num_warps=1), + ], + key=["n"], + ) + @triton.jit + def strided_copy(x_ptr, out_ptr, n, S, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK * S + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask), mask=mask) + + return strided_copy + + +def _grid_per_stride(meta): + return (64 // (16 * meta["S"]),) # S=2: 2 programs, S=3: 1 + + +def test_configs_compiling_to_one_kernel_each_get_an_event(): + """The dedup key holds the binding, so the second config is not + lost behind the first one's specialization.""" + ir = _SkipIRClient() + traced = tilelens.trace(ir)(_make_runtime_stride_configs()) + x, out = _inputs() + + traced[_grid_per_stride](x, out, 64, BLOCK=16) + + (events,) = ir.finalized + assert [(e.kwargs["S"], e.resolved_grid) for e in events] == [ + (2, (2, 1, 1)), + (3, (1, 1, 1)), + ] + assert len({e.specialization for e in events}) == 1 + + +def _compiled_sanitizer(): + from tilelens.clients import Sanitizer + + return Sanitizer(compile=True, abort_on_error=False) + + +def test_ir_only_launches_retain_no_tensor(): + """Without a GPU: a harness launching an IR-only trace in a loop + with fresh tensors, calling tilelens.clear() after each launch, holds on + to none of them (the grid callable closes over them, too); neither does + the trace's host-compile cache.""" + import gc + import weakref + + traced = tilelens.trace(_compiled_sanitizer())(_make_add_one()) + refs = [] + + def launch(): + x = torch.empty(4096) + out = torch.empty_like(x) + + def grid(meta): + return (triton.cdiv(x.numel(), meta["BLOCK"]),) + + traced[grid](x, out, 4096, BLOCK=1024) + assert not trace_module.launches[-1].tensors + tilelens.clear() + refs.extend((weakref.ref(x), weakref.ref(out))) + + for _ in range(3): + launch() + gc.collect() + + assert [ref() for ref in refs] == [None] * 6 + + +@needs_gpu +def test_ir_only_launches_retain_no_device_memory(): + """A harness launching an IR-only trace in a loop with fresh + device tensors, calling tilelens.clear() after each launch, holds on to + none of them (the grid callable closes over them, too).""" + traced = tilelens.trace(_compiled_sanitizer())(_make_add_one()) + n = 16 * 2**20 # 64 MiB of float32 + + def launch(): + x = torch.empty(n, device="cuda") + out = torch.empty_like(x) + + def grid(meta): + return (triton.cdiv(x.numel(), meta["BLOCK"]),) + + traced[grid](x, out, n, BLOCK=1024) + assert not trace_module.launches[-1].tensors + tilelens.clear() + + # Measured from before the first launch: the manager's Launch is + # replaced per launch, so tensors it held would be the last launch's. + torch.cuda.synchronize() + base = torch.cuda.memory_allocated() + for _ in range(10): + launch() + torch.cuda.synchronize() + + assert torch.cuda.memory_allocated() - base < 2**20 + + +def test_plain_heuristics_kernel_fires_events(): + @triton.heuristics({"BLOCK": lambda args: 16}) + @triton.jit + def heur_add_one(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask) + 1, mask=mask) + + ir = _SkipIRClient() + traced = tilelens.trace(ir)(heur_add_one) + x, out = _inputs() + + traced[_grid](x, out, 64) + + (events,) = ir.finalized + assert [e.kwargs["BLOCK"] for e in events] == [16] + assert events[0].resolved_grid == (4, 1, 1) + assert torch.equal(out, torch.zeros_like(x)) + + +def test_a_compile_error_is_data_and_the_next_launch_is_clean(): + @triton.jit + def bounded(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + tl.static_assert(BLOCK <= 32) + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask) + 1, mask=mask) + + ir = _SkipIRClient() + traced = tilelens.trace(ir)(bounded) + x, out = _inputs() + + # The only config failed to host-compile: reported as data, and the + # skipped launch then ends normally. + assert traced[(1,)](x, out, 64, BLOCK=64) is None + assert ir.log == ["begin", "compile_failed", "finalize"] + assert ir.finalized == [[]] + (failure,) = ir.failures + assert isinstance(failure.error, CompileTimeAssertionFailure) + assert failure.target == CUDA89 + assert "run" not in vars(traced.jit_fn) + + ir.log.clear() + traced[(4,)](x, out, 64, BLOCK=16) + + assert ir.log == ["begin", "before", "after", "finalize"] + # The failed launch finalized with nothing compiled. + assert [len(events) for events in ir.finalized] == [0, 1] + + +def _make_bounded_autotuned(): + @triton.autotune( + configs=[ + triton.Config({"BLOCK": 16}, num_warps=1), + # Fails to compile (static_assert) ... + triton.Config({"BLOCK": 64}, num_warps=1), + # ... compiles, but needs more threads than a block can have. + triton.Config({"BLOCK": 16}, num_warps=64), + ], + key=["n"], + ) + @triton.jit + def bounded(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + tl.static_assert(BLOCK <= 32) + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask) + 1, mask=mask) + + return bounded + + +def test_a_config_that_fails_to_compile_is_reported_and_one_too_big_is_analyzed(): + """Nothing is loaded, so a config some device could not run + (num_warps=64: more threads than a block can have) is analyzed like any + other; it can only add findings, never hide one.""" + ir = _SkipIRClient() + traced = tilelens.trace(ir)(_make_bounded_autotuned()) + x, out = _inputs() + + traced[_grid](x, out, 64) + + (events,) = ir.finalized + assert [(e.kwargs["BLOCK"], e.kwargs["num_warps"]) for e in events] == [ + (16, 1), + (16, 64), + ] + ((config, error),) = [ + ((f.kwargs["BLOCK"], f.kwargs["num_warps"]), f.error) for f in ir.failures + ] + assert config == (64, 1) and isinstance(error, CompileTimeAssertionFailure) + assert torch.equal(out, torch.zeros_like(x)) + + +def _times_two(): + # `other=` exercises the semantic._load_legacy path (`other.handle if + # other else None`) that a leaked tensor.__bool__ would break. + @triton.jit + def times_two(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.arange(0, BLOCK) + mask = offs < n + loaded = tl.load(x_ptr + offs, mask=mask, other=-1.0) + tl.store(out_ptr + offs, loaded * 2, mask=mask) + + return times_two + + +def test_mixed_trace_compiles_for_ir_and_interprets_for_eager(): + kernel = _make_add_one() + ir, eager = _SkipIRClient(), _EagerCounter() + traced = tilelens.trace(eager)(tilelens.trace(ir)(kernel)) + x, out = _inputs() + + traced[(4,)](x, out, 64, BLOCK=16) + + (events,) = ir.finalized + assert len(events) == 1 + assert list(events[0].kernel.asm) == ["ttir"] + assert eager.stores == 4 + torch.testing.assert_close(out, x + 1) + # Launch.tensors: the interpreter's copies only, not the caller's + # tensors on top. + tensors = trace_module.launches[-1].tensors + assert len(tensors) == 2 and not {id(t) for t in tensors} & {id(x), id(out)} + + # Interpreter patches must not leak into later compiles, + # of another kernel or of this one (a new BLOCK forces a recompile). + compiler = HostCompiler() + for jit_fn, block in ((_times_two(), 64), (kernel, 32)): + compiled = compiler.compile( + jit_fn, (x, out, 64), {"BLOCK": block}, target=CUDA80 + ) + assert "tt.store" in compiled.asm["ttir"] + + +@needs_gpu +def test_interpreter_patches_do_not_leak_into_later_real_launches(): + kernel = _make_add_one() + traced = tilelens.trace(_EagerCounter())(tilelens.trace(_SkipIRClient())(kernel)) + x, out = _inputs(device="cuda") + traced[(4,)](x, out, 64, BLOCK=16) + + doubled = torch.zeros_like(x) + _times_two()[(1,)](x, doubled, 64, BLOCK=64) + again = torch.zeros_like(x) + kernel[(2,)](x, again, 64, BLOCK=32) + torch.cuda.synchronize() + torch.testing.assert_close(doubled, x * 2) + torch.testing.assert_close(again, x + 1) + + +@pytest.mark.parametrize("client_cls", [_SkipIRClient, _EagerCounter]) +def test_trace_leaves_the_users_autotuner_unchanged(client_cls): + user = _make_autotuned() + heuristics, jit_fn = user.fn, user.fn.fn + before = dict(vars(user)) + before_heuristics = dict(vars(heuristics)) + traced = tilelens.trace(client_cls())(user) + x, out = _inputs() + + traced[_grid64](x, out, 64) + + assert vars(user).keys() == before.keys() + assert all(vars(user)[k] is v for k, v in before.items()) + assert vars(heuristics).keys() == before_heuristics.keys() + assert all(vars(heuristics)[k] is v for k, v in before_heuristics.items()) + assert user.fn is heuristics and heuristics.fn is jit_fn + assert user.cache == {} + + +@needs_gpu +def test_the_users_autotuner_still_autotunes_for_real_after_a_trace(): + user = _make_autotuned() + x, out = _inputs(device="cuda") + tilelens.trace(_SkipIRClient())(user)[_grid64](x, out, 64) + + untraced = torch.zeros_like(x) + user[_grid64](x, untraced, 64) + torch.cuda.synchronize() + torch.testing.assert_close(untraced, x + 1) + assert len(user.cache) == 1 + assert user.best_config in user.configs + + +def _add_one_helper(x): + return x + 1 + + +def _kernel_with_helper(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store( + out_ptr + offs, + _traced_add_one_helper(tl.load(x_ptr + offs, mask=mask)), # noqa: F821 + mask=mask, + ) + + +@pytest.fixture +def helper_kernel(): + """A kernel whose device function is a module global wrapped by + tilelens.trace, as the CLI wrappers do for every @triton.jit function. + The real code generator only accepts JITFunctions, so real compiles must + see the unwrapped binding. Built here, not at import, so the jits are real + even when TRITON_INTERPRET was set during collection.""" + module_globals = globals() + helper = tilelens.trace(_EagerCounter())(triton.jit(_add_one_helper)) + module_globals["_traced_add_one_helper"] = helper + try: + yield triton.jit(_kernel_with_helper), helper + finally: + module_globals.pop("_traced_add_one_helper", None) + + +def test_traced_device_function_compiles_under_ir_capture(helper_kernel): + kernel, helper = helper_kernel + ir = _SkipIRClient() + traced = tilelens.trace(ir)(kernel) + x, out = _inputs() + + traced[(4,)](x, out, 64, BLOCK=16) + + (events,) = ir.finalized + assert len(events) == 1 + assert "tt.func" in events[0].kernel.asm["ttir"] + assert globals()["_traced_add_one_helper"] is helper + + +def _kernel_with_package_helper(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store( + out_ptr + offs, + _helper_pkg.api.add_one(tl.load(x_ptr + offs, mask=mask)), # noqa: F821 + mask=mask, + ) + + +def test_traced_device_function_behind_a_package_path_compiles(): + # `pkg.api.add_one`, where `api` re-exports the traced helper: Triton + # resolves it through two module attributes. + module_globals = globals() + helper = tilelens.trace(_EagerCounter())(triton.jit(_add_one_helper)) + pkg = types.ModuleType("tilelens_test_helper_pkg") + pkg.api = types.ModuleType("tilelens_test_helper_pkg.api") + pkg.api.add_one = helper + module_globals["_helper_pkg"] = pkg + try: + ir = _SkipIRClient() + traced = tilelens.trace(ir)(triton.jit(_kernel_with_package_helper)) + x, out = _inputs() + + traced[(4,)](x, out, 64, BLOCK=16) + + (events,) = ir.finalized + assert "tt.func" in events[0].kernel.asm["ttir"] + assert pkg.api.add_one is helper + assert torch.equal(out, torch.zeros_like(x)) + finally: + module_globals.pop("_helper_pkg", None) + + +@needs_gpu +def test_traced_device_function_compiles_in_the_interpreted_warmup(helper_kernel): + # Interpreting clients that vote for a warmup compile for real, through + # the JIT (e.g. the profiler); the traced helper must resolve there too. + kernel, helper = helper_kernel + eager = _EagerCounter(warmup_vote=True) + traced = tilelens.trace(eager)(kernel) + x, out = _inputs(device="cuda") + + traced[(4,)](x, out, 64, BLOCK=16) + torch.cuda.synchronize() + + assert len(eager.warmups) == 1 and "ttir" in eager.warmups[0].asm + assert eager.stores == 4 + torch.testing.assert_close(out, x + 1) + assert globals()["_traced_add_one_helper"] is helper + + +def _make_apply_fn(): + @triton.jit + def apply_fn(x_ptr, out_ptr, n, FN: tl.constexpr, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, FN(tl.load(x_ptr + offs, mask=mask)), mask=mask) + + return apply_fn + + +def _make_apply_default(): + # Triton's code generator evaluates parameter defaults in the kernel's + # globals, so the default is a module global (bound by the caller). + @triton.jit + def apply_default( + x_ptr, + out_ptr, + n, + BLOCK: tl.constexpr, + FN: tl.constexpr = _traced_default_helper, # noqa: F821 + ): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, FN(tl.load(x_ptr + offs, mask=mask)), mask=mask) + + return apply_default + + +def _make_apply_first(): + @triton.jit + def apply_first(x_ptr, out_ptr, n, FNS: tl.constexpr, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, FNS[0](tl.load(x_ptr + offs, mask=mask)), mask=mask) + + return apply_first + + +def _assert_real_compiles_still_work(monkeypatch, *, launch: bool): + # The interpreter's triton.language patches must not have leaked. A + # never-seen kernel, host-compiled (never disk-cached); with ``launch`` + # (a needs_gpu test) also compiled and launched by the JIT (no disk-cache + # shortcut). + monkeypatch.setenv("TRITON_ALWAYS_COMPILE", "1") + + @triton.jit + def times_two(x_ptr, out_ptr, n, BLOCK: tl.constexpr): + offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + mask = offs < n + tl.store(out_ptr + offs, tl.load(x_ptr + offs, mask=mask) * 2, mask=mask) + + x, out = _inputs() + compiled = HostCompiler().compile( + times_two, (x, out, 64), {"BLOCK": 16}, target=CUDA80 + ) + assert "tt.store" in compiled.asm["ttir"] + if not launch: + return + x, out = _inputs(device="cuda") + times_two[(4,)](x, out, 64, BLOCK=16) + torch.cuda.synchronize() + torch.testing.assert_close(out, x * 2) + + +@pytest.mark.parametrize("passing", ["keyword", "default", "tuple"]) +def test_traced_helper_passed_as_an_argument_compiles(passing, monkeypatch): + # The CLI shape (every @triton.jit is a TritonTrace), with the helper + # reaching the real compile as a constexpr argument, which the globals + # unwrap cannot see. + helper = tilelens.trace(_EagerCounter())(triton.jit(_add_one_helper)) + if passing == "keyword": + kernel, extra = _make_apply_fn(), {"FN": helper} + elif passing == "default": + monkeypatch.setitem(globals(), "_traced_default_helper", helper) + kernel, extra = _make_apply_default(), {} + else: + kernel, extra = _make_apply_first(), {"FNS": (helper,)} + ir = _SkipIRClient() + x, out = _inputs() + + tilelens.trace(ir)(kernel)[(4,)](x, out, 64, BLOCK=16, **extra) + + assert ir.failures == [] + (events,) = ir.finalized + assert "tt.func" in events[0].kernel.asm["ttir"] + _assert_real_compiles_still_work(monkeypatch, launch=False) + + +@needs_gpu +def test_voted_warmup_compiles_with_a_traced_helper_argument(monkeypatch): + # Interpreting clients that vote for a warmup compile for real, through + # the JIT. + helper = tilelens.trace(_EagerCounter())(triton.jit(_add_one_helper)) + eager = _EagerCounter(warmup_vote=True) + x, out = _inputs(device="cuda") + + tilelens.trace(eager)(_make_apply_fn())[(4,)](x, out, 64, BLOCK=16, FN=helper) + torch.cuda.synchronize() + + assert len(eager.warmups) == 1 and "ttir" in eager.warmups[0].asm + torch.testing.assert_close(out, x + 1) + _assert_real_compiles_still_work(monkeypatch, launch=True) + + +def test_a_traced_callee_the_compile_still_reaches_fails_it_cleanly(monkeypatch): + # Should a compile reach a TritonTrace anyway (here: with the argument + # mapping switched off), the trace refuses to interpret there. The + # compile fails; triton.language is left alone for later compiles. + monkeypatch.setattr( + trace_module, "_untraced_call_args", lambda jit_fn, args, kwargs: (args, kwargs) + ) + helper = tilelens.trace(_EagerCounter())(triton.jit(_add_one_helper)) + ir = _SkipIRClient() + x, out = _inputs() + + # A compile failure like any other: data, and the skipped launch + # ends normally. + tilelens.trace(ir)(_make_apply_fn())[(4,)](x, out, 64, BLOCK=16, FN=helper) + + assert ir.log == ["begin", "compile_failed", "finalize"] + (failure,) = ir.failures + assert isinstance(failure.error, CompilationError) + assert "outside a traced launch" in str(failure.error) + _assert_real_compiles_still_work(monkeypatch, launch=False) diff --git a/tests/unit/ir/test_verdict_io.py b/tests/unit/ir/test_verdict_io.py new file mode 100644 index 000000000..9140e2f3b --- /dev/null +++ b/tests/unit/ir/test_verdict_io.py @@ -0,0 +1,305 @@ +"""Persistence of the IR-mode records: the tilelens.ir.verdict records +and the compiled sanitizer's findings are plain data that tilelens.save() / +tilelens.load() round-trip, with a public source-location type in place of +the reader's private one. No GPU. +""" + +from __future__ import annotations + +import dataclasses +import enum +import gc +import importlib +import subprocess +import sys +import weakref +import zipfile +from pathlib import Path +from types import MappingProxyType, SimpleNamespace + +import numpy as np +import pytest +import torch + +import tilelens +from tilelens.clients.sanitizer.data import CompiledSanitizerRecord +from tilelens.core.data import Launch, Load, Store +from tilelens.ir import _mlir_walk as W +from tilelens.ir.launch import TensorFacts, tensor_facts +from tilelens.ir.ttir_reader import TTIRKind, UnsupportedTTIR +from tilelens.ir.verdict import ConfigVerdict, IRVerdict, Refusal, SourceLocation +from tilelens.utils.traceback_utils import TracebackInfo + +trace_module = importlib.import_module("tilelens.core.trace") +REPO = Path(__file__).resolve().parents[3] + + +@pytest.mark.parametrize( + "module", ["tilelens.ir.verdict", "tilelens.clients.sanitizer.data"] +) +def test_record_modules_import_without_triton(module): + code = ( + "import sys\n" + f"import {module}\n" + "loaded = sorted(m for m in sys.modules if m.split('.')[0] == 'triton')\n" + "assert not loaded, loaded\n" + ) + subprocess.run([sys.executable, "-c", code], cwd=REPO, check=True) + + +# ======== source locations and refusals ========= + + +def test_refusal_holds_the_reader_loc_as_a_public_source_location(): + exc = UnsupportedTTIR( + TTIRKind.CONTROL_FLOW, "scf.while", line_no=7, loc=W.SourceLoc("k.py", 3, 5) + ) + refusal = Refusal.from_exception(exc) + + assert refusal == Refusal( + "control-flow", "scf.while", 7, SourceLocation("k.py", 3, 5) + ) + assert type(refusal.loc) is SourceLocation + # The reader's enum kind is held as its plain string, as a load gives it. + assert refusal.kind == TTIRKind.CONTROL_FLOW + assert not isinstance(refusal.kind, enum.Enum) + # Building one directly converts the reader's loc the same way; an + # object without a column gets None. + assert Refusal("call", "m", loc=W.SourceLoc("k.py", 3, 5)).loc == refusal.loc + no_col = Refusal("call", "m", loc=SimpleNamespace(file="k.py", line=2)) + assert no_col.loc == SourceLocation("k.py", 2, None) + assert Refusal("call", "m").loc is None + + +@pytest.mark.parametrize( + "fields, match", + [ + ({"loc": ("k.py", 3, 5)}, "Refusal.loc must be a SourceLocation"), + ({"loc": "k.py:3"}, "Refusal.loc must be a SourceLocation"), + ({"kind": None}, "Refusal.kind must be a str"), + ({"message": ValueError("m")}, "Refusal.message must be a str"), + ], +) +def test_refusal_rejects_fields_a_trace_cannot_hold(fields, match): + with pytest.raises(TypeError, match=match): + Refusal(**{"kind": "call", "message": "m", **fields}) + + +def test_record_hashes_agree_with_equality(): + loc = SourceLocation("k.py", 3, 5) + from_enum = Refusal(TTIRKind.CALL, "m", 4, W.SourceLoc("k.py", 3, 5)) + plain = Refusal("call", "m", 4, loc) + assert from_enum == plain and hash(from_enum) == hash(plain) + assert len({from_enum, plain}) == 1 + assert hash(loc) == hash(SourceLocation("k.py", 3, 5)) + # The verdicts hold a config dict: explicitly unhashable, not a hash + # that fails only for some field values. + for verdict in (ConfigVerdict("h", {}, "ok"), IRVerdict("toy_ir", "ok")): + assert type(verdict).__hash__ is None + with pytest.raises(TypeError, match="unhashable"): + hash(verdict) + + +class _Note(str, enum.Enum): + TIMEOUT = "solver timed out" + + +def test_verdict_notes_are_a_tuple_of_plain_strings(): + verdict = IRVerdict("toy_ir", "ok", notes=(n for n in ["a", _Note.TIMEOUT])) + assert verdict.notes == ("a", "solver timed out") + assert not any(isinstance(note, enum.Enum) for note in verdict.notes) + assert IRVerdict("toy_ir", "ok").notes == () + with pytest.raises(TypeError, match="notes takes a sequence, not a str"): + IRVerdict("toy_ir", "ok", notes="solver timed out") + with pytest.raises(TypeError, match="IRVerdict note must be a str, not bytes"): + IRVerdict("toy_ir", "ok", notes=[b"raw"]) # type: ignore[list-item] + with pytest.raises(TypeError, match="IRVerdict note must be a str, not int"): + IRVerdict("toy_ir", "ok", notes=b"ab") # type: ignore[arg-type] + with pytest.raises(TypeError, match="per_config items must be ConfigVerdicts"): + IRVerdict("toy_ir", "ok", per_config=[{"BLOCK": 16}]) # type: ignore[list-item] + + +# ======== save / load ========= + + +def _save_and_load(tmp_path, monkeypatch, records): + monkeypatch.setattr( + trace_module, "launches", [Launch(grid=(4, 1, 1), records=records)] + ) + path = tilelens.save(tmp_path / "trace.tvz") + with zipfile.ZipFile(path) as archive: + manifest = archive.read("manifest.json").decode() + (launch,) = tilelens.load(path) + return launch, manifest + + +def _refused_verdict() -> IRVerdict: + refusal = Refusal.from_exception( + UnsupportedTTIR( + "control-flow", "scf.while", line_no=7, loc=W.SourceLoc("k.py", 3, 5) + ) + ) + return IRVerdict( + "sanitizer_ir", + "unsupported", + scope="launch", + refusal=refusal, + per_config=[ + ConfigVerdict( + "hash-a", MappingProxyType({"BLOCK": 16, "num_warps": 4}), "proved" + ), + ConfigVerdict(None, {"BLOCK": 64}, "refused", refusal, n_reports=2), + ConfigVerdict( + "hash-c", + {"BLOCK": 32}, + "refused", + Refusal("solver-unknown", "timeout after 10 s"), + ), + ], + notes=["1 of 3 configs proved"], + ) + + +def test_a_saved_trace_holds_ir_verdicts(tmp_path, monkeypatch): + verdict = _refused_verdict() + + launch, manifest = _save_and_load(tmp_path, monkeypatch, ["report", verdict]) + + assert launch.records == ["report", verdict] + loaded = launch.records[1] + assert type(loaded) is IRVerdict and loaded is not verdict + assert type(loaded.per_config) is tuple and type(loaded.notes) is tuple + assert [type(c.config) for c in loaded.per_config] == [dict] * 3 + assert type(loaded.refusal.loc) is SourceLocation + assert not isinstance(loaded.refusal.kind, enum.Enum) + assert hash(loaded.refusal) == hash(verdict.refusal) + assert loaded.per_config[1].refusal == loaded.refusal + # The manifest names only public record types. + assert "tilelens.ir.verdict:SourceLocation" in manifest + assert "_mlir_walk" not in manifest + + +def _findings(facts: TensorFacts) -> list[CompiledSanitizerRecord]: + traceback = TracebackInfo("k.py", 3, "kernel", "x = tl.load(x_ptr + offs)") + return [ + CompiledSanitizerRecord( + kind="out-of-bounds", + op_type=Load, + tensor_name="x_ptr", + tensor_facts=facts, + witness={"pid_0": np.int64(3), "arange_0_d0": 5}, + config=MappingProxyType({"BLOCK": 64}), + user_code_tracebacks=(traceback,), + violation_offset=np.int64(-2), + violation_address=facts.data_ptr - 2 * facts.elem_size, + detail="mask is live at the witness", + ), + CompiledSanitizerRecord( + kind="integer-overflow", + op_type=Store, + tensor_name="out_ptr", + tensor_facts=facts, + witness={"pid_0": 2**20}, + config={}, + user_code_tracebacks=[traceback], + detail="pid * 4096 overflows i32", + ), + CompiledSanitizerRecord( + kind="division-by-zero", + op_type=Load, + tensor_name="x_ptr", + tensor_facts=facts, + witness={"pid_0": 0}, + config={"BLOCK": 16}, + user_code_tracebacks=[], + ), + ] + + +def test_compiled_sanitizer_records_hold_no_tensor(): + tensor = torch.arange(24, dtype=torch.float32).reshape(4, 6)[:, ::2] + tensor_ref = weakref.ref(tensor) + records = _findings(tensor_facts(tensor)) + + first = records[0] + assert first.tensor_facts.shape == (4, 3) + assert first.tensor_facts.strides == (6, 2) + assert first.tensor_facts.dtype == "torch.float32" + assert first.witness == {"pid_0": 3, "arange_0_d0": 5} + assert all(isinstance(value, int) for value in first.witness.values()) + assert isinstance(first.violation_offset, int) + assert isinstance(first.config, dict) + assert isinstance(first.user_code_tracebacks, list) + # A record outlives its launch without keeping the tensor alive. + del tensor + gc.collect() + assert tensor_ref() is None + with pytest.raises(TypeError, match="unhashable"): + hash(first) + + +@pytest.mark.parametrize( + "overrides, error, match", + [ + ({"kind": "oob"}, ValueError, "unknown compiled sanitizer finding"), + ({"op_type": object}, TypeError, "op_type must be Load or Store"), + ({"witness": {"pid_0": 1.5}}, TypeError, "float"), + ({"violation_offset": 2.0}, TypeError, "float"), + ], +) +def test_compiled_sanitizer_records_reject_values_a_trace_cannot_hold( + overrides, error, match +): + facts = tensor_facts(torch.zeros(4)) + fields = { + "kind": "out-of-bounds", + "op_type": Load, + "tensor_name": "x_ptr", + "tensor_facts": facts, + "witness": {}, + "config": {}, + "user_code_tracebacks": [], + **overrides, + } + with pytest.raises(error, match=match): + CompiledSanitizerRecord(**fields) + + +def test_a_saved_trace_holds_compiled_sanitizer_records(tmp_path, monkeypatch): + records = _findings(tensor_facts(torch.arange(24.0).reshape(4, 6)[:, ::2])) + verdict = IRVerdict( + "sanitizer_ir", + "findings", + per_config=[ + ConfigVerdict("hash-a", {"BLOCK": 64}, "findings", n_reports=len(records)) + ], + ) + + launch, manifest = _save_and_load(tmp_path, monkeypatch, [*records, verdict]) + + assert launch.records == [*records, verdict] + loaded = launch.records[:-1] + assert [type(r) for r in loaded] == [CompiledSanitizerRecord] * 3 + assert [r.op_type for r in loaded] == [Load, Store, Load] + for record in loaded: + assert type(record.tensor_facts) is TensorFacts + assert all(type(tb) is TracebackInfo for tb in record.user_code_tracebacks) + assert not any( + isinstance(getattr(record, f.name), torch.Tensor) + for f in dataclasses.fields(record) + ) + # Nothing was stored as a tensor payload. + assert '"kind": "tensor"' not in manifest + + +def test_trace_io_registers_tensor_facts_in_its_own_right(): + """TensorFacts is registered from tilelens.ir.launch, not only as a name + tilelens.clients.sanitizer.data happens to import.""" + code = ( + "import tilelens.clients.sanitizer.data as data\n" + "del data.TensorFacts\n" + "from tilelens.core import trace_io\n" + "from tilelens.ir.launch import TensorFacts\n" + "assert trace_io._TRACE_CLASSES['tilelens.ir.launch:TensorFacts'] is TensorFacts\n" + ) + subprocess.run([sys.executable, "-c", code], cwd=REPO, check=True) diff --git a/tests/unit/sanitizer_compiled/__init__.py b/tests/unit/sanitizer_compiled/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/unit/sanitizer_compiled/test_client.py b/tests/unit/sanitizer_compiled/test_client.py new file mode 100644 index 000000000..a9b81c9c9 --- /dev/null +++ b/tests/unit/sanitizer_compiled/test_client.py @@ -0,0 +1,1053 @@ +"""tilelens.clients.sanitizer.compiled.client: the CompiledSanitizer, its +factory and its reports, on fake launches. + +CPU only: fake compiled kernels hold the IR tests' TTIR (tests/unit/ir/ +ttir_corpus.py, compiled at test time) or small +TTIR texts whose locs point into a kernel source the test writes, and the +launch events are the core's own, built from CPU tensors. The real-kernel +counterparts (host-compiled, CPU too) live in +tests/end_to_end/test_compiled_sanitizer.py. +""" + +from __future__ import annotations + +import ast +import importlib +import inspect +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +import torch +import triton.language as tl + +import tilelens +import tilelens.ir +from tilelens.clients import CompiledSanitizer as ExportedCompiledSanitizer +from tilelens.clients import CompiledSanitizerRecord as ExportedRecord +from tilelens.clients import Sanitizer +from tilelens.clients.sanitizer.compiled import CompiledSanitizer, SanitizerKind +from tilelens.clients.sanitizer.data import CompiledSanitizerRecord +from tilelens.clients.sanitizer.sanitizer import NullSanitizer, SymbolicSanitizer +from tilelens.core.client import ClientManager, LaunchCall +from tilelens.core.config import config as cfg +from tilelens.core.data import Load, Store +from tilelens.ir import ConfigVerdict, IRClient, IRVerdict, ParseCache +from tilelens.ir.capture import CompiledArtifacts, CompiledSpecialization +from tilelens.ir.launch import tensor_facts +from tilelens.ir.ttir_reader import TTIRKind, UnsupportedTTIR, parse_ttir +from tilelens.ir.verdict import SourceLocation + +from ..ir import ttir_corpus + +client_module = importlib.import_module("tilelens.clients.sanitizer.compiled.client") +trace_module = importlib.import_module("tilelens.core.trace") + +ADD_TTIR = ttir_corpus.text("ttir/golden_add_sm80.ttir") +GATHER_TTIR = ttir_corpus.text("ttir/golden_gather_sm80.ttir") + + +# ======== fakes ========= + + +class _FakeJit: + """What the core reads of a JITFunction to bind a launch.""" + + def __init__(self, fn, constexprs=()): + self.signature = inspect.signature(fn) + self.params = [ + SimpleNamespace(name=name, is_constexpr=name in constexprs) + for name in self.signature.parameters + ] + + +def _add_kernel(x_ptr, y_ptr, out_ptr, n_elements, BLOCK_SIZE): + pass + + +def _gather_kernel(idx_ptr, src_ptr, out_ptr, n_elements, BLOCK_SIZE): + pass + + +def _div_kernel(p_ptr, d): + pass + + +ADD = _FakeJit(_add_kernel, {"BLOCK_SIZE"}) +GATHER = _FakeJit(_gather_kernel, {"BLOCK_SIZE"}) +DIV = _FakeJit(_div_kernel) + + +class _FakeKernel: + """A CompiledKernel as the artifact log reads it.""" + + def __init__(self, key, asm): + self.hash = f"hash-{key}" + self.asm = asm + self.metadata = SimpleNamespace( + target=SimpleNamespace(backend="cuda", arch=89), + num_warps=4, + num_stages=3, + shared=0, + name=f"kernel_{key}", + ) + + +def _compiled(jit, args, kwargs=None, *, ttir, key="a", grid=(1,)): + """A before_launch event: ``jit`` compiled (to ``ttir``) for this call.""" + asm = {} if ttir is None else {"ttir": ttir} + return ClientManager._launch_event( + jit, args, dict(kwargs or {}), grid, _FakeKernel(key, asm) + ) + + +def _failed(jit, args, kwargs, error, *, target=None): + """A compile_failed event.""" + return ClientManager._launch_event( + jit, args, dict(kwargs), (1,), None, error=error, target=target + ) + + +def _call(jit, *, capture=True): + return LaunchCall(jit_fn=jit, args=(), kwargs={}, grid=None, capture=capture) + + +def _launch(san, jit=ADD, *, compiled=(), failures=(), capture=True, peers=()): + """One traced launch through a ClientManager, as the core drives it.""" + manager = ClientManager([san, *peers]) + manager.begin_launch(_call(jit, capture=capture)) + for event in compiled: + manager._dispatch_ir("before_launch", event, manager.ir_clients()) + for event in failures: + manager._dispatch_ir("compile_failed", event, manager.ir_clients()) + manager.finalize() + return manager.launch + + +def _add_args(numel=4096, n=4096): + x, y, out = (torch.zeros(numel) for _ in range(3)) + return (x, y, out, n) + + +def _add_event(numel=4096, n=4096, *, blocks=4, key="a", kwargs=None): + """golden add (BLOCK_SIZE 1024 folded, masked by n) over ``blocks`` + programs.""" + kwargs = {"BLOCK_SIZE": 1024} if kwargs is None else kwargs + args = _add_args(numel, n) + return _compiled(ADD, args, kwargs, ttir=ADD_TTIR, key=key, grid=(blocks,)) + + +def _oob_event(**kwargs): + """golden add with every lane of 5 * 1024 active over 4096 elements.""" + return _add_event(n=10**6, blocks=5, **kwargs) + + +def _gather_event(key="g"): + args = (torch.zeros(64, dtype=torch.int32), torch.zeros(64), torch.zeros(64), 64) + return _compiled(GATHER, args, {"BLOCK_SIZE": 1024}, ttir=GATHER_TTIR, key=key) + + +# The divide kernel: a source file whose lines the TTIR locs point at. +DIV_SOURCE = """\ +@triton.jit +def div_kernel(p_ptr, d): + pid = tl.program_id(0) + q = pid // d + tl.load(p_ptr + q) +""" + + +def _div_ttir(path: Path) -> str: + path.write_text(DIV_SOURCE, encoding="utf-8") + f = str(path) + return f"""module {{ + tt.func public @div_kernel(%p: !tt.ptr loc("p_ptr"("{f}":2:0)), %d: i32 loc("d"("{f}":2:0))) attributes {{noinline = false}} {{ + %pid = tt.get_program_id x : i32 loc("{f}":3:10) + %q = arith.divsi %pid, %d : i32 loc("{f}":4:8) + %a = tt.addptr %p, %q : !tt.ptr, i32 loc("{f}":5:12) + %v = tt.load %a : !tt.ptr loc("{f}":5:4) + tt.return loc("{f}":5:4) + }} loc("{f}":2:0) +}} loc("{f}":2:0) +""" + + +def _div_event(ttir, d, *, numel=4, blocks=4): + return _compiled( + DIV, (torch.zeros(numel, dtype=torch.int32), d), ttir=ttir, grid=(blocks,) + ) + + +class _Peer(IRClient): + """Another IR client in the same trace.""" + + NAME = "peer" + IR_STAGES = frozenset() + LAUNCH = "skip" + + def __init__(self): + super().__init__() + self.finalized = 0 + + def analyze_launch(self, log): + self.finalized += 1 + return [], IRVerdict(self.NAME, "ok") + + def on_analysis_error(self, exc): + raise AssertionError(exc) + + +@pytest.fixture +def _isolate_sanitizer_cfg(): + saved = cfg.enable_sanitizer + yield + cfg.enable_sanitizer = saved + + +@pytest.fixture +def quiet(): + return CompiledSanitizer(abort_on_error=False) + + +# ======== the factory and the declarations ========= + + +def test_factory_dispatches_on_compile(_isolate_sanitizer_cfg): + cfg.enable_sanitizer = True + compiled = Sanitizer(compile=True) + assert type(compiled) is CompiledSanitizer + # A virtual subclass: a sanitizer mode, not an eager one. + assert isinstance(compiled, Sanitizer) + assert not issubclass(CompiledSanitizer, SymbolicSanitizer) + assert (compiled.abort_on_error, compiled.timeout_ms) == (True, 10_000) + tuned = Sanitizer(compile=True, abort_on_error=False, timeout_ms=50) + assert (tuned.abort_on_error, tuned.timeout_ms) == (False, 50) + + for eager in (Sanitizer(), Sanitizer(compile=False, abort_on_error=False)): + assert type(eager) is SymbolicSanitizer + assert Sanitizer(compile=False, abort_on_error=False).abort_on_error is False + # ``compile`` is a keyword, and the eager class is never the compiled one. + with pytest.raises(TypeError): + Sanitizer(True, True) + with pytest.raises(TypeError, match="Sanitizer\\(compile=True\\)"): + SymbolicSanitizer(compile=True) + for timeout_ms in (0, -1, 1.5, True): + with pytest.raises(ValueError, match="positive int"): + CompiledSanitizer(timeout_ms=timeout_ms) + + +def test_factory_initializes_the_eager_sanitizer_once( + _isolate_sanitizer_cfg, monkeypatch +): + cfg.enable_sanitizer = True + calls = [] + original = SymbolicSanitizer.__init__ + + def counting(self, *args, **kwargs): + calls.append(kwargs) + original(self, *args, **kwargs) + + monkeypatch.setattr(SymbolicSanitizer, "__init__", counting) + Sanitizer(compile=False, abort_on_error=False) + assert calls == [{"compile": False, "abort_on_error": False}] + + +def test_the_disable_flag_wins_over_compile(_isolate_sanitizer_cfg): + cfg.enable_sanitizer = False + off = Sanitizer(compile=True, abort_on_error=False) + assert type(off) is NullSanitizer + # trace() leaves a kernel traced with it untraced. + kernel = MagicMock() + assert tilelens.trace(off)(kernel) is kernel + # So it does with a CompiledSanitizer built directly, like an explicit + # SymbolicSanitizer(): the kernel then runs, and is not silently skipped. + assert tilelens.trace(CompiledSanitizer(abort_on_error=False))(kernel) is kernel + assert tilelens.trace(SymbolicSanitizer(abort_on_error=False))(kernel) is kernel + + +def test_declarations_and_composition(): + san = CompiledSanitizer() + assert san.NAME == "compiled_sanitizer" != SymbolicSanitizer.NAME + assert (san.IR_STAGES, san.LAUNCH, san.NEEDS_INTERPRETER) == ( + frozenset({"ttir"}), + "skip", + False, + ) + assert ExportedCompiledSanitizer is CompiledSanitizer + assert ExportedRecord is CompiledSanitizerRecord + assert tilelens.ir.SourceLocation is SourceLocation + # The eager sanitizer and another IR client can share its trace. + manager = ClientManager([san, SymbolicSanitizer(abort_on_error=False), _Peer()]) + assert set(manager.clients) == {"compiled_sanitizer", "sanitizer", "peer"} + + +def test_the_status_is_none_until_a_launch_is_finalized(quiet): + assert (quiet.last_status, quiet.last_verdict, quiet.records) == (None, None, []) + _launch(quiet, compiled=[_oob_event()]) + assert quiet.last_status == "violations" + manager = ClientManager([quiet]) + manager.begin_launch(_call(ADD)) + assert (quiet.last_status, quiet.last_verdict, quiet.records) == (None, None, []) + + +# ======== proofs and findings ========= + + +def test_an_in_bounds_launch_is_ok(capsys): + san = CompiledSanitizer() # abort_on_error: nothing to report + launch = _launch(san, compiled=[_add_event()]) + + verdict = san.last_verdict + assert san.last_status == "ok" and san.records == [] + assert launch.records == [verdict] + assert verdict == IRVerdict( + "compiled_sanitizer", + "ok", + # the proof holds for this launch's arguments, grid and tensors + scope="launch", + per_config=(ConfigVerdict("hash-a", {"BLOCK_SIZE": 1024}, "ok"),), + ) + assert capsys.readouterr().out == "" + + +def test_findings_become_records_after_which_the_verdict_follows(quiet, capsys): + event = _oob_event() + launch = _launch(quiet, compiled=[event]) + + records = quiet.records + assert launch.records == [*records, quiet.last_verdict] + assert [(r.kind, r.op_type, r.tensor_name) for r in records] == [ + ("out-of-bounds", Load, "x_ptr"), + ("out-of-bounds", Load, "y_ptr"), + ("out-of-bounds", Store, "out_ptr"), + ] + accesses = parse_ttir(ADD_TTIR).accesses + for record, access, tensor in zip(records, accesses, event.bound_args.values()): + assert record.tensor_facts == tensor_facts(tensor) + assert record.config == {"BLOCK_SIZE": 1024} + offset = record.violation_offset + assert 4096 <= offset < 5120 + assert record.violation_address == tensor.data_ptr() + offset * 4 + assert record.witness["pid_0"] == 4 + (lane,) = [v for k, v in record.witness.items() if k.startswith("arange_")] + assert 4 * 1024 + lane == offset + (tb,) = record.user_code_tracebacks + assert (tb.filename, tb.lineno, tb.func_name) == ( + access.loc.file, + access.loc.line, + "add_kernel", + ) + assert f"{offset}" in record.detail + (config,) = quiet.last_verdict.per_config + assert (config.status, config.n_reports, config.refusal) == ("violations", 3, None) + # Printed only with abort_on_error or TILELENS_VERBOSE. + assert capsys.readouterr().out == "" + + +def test_a_division_by_zero_record_points_at_the_division(quiet, tmp_path): + ttir = _div_ttir(tmp_path / "k.py") + _launch(quiet, DIV, compiled=[_div_event(ttir, 0)]) + + (record,) = quiet.records + assert (record.kind, record.op_type, record.tensor_name) == ( + "division-by-zero", + Load, + "p_ptr", + ) + assert (record.violation_offset, record.violation_address) == (None, None) + assert set(record.witness) == {"pid_0", "pid_1", "pid_2"} + (tb,) = record.user_code_tracebacks + assert (tb.lineno, tb.func_name, tb.line_of_code.strip()) == ( + 4, + "div_kernel", + "q = pid // d", + ) + assert "divisor" in record.detail + + # d = 1: no division by zero, but pid 1..3 read past the one element. + _launch(quiet, DIV, compiled=[_div_event(ttir, 1, numel=1)]) + (record,) = quiet.records + assert (record.kind, record.violation_offset) == ("out-of-bounds", 1) + assert record.user_code_tracebacks[0].line_of_code.strip() == "tl.load(p_ptr + q)" + + +def test_a_finding_without_a_source_location_names_its_ttir_line(quiet): + text = ( + "module {\n tt.func public @k(%p: !tt.ptr, %d: i32) " + "attributes {noinline = false} {\n" + " %pid = tt.get_program_id x : i32\n" + " %q = arith.divsi %pid, %d : i32\n" + " %a = tt.addptr %p, %q : !tt.ptr, i32\n" + " %v = tt.load %a : !tt.ptr\n" + " tt.return\n }\n}\n" + ) + nameless = _FakeJit(lambda arg0, arg1: None) + event = _compiled(nameless, (torch.zeros(4, dtype=torch.int32), 0), ttir=text) + _launch(quiet, nameless, compiled=[event]) + + (record,) = quiet.records + assert record.user_code_tracebacks == [] + assert record.detail.endswith("(TTIR line 4)") + + +def test_a_loop_finding_on_an_unbound_pointer_has_no_tensor_facts(quiet): + # The loop's bound divides by zero before its store reads anything; the + # store's pointer is bound to no tensor (e.g. a tuple argument's part). + text = ( + "module {\n tt.func public @k(%p: !tt.ptr, %d: i32) " + "attributes {noinline = false} {\n" + " %c0 = arith.constant 0 : i32\n" + " %c1 = arith.constant 1 : i32\n" + " %c8 = arith.constant 8 : i32\n" + " %u = arith.divsi %c8, %d : i32\n" + " scf.for %i = %c0 to %u step %c1 : i32 {\n" + " %a = tt.addptr %p, %i : !tt.ptr, i32\n" + " tt.store %a, %c0 : !tt.ptr\n" + " }\n" + " tt.return\n }\n}\n" + ) + nameless = _FakeJit(lambda arg0, arg1: None) + _launch(quiet, nameless, compiled=[_compiled(nameless, (None, 0), ttir=text)]) + + (record,) = quiet.records + assert (record.kind, record.op_type, record.tensor_name) == ( + "division-by-zero", + Store, + "arg0", + ) + assert record.tensor_facts is None + # The store itself could not be checked. + assert quiet.last_verdict.refusal.kind == "missing-binding" + assert quiet.last_status == "violations" + + +# ======== configs, refusals and the union ========= + + +def test_the_status_is_the_union_over_configs(quiet): + _launch( + quiet, + compiled=[ + _add_event(key="a", kwargs={"BLOCK_SIZE": 1024, "num_warps": 4}), + _oob_event(key="b", kwargs={"BLOCK_SIZE": 1024, "num_warps": 8}), + ], + ) + verdict = quiet.last_verdict + assert verdict.status == "violations" and verdict.refusal is None + assert [(c.specialization, c.status, c.n_reports) for c in verdict.per_config] == [ + ("hash-a", "ok", 0), + ("hash-b", "violations", 3), + ] + assert {r.config["num_warps"] for r in quiet.records} == {8} + + # An unsupported config makes an otherwise clean launch unsupported, and + # its refusal is the verdict's; a finding elsewhere still wins. + _launch(quiet, compiled=[_add_event(), _gather_event()]) + assert quiet.last_status == "unsupported" + assert quiet.last_verdict.refusal == quiet.last_verdict.per_config[1].refusal + _launch(quiet, compiled=[_gather_event(), _oob_event()]) + verdict = quiet.last_verdict + assert [c.status for c in verdict.per_config] == ["unsupported", "violations"] + assert verdict.status == "violations" + # Not all there is: part of the launch was not checked. + assert verdict.refusal.kind == "indirect-address" + + +def test_configs_sharing_one_kernel_are_checked_and_reported_apart(quiet): + """Two configs that compile to one kernel ("hash-a") but set a + runtime int kwarg differently each get their own verdict, checked + against their own binding; findings carry their config.""" + x, y, out, _ = _add_args() + + def config_event(n, blocks): + kwargs = {"BLOCK_SIZE": 1024, "n_elements": n} + return _compiled(ADD, (x, y, out), kwargs, ttir=ADD_TTIR, grid=(blocks,)) + + _launch(quiet, compiled=[config_event(4096, 4), config_event(10**6, 5)]) + + verdict = quiet.last_verdict + assert verdict.status == "violations" + assert [ + (c.specialization, c.config["n_elements"], c.status, c.n_reports) + for c in verdict.per_config + ] == [("hash-a", 4096, "ok", 0), ("hash-a", 10**6, "violations", 3)] + assert {r.config["n_elements"] for r in quiet.records} == {10**6} + + # A kernel the reader refuses leaves every config sharing it unchecked. + args = (torch.zeros(64, dtype=torch.int32), torch.zeros(64), torch.zeros(64)) + events = [ + _compiled(GATHER, args, {"BLOCK_SIZE": 1024, "n_elements": n}, ttir=GATHER_TTIR) + for n in (32, 64) + ] + _launch(quiet, GATHER, compiled=events) + first, second = quiet.last_verdict.per_config + assert [c.config["n_elements"] for c in (first, second)] == [32, 64] + assert first.status == second.status == "unsupported" + assert first.refusal == second.refusal and first.refusal.kind == "indirect-address" + + +def test_configs_are_told_apart_by_their_values_not_their_reprs(quiet): + """Bindings of one kernel group into configs by value: a tensor config + kwarg (e.g. a heuristic's view) by its data_ptr, shape, strides and + dtype, and reported by those facts, never its data; other values that + print alike stay apart unless they are equal.""" + + def event(value, *, oob): + args = _add_args(n=10**6 if oob else 4096) + kwargs = {"BLOCK_SIZE": 1024, "V": value} + return _compiled(ADD, args, kwargs, ttir=ADD_TTIR, grid=(5 if oob else 4,)) + + def per_config(first, second): + _launch(quiet, compiled=[event(first, oob=False), event(second, oob=True)]) + return [(c.status, c.n_reports) for c in quiet.last_verdict.per_config] + + apart = [("ok", 0), ("violations", 3)] + together = [("violations", 3)] + view = torch.zeros(4) + assert per_config(view, torch.zeros(4)) == apart # one repr, two tensors + assert per_config(_Opaque(), _Opaque()) == apart # one repr, two objects + assert per_config(view, view) == together + assert per_config(tl.dtype("fp32"), tl.dtype("fp32")) == together # equal + + per_config(view, view) + (config,) = quiet.last_verdict.per_config + assert config.config["V"] == ( + f"" + ) + assert {r.config["V"] for r in quiet.records} == {config.config["V"]} + + +def test_a_kernel_delivered_without_a_binding_is_never_ok(quiet): + """Nothing was checked, so nothing is "ok": ArtifactLog delivers every + kernel with a binding, but another producer might not.""" + artifacts = CompiledArtifacts(stages={"ttir": ADD_TTIR}, meta={"config": {}}) + spec = CompiledSpecialization("hash-x", artifacts, bindings=()) + ((records, verdict),) = quiet._check_specialization(spec) + assert records == [] and verdict.n_reports == 0 + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "internal-error") + + +def test_a_reader_refusal_keeps_its_kind_and_loc_on_cache_hits(quiet): + texts = [] + + def reader(text): + texts.append(text) + return parse_ttir(text) + + quiet.parses = ParseCache(reader, refusal=UnsupportedTTIR) + _launch(quiet, GATHER, compiled=[_gather_event()]) + first = quiet.last_verdict + _launch(quiet, GATHER, compiled=[_gather_event()]) + + assert texts == [GATHER_TTIR] # read once, across launches + assert quiet.last_verdict == first + assert (first.status, quiet.records) == ("unsupported", []) + refusal = first.refusal + assert refusal == first.per_config[0].refusal + # The reader's enum kind, held as its plain string. + assert refusal.kind == "indirect-address" and not isinstance(refusal.kind, TTIRKind) + assert isinstance(refusal.loc, SourceLocation) and refusal.line_no is not None + assert refusal.message.startswith(f"{refusal.loc.file}:{refusal.loc.line}: ") + + +def test_an_abstention_is_unsupported_with_the_sanitizers_kind(quiet): + # A float n_elements has no binding: the mask reading it is unknown. + x, y, out, _ = _add_args() + event = _compiled(ADD, (x, y, out, 4096.0), {"BLOCK_SIZE": 1024}, ttir=ADD_TTIR) + _launch(quiet, compiled=[event]) + + verdict = quiet.last_verdict + assert verdict.status == "unsupported" and quiet.records == [] + assert verdict.refusal.kind == SanitizerKind.MISSING_BINDING == "missing-binding" + assert "n_elements" in verdict.refusal.message + + +def _static_assert_failure(): + from triton.compiler.errors import CompileTimeAssertionFailure + + return CompileTimeAssertionFailure(None, ast.Pass(), "BLOCK_SIZE <= 1024") + + +def _wrapped(cause): + """``cause`` as Triton's code generator re-raises an error of a called + @jit helper or builtin: a CompilationError raised from it.""" + from triton.compiler.errors import CompilationError + + wrapped = CompilationError(None, ast.Pass(), None) + wrapped.__cause__ = cause + return wrapped + + +def test_a_compile_failure_the_target_may_explain_is_unsupported(quiet): + """A config that failed to compile for the IR target may compile for + (and launch on) the user's GPU, unchecked: unsupported, kind + compile-failed, naming the target and how to name another; never "ok", + never a note.""" + from triton.backends.compiler import GPUTarget + + args = _add_args() + cuda89 = GPUTarget("cuda", 89, 32) + failures = [ + _failed( + ADD, + args, + {"BLOCK_SIZE": 2048}, + ValueError("num_ctas > 1 requires NVIDIA SM90+ (Hopper)"), + target=cuda89, + ), + # An fp8 type the target lacks: the code generator wraps it. + _failed( + ADD, + args, + {"BLOCK_SIZE": 4096}, + _wrapped(ValueError("type fp8e4nv not supported in this architecture")), + target=cuda89, + ), + # An event built outside the core names no target. + _failed(ADD, args, {"BLOCK_SIZE": 512}, RuntimeError("bad option")), + ] + _launch(quiet, compiled=[_add_event()], failures=failures) + + verdict = quiet.last_verdict + assert (verdict.status, verdict.scope, verdict.notes) == ("unsupported", None, ()) + ok, *failed = verdict.per_config + assert ok.status == "ok" + assert [(c.specialization, c.config, c.status, c.refusal.kind) for c in failed] == [ + (None, {"BLOCK_SIZE": 2048}, "unsupported", "compile-failed"), + (None, {"BLOCK_SIZE": 4096}, "unsupported", "compile-failed"), + (None, {"BLOCK_SIZE": 512}, "unsupported", "compile-failed"), + ] + assert verdict.refusal == failed[0].refusal + assert failed[0].refusal.message == ( + "it failed to compile for cuda:89 (ValueError: num_ctas > 1 requires " + "NVIDIA SM90+ (Hopper)), so it was not checked; a kernel can compile for " + "one target and fail for another, so it may launch on a GPU of another " + "kind: to check it, name a target it compiles for " + "(Sanitizer(compile=True, target=...), or TILELENS_IR_TARGET)" + ) + # The innermost error, not the wrapper's source excerpt. + assert failed[1].refusal.message.startswith( + "it failed to compile for cuda:89 (ValueError: type fp8e4nv not supported " + "in this architecture), so it was not checked;" + ) + assert failed[2].refusal.message.startswith( + "it failed to compile (RuntimeError: bad option), so it was not checked;" + ) + + +def test_a_failure_no_target_compiles_past_is_a_note(quiet): + """A failing tl.static_assert (also in a called helper) or a construct + Triton never compiles fails for every target: the config never launches + anywhere, so it is only noted, and the launch can be "ok". Not so once + the compile had asked for its target (e.g. tl.target_info), or for any + other error.""" + from triton.backends.compiler import GPUTarget + from triton.compiler.errors import UnsupportedLanguageConstruct + + from tilelens.core import host_compile + + cuda89 = GPUTarget("cuda", 89, 32) + args = _add_args() + nowhere = [ + _static_assert_failure(), + _wrapped(_static_assert_failure()), + UnsupportedLanguageConstruct(None, ast.Pass(), "nested function"), + ] + failures = [ + _failed(ADD, args, {"BLOCK_SIZE": 2048 * (i + 1)}, error, target=cuda89) + for i, error in enumerate(nowhere) + ] + _launch(quiet, compiled=[_add_event()], failures=failures) + + verdict = quiet.last_verdict + assert verdict.status == "ok" and len(verdict.per_config) == 1 + assert len(verdict.notes) == 3 + assert verdict.notes[0] == ( + "config {'BLOCK_SIZE': 2048} was not checked: it failed to compile for " + "cuda:89 (CompileTimeAssertionFailure: BLOCK_SIZE <= 1024), an error of " + "its own code whatever the target, so it never launches" + ) + # Through the helper's wrapper: the assertion's own message. + assert "(CompileTimeAssertionFailure: BLOCK_SIZE <= 1024)" in verdict.notes[1] + assert "(UnsupportedLanguageConstruct: nested function)" in verdict.notes[2] + + # The same errors after a target query, or wrapped around another error, + # may be the target's. + queried = _static_assert_failure() + host_compile._mark_target_queried(queried) + maybe = [queried, _wrapped(ValueError("x")), None] + failures = [ + _failed(ADD, args, {"BLOCK_SIZE": 2048 * (i + 1)}, error, target=cuda89) + for i, error in enumerate(maybe) + ] + _launch(quiet, compiled=[_add_event()], failures=failures) + verdict = quiet.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + assert [c.refusal.kind for c in verdict.per_config[1:]] == ["compile-failed"] * 3 + + +def test_a_host_compile_that_could_not_run_is_unsupported_not_a_note(quiet): + """HostCompileUnavailable (Triton's compile API, or a compile that asked + for a device) says nothing about the kernel: the config may launch, so + it is unsupported, never a note that it cannot launch there.""" + from triton.backends.compiler import GPUTarget + from triton.compiler.errors import CompilationError + + from tilelens.core.host_compile import HostCompileUnavailable + + cuda80 = GPUTarget("cuda", 80, 32) + unavailable = HostCompileUnavailable("Triton asked its driver for 'utils'") + # Triton's code generator re-raises what the kernel's code raised. + wrapped = CompilationError("def add(...)", None, repr(unavailable)) + wrapped.__cause__ = unavailable + failures = [ + _failed(ADD, _add_args(), {"BLOCK_SIZE": 2048}, unavailable, target=cuda80), + _failed(ADD, _add_args(), {"BLOCK_SIZE": 4096}, wrapped, target=cuda80), + ] + _launch(quiet, compiled=[_add_event()], failures=failures) + verdict = quiet.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + ok, *refused = verdict.per_config + assert ok.status == "ok" + assert [(c.config, c.status, c.refusal.kind) for c in refused] == [ + ({"BLOCK_SIZE": 2048}, "unsupported", "host-compile-unavailable"), + ({"BLOCK_SIZE": 4096}, "unsupported", "host-compile-unavailable"), + ] + assert all("asked its driver for 'utils'" in c.refusal.message for c in refused) + + # Nothing compiled: the launch's refusal says why. + _launch(quiet, failures=failures[:1]) + verdict = quiet.last_verdict + assert (verdict.status, verdict.refusal.kind) == ( + "unsupported", + SanitizerKind.HOST_COMPILE_UNAVAILABLE, + ) + assert len(verdict.per_config) == 1 and verdict.notes == () + + +def test_a_launch_that_compiled_nothing_is_unsupported(quiet): + # No JITFunction (TRITON_INTERPRET, Gluon, NKI): nothing was captured. + _launch(quiet, capture=False) + verdict = quiet.last_verdict + assert (verdict.status, verdict.refusal.kind) == ( + "unsupported", + "no-compiled-kernel", + ) + assert "TRITON_INTERPRET" in verdict.refusal.message + + # Every config failed to compile for the target (a mixed trace goes + # on to interpret the launch): the first one's refusal. + from triton.backends.compiler import GPUTarget + + cuda89 = GPUTarget("cuda", 89, 32) + failure = _failed( + ADD, _add_args(), {"BLOCK_SIZE": 2048}, RuntimeError("bad"), target=cuda89 + ) + _launch(quiet, failures=[failure]) + verdict = quiet.last_verdict + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "compile-failed") + (config,) = verdict.per_config + assert config.refusal == verdict.refusal and verdict.notes == () + assert "for cuda:89 (RuntimeError: bad)" in verdict.refusal.message + assert "TILELENS_IR_TARGET" in verdict.refusal.message + + # Every config failed with an error no target compiles past: compile-failed + # too, the notes saying why. + failures = [ + _failed(ADD, _add_args(), {"BLOCK_SIZE": b}, _static_assert_failure(), target=t) + for b, t in ((2048, cuda89), (4096, None)) + ] + _launch(quiet, failures=failures) + verdict = quiet.last_verdict + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "compile-failed") + assert verdict.per_config == () and len(verdict.notes) == 2 + assert verdict.refusal.message == ( + "no config of the launch compiled for cuda:89, so nothing was checked: " + "each failed with an error of its own code whatever the target (see the " + "notes)" + ) + + +def test_errors_no_target_decides_name_no_target(quiet): + """A call that does not bind the kernel's parameters (the host compile + marks it, see bind_failed) and a compile refused while the language is + patched (no compile ran) are unsupported, never notes, but their + refusals send nobody to another target: the untraced call raises the + bind error on any GPU, and the refused compile says nothing about the + kernel. (A traced launch raises a bind failure instead; only an + event built outside the core, as here, delivers one.)""" + from triton.backends.compiler import GPUTarget + + from tilelens.core import host_compile + from tilelens.core.client import LanguagePatchedError + + cuda89 = GPUTarget("cuda", 89, 32) + unbound = TypeError("dynamic_func() missing 1 required positional argument: 'n'") + host_compile._mark_bind_failed(unbound) + patched = LanguagePatchedError("a Triton compile cannot run while ...") + failures = [ + _failed(ADD, _add_args(), {"BLOCK_SIZE": b}, error, target=cuda89) + for b, error in ((2048, unbound), (4096, patched)) + ] + _launch(quiet, compiled=[_add_event()], failures=failures) + + verdict = quiet.last_verdict + assert verdict.status == "unsupported" and verdict.notes == () + ok, bind, refused = verdict.per_config + assert ok.status == "ok" + assert (bind.status, bind.refusal.kind) == ("unsupported", "compile-failed") + assert bind.refusal.message == ( + "the call does not bind to the kernel's parameters (TypeError: " + "dynamic_func() missing 1 required positional argument: 'n'), so it was " + "not checked; Triton raises this error for the call whatever the GPU" + ) + assert (refused.status, refused.refusal.kind) == ( + "unsupported", + "host-compile-unavailable", + ) + assert refused.refusal.message == "a Triton compile cannot run while ..." + for config in (bind, refused): + assert "TILELENS_IR_TARGET" not in config.refusal.message + + +def test_a_launch_no_config_of_which_compiled_prints_its_notes(capsys): + """The refusal of a launch every config of which failed only as a note + says "see the notes": they are printed with it, each once.""" + from triton.backends.compiler import GPUTarget + + cuda89 = GPUTarget("cuda", 89, 32) + det = CompiledSanitizer() # abort_on_error: prints what was not checked + failures = [ + _failed( + ADD, _add_args(), {"BLOCK_SIZE": b}, _static_assert_failure(), target=cuda89 + ) + for b in (2048, 4096) + ] + for _ in range(2): + _launch(det, failures=failures) + assert det.last_verdict.per_config == () and len(det.last_verdict.notes) == 2 + + lines = capsys.readouterr().out.splitlines() + assert lines == [ + "[CompiledSanitizer] not checked: compile-failed: no config of the launch " + "compiled for cuda:89, so nothing was checked: each failed with an error " + "of its own code whatever the target (see the notes)", + *(f"[CompiledSanitizer] note: {note}" for note in det.last_verdict.notes), + ] + assert "(CompileTimeAssertionFailure: BLOCK_SIZE <= 1024)" in lines[1] + + +def test_a_kernel_without_ttir_is_unsupported(quiet): + _launch(quiet, compiled=[_compiled(ADD, _add_args(), ttir=None)]) + (config,) = quiet.last_verdict.per_config + assert (config.status, config.refusal.kind) == ("unsupported", "no-ttir") + + +def test_analysis_bugs_are_contained_per_config(quiet, monkeypatch): + def reader(text): + if "gather_kernel" in text: + raise KeyError("reader bug") + return parse_ttir(text) + + quiet.parses = ParseCache(reader, refusal=UnsupportedTTIR) + real_check = client_module.check_graph + + def check_graph(graph, binding, **kw): + if binding.grid == (2, 1, 1): + raise ZeroDivisionError("evaluator bug") + return real_check(graph, binding, **kw) + + monkeypatch.setattr(client_module, "check_graph", check_graph) + _launch( + quiet, + compiled=[_gather_event(), _add_event(blocks=2, key="b"), _oob_event()], + ) + reader_bug, evaluator_bug, checked = quiet.last_verdict.per_config + assert (reader_bug.refusal.kind, evaluator_bug.refusal.kind) == ( + "internal-error", + "internal-error", + ) + assert "KeyError" in reader_bug.refusal.message + assert evaluator_bug.refusal.message == "ZeroDivisionError: evaluator bug" + # The other config's findings still count. + assert checked.status == "violations" and len(quiet.records) == 3 + assert quiet.last_status == "violations" + + +def test_a_bug_outside_the_configs_is_an_internal_error(quiet, monkeypatch): + def broken(failure): + raise RuntimeError("note bug") + + monkeypatch.setattr(client_module, "_failure_refusal", broken) + failure = _failed(ADD, _add_args(), {"BLOCK_SIZE": 2048}, RuntimeError("bad")) + launch = _launch(quiet, compiled=[_oob_event()], failures=[failure]) + + verdict = quiet.last_verdict + assert (verdict.status, verdict.refusal.kind) == ("unsupported", "internal-error") + assert verdict.refusal.message == "RuntimeError: note bug" + assert quiet.records == [] and launch.records == [verdict] + + +def test_a_repeated_launch_reuses_its_check(quiet, monkeypatch): + """A training loop: fresh tensors of the same shapes each launch. The + check runs once; each launch's records hold its own tensors' addresses.""" + calls = [] + real = client_module.check_graph + + def counting(graph, binding, **kwargs): + calls.append(binding) + return real(graph, binding, **kwargs) + + monkeypatch.setattr(client_module, "check_graph", counting) + addresses, events = [], [] # the events hold the tensors: fresh addresses + for _ in range(3): + event = _oob_event() + events.append(event) + _launch(quiet, compiled=[event]) + x = event.bound_args["x_ptr"] + (record,) = [r for r in quiet.records if r.tensor_name == "x_ptr"] + assert record.violation_address == x.data_ptr() + record.violation_offset * 4 + assert record.tensor_facts == tensor_facts(x) + addresses.append(record.violation_address) + assert len(calls) == 1 and len(set(addresses)) == 3 + assert quiet.last_status == "violations" + # Another argument (or shape) is another check. + _launch(quiet, compiled=[_add_event()]) + _launch(quiet, compiled=[_add_event(numel=8192, n=8192, blocks=8)]) + assert len(calls) == 3 and quiet.last_status == "ok" + + +def test_an_interrupt_during_the_check_is_not_contained(quiet, monkeypatch): + """Ctrl+C is the user's: never an internal-error verdict.""" + + def interrupted(graph, binding, **kwargs): + raise KeyboardInterrupt + + monkeypatch.setattr(client_module, "check_graph", interrupted) + with pytest.raises(KeyboardInterrupt): + _launch(quiet, compiled=[_oob_event()]) + + +def test_records_belong_to_the_last_launch(quiet): + _launch(quiet, compiled=[_oob_event()]) + assert len(quiet.records) == 3 + _launch(quiet, compiled=[_add_event()]) + assert (quiet.last_status, quiet.records) == ("ok", []) + + +# ======== reporting and abort_on_error ========= + + +def test_abort_on_error_reports_then_exits_once_every_client_finalized(capsys): + san, peer = CompiledSanitizer(), _Peer() + with pytest.raises(SystemExit) as exc_info: + _launch(san, compiled=[_oob_event()], peers=[peer]) + + assert exc_info.value.code == 1 + assert peer.finalized == 1 + assert san.last_status == "violations" and len(san.records) == 3 + out = capsys.readouterr().out + assert out.count("Out-Of-Bounds Access Detected") == 3 + assert "Tensor Arg: x_ptr" in out and "Operation: Store" in out + assert "Witness: pid_0=4" in out and "Config: {'BLOCK_SIZE': 1024}" in out + assert f"{san.records[0].violation_address:#x}" in out + + +def test_abort_on_error_reports_but_never_exits_for_unchecked_parts(capsys): + san = CompiledSanitizer() + _launch(san, GATHER, compiled=[_gather_event()]) + assert san.last_status == "unsupported" + (line,) = capsys.readouterr().out.splitlines() + assert line.startswith( + "[CompiledSanitizer] not checked (config {'BLOCK_SIZE': 1024}): " + "indirect-address: " + ) + # Once per client: a kernel launched in a loop does not repeat it. + _launch(san, GATHER, compiled=[_gather_event()]) + assert san.last_status == "unsupported" + assert capsys.readouterr().out == "" + # A launch-level refusal is printed too. + _launch(san, capture=False) + assert "not checked: no-compiled-kernel: " in capsys.readouterr().out + + +def test_every_config_that_failed_to_compile_is_printed_once(capsys): + """Failed configs have no kernel: each is printed once per config and + message, and the launch never exits for them.""" + san = CompiledSanitizer() + failures = [ + _failed(ADD, _add_args(), {"BLOCK_SIZE": b}, ValueError("not here")) + for b in (2048, 4096) + ] + _launch(san, compiled=[_add_event()], failures=failures) + assert san.last_status == "unsupported" + lines = capsys.readouterr().out.splitlines() + assert [line.split(": compile-failed: ")[0] for line in lines] == [ + "[CompiledSanitizer] not checked (config {'BLOCK_SIZE': 2048})", + "[CompiledSanitizer] not checked (config {'BLOCK_SIZE': 4096})", + ] + _launch(san, compiled=[_add_event()], failures=failures) + assert capsys.readouterr().out == "" + + +def test_an_unchecked_op_prints_once_whatever_its_message(capsys): + """Once per compiled kernel, kind and op: a withheld finding's message + names an element offset that changes with the tensors' sizes.""" + printed: set = set() + loc = SourceLocation("k.py", 7) + + def verdict(message, kind="unmodelable-condition", line_no=3): + refusal = client_module.Refusal(kind, message, line_no, loc) + return IRVerdict( + "compiled_sanitizer", + "unsupported", + refusal=refusal, + per_config=(ConfigVerdict("hash-a", {}, "unsupported", refusal),), + ) + + client_module.print_unchecked(verdict("at element offset 4096"), printed) + client_module.print_unchecked(verdict("at element offset 8192"), printed) + assert len(capsys.readouterr().out.splitlines()) == 1 + client_module.print_unchecked(verdict("offset 1", line_no=4), printed) + client_module.print_unchecked(verdict("offset 1", "data-dependent-mask"), printed) + assert len(capsys.readouterr().out.splitlines()) == 2 + + +def test_verbose_prints_without_aborting(quiet, monkeypatch, capsys, tmp_path): + monkeypatch.setattr(cfg, "verbose", True) + _launch(quiet, DIV, compiled=[_div_event(_div_ttir(tmp_path / "k.py"), 0)]) + out = capsys.readouterr().out + assert "Division By Zero Detected" in out + assert "Code: q = pid // d" in out + assert "Invalid access detected" not in out + + +# ======== persistence ========= + + +class _Opaque: + def __repr__(self): + return "" + + +def test_a_launch_round_trips_through_a_saved_trace(quiet, tmp_path): + kwargs = {"BLOCK_SIZE": 1024, "DTYPE": _Opaque()} + launch = _launch(quiet, compiled=[_oob_event(kwargs=kwargs), _gather_event()]) + # A config value a trace cannot hold is kept as its repr. + assert quiet.records[0].config == {"BLOCK_SIZE": 1024, "DTYPE": ""} + assert quiet.last_verdict.per_config[0].config["DTYPE"] == "" + + saved = list(trace_module.launches) + trace_module.launches[:] = [launch] + try: + path = tilelens.save(tmp_path / "trace.zip") + (loaded,) = tilelens.load(path) + finally: + trace_module.launches[:] = saved + assert loaded.records == launch.records + assert all(isinstance(r, CompiledSanitizerRecord) for r in loaded.records[:-1]) + assert loaded.records[-1].per_config[1].refusal.loc == ( + launch.records[-1].per_config[1].refusal.loc + ) diff --git a/tests/unit/sanitizer_compiled/test_oob.py b/tests/unit/sanitizer_compiled/test_oob.py new file mode 100644 index 000000000..ac2cf24c6 --- /dev/null +++ b/tests/unit/sanitizer_compiled/test_oob.py @@ -0,0 +1,1233 @@ +"""tilelens.clients.sanitizer.compiled.oob: the compiled sanitizer's checks. + +CPU only: graphs come from the TTIR reader on the IR tests' corpus +(tests/unit/ir/ttir_corpus.py: ttir/ and reader_ttir/, compiled at test +time) or on small TTIR texts, or are built by hand; +launches are synthetic LaunchBindings, with TensorFacts read from CPU +tensors where a view's layout matters. Ports the #361 cases of +tests/unit/test_compiled_sanitizer_oob.py onto the new reader and binding. +""" + +from __future__ import annotations + +import pickle +import subprocess +import sys +from pathlib import Path +from types import MappingProxyType + +import pytest +import torch +import z3 + +from tilelens.clients.sanitizer.compiled import oob +from tilelens.clients.sanitizer.compiled.oob import ( + CheckResult, + Finding, + SanitizerKind, + _in_view, + check_graph, + launch_key, + readdressed, +) +from tilelens.clients.symbolic_engine import SymbolicClient +from tilelens.ir.launch import LaunchBinding, TensorFacts, tensor_facts +from tilelens.ir.ttir_reader import ( + AccessEvent, + AccessGraph, + Arange, + Bin, + BoolBin, + Cmp, + Const, + DataDep, + FuncArg, + LoopInfo, + LoopVar, + Observed, + Param, + Pid, + TTIRKind, + parse_ttir, +) +from tilelens.ir.verdict import Refusal, SourceLocation + +from ..ir import ttir_corpus + +K = SanitizerKind +REPO = Path(__file__).resolve().parents[3] +ADD = "ttir/golden_add_sm80.ttir" +MATMUL = "ttir/golden_matmul_s3_sm80.ttir" +TILE2D = "ttir/golden_tile2d_sm80.ttir" +PTR = 0x1000 + + +def _text(name: str) -> str: + return ttir_corpus.text(name) + + +def _graph(name: str) -> AccessGraph: + return parse_ttir(_text(name)) + + +def _module( + body: str, args: str = "%p: !tt.ptr, %n: i32" +) -> tuple[AccessGraph, str]: + """A minimal TTIR module (no locs, so its parameters read as arg0, + arg1, ...) around ``body``'s op lines, and its text.""" + lines = "\n ".join(line.strip() for line in body.strip().splitlines()) + text = ( + f"module {{\n tt.func public @k({args}) attributes {{noinline = false}} {{\n" + f" {lines}\n tt.return\n }}\n}}\n" + ) + return parse_ttir(text), text + + +def _line(text: str, needle: str) -> int: + (line,) = [i for i, t in enumerate(text.splitlines(), 1) if needle in t] + return line + + +def _facts(numel: int, elem_size: int = 4, data_ptr: int = PTR) -> TensorFacts: + """A contiguous 1D tensor.""" + return TensorFacts( + data_ptr, elem_size, numel, (numel,), (1,), "torch.float32", True + ) + + +def _bind(grid=(1, 1, 1), params=None, **tensors) -> LaunchBinding: + return LaunchBinding( + params=MappingProxyType(dict(params or {})), + tensors=MappingProxyType(tensors), + constexprs=MappingProxyType({}), + raw_grid=grid, + grid=grid, + config=MappingProxyType({}), + ) + + +def _add(grid, n, **tensors) -> CheckResult: + return check_graph(_graph(ADD), _bind(grid, {"n_elements": n}, **tensors)) + + +def _all(numel: int, *names: str, elem_size: int = 4) -> dict[str, TensorFacts]: + return {name: _facts(numel, elem_size) for name in names} + + +ADD_TENSORS = ("x_ptr", "y_ptr", "out_ptr") + + +def _clean(result: CheckResult) -> None: + assert result.findings == () + assert result.refusal is None and result.abstained == () + + +def _found(result: CheckResult) -> list[tuple[str, int]]: + return [(f.kind, f.access_index) for f in result.findings] + + +def _synthetic(*accesses: AccessEvent, loop: LoopInfo | None = None, args=("p",)): + return AccessGraph( + kernel_name="synthetic", + func_args=[FuncArg(a, True, 32) for a in args], + accesses=accesses, + loop=loop, + ) + + +def _access(offset, *, kind="load", base="p", mask=None, line=1, **kw) -> AccessEvent: + return AccessEvent(kind, base, offset, mask, 32, None, line, **kw) + + +# ─────────────────────────── add (1D, masked) ─────────────────────────── + + +def test_add_in_bounds_is_clean(): + _clean(_add((4, 1, 1), 4096, **_all(4096, *ADD_TENSORS))) + + +def test_add_unmasked_tail_is_out_of_bounds(): + """A mask bound (n_elements) past the tensor leaves the last block's + tail unguarded.""" + r = _add((5, 1, 1), 10**9, **_all(4096, *ADD_TENSORS)) + assert r.refusal is None and r.abstained == () + assert _found(r) == [ + ("out-of-bounds", 0), + ("out-of-bounds", 1), + ("out-of-bounds", 2), + ] + text = _text(ADD).splitlines() + for f, kind in zip(r.findings, ("load", "load", "store")): + assert f.access_kind == kind and f.base_param == ADD_TENSORS[f.access_index] + assert f.violation_offset >= 4096 + assert f.violation_address == PTR + f.violation_offset * 4 + # the witness is the state that reaches the offset + lane = next(v for k, v in f.witness.items() if k.startswith("arange_")) + assert f.witness["pid_0"] * 1024 + lane == f.violation_offset + assert f"tt.{kind}" in text[f.line_no - 1] + assert isinstance(f.loc, SourceLocation) and f.loc.line > 0 + assert "outside the tensor's 4096 elements" in f.detail + + +# ─────────────────────────── matmul (loop, 2D) ─────────────────────────── + +_MATMUL_PARAMS = { + "M": 128, + "N": 128, + "K": 128, + "stride_am": 128, + "stride_bk": 128, + "stride_cm": 128, +} + + +def test_matmul_in_bounds_and_oversized_grid(): + tensors = _all(128 * 128, "a_ptr", "b_ptr", "c_ptr", elem_size=2) + _clean(check_graph(_graph(MATMUL), _bind((2, 2, 1), _MATMUL_PARAMS, **tensors))) + # The A load has no row mask (only a K mask): too many row blocks + # (pid_m = 2 with M = 128, BLOCK_M = 64) read rows past M. + r = check_graph(_graph(MATMUL), _bind((3, 2, 1), _MATMUL_PARAMS, **tensors)) + assert ("out-of-bounds", 0) in _found(r) and r.refusal is None + a = r.findings[0] + assert a.base_param == "a_ptr" and a.witness["pid_0"] == 2 + assert 0 <= a.witness["iter_loop"] < 4 # K / BLOCK_K iterations run + assert a.violation_address == PTR + a.violation_offset * 2 + + +def test_modeled_branch_path_constrains_the_witness(): + """Only program 0 stores (``if pid == 0``): the path is modeled, so a + later program's out-of-range offsets are no witness.""" + g = _graph("ttir/golden_pid_branch_sm80.ttir") + b = _bind((4, 1, 1), {"n_elements": 1024}, x_ptr=_facts(1024), out_ptr=_facts(256)) + _clean(check_graph(g, b)) + b = _bind((4, 1, 1), {"n_elements": 1024}, x_ptr=_facts(1024), out_ptr=_facts(255)) + ((f,),) = [check_graph(g, b).findings] + assert f.access_index == 1 and f.witness["pid_0"] == 0 and f.violation_offset == 255 + + +def test_grid_stride_loop_with_a_pid_lower_bound(): + """``for row in range(pid, n_rows, 4)``: a loop bound that depends on + the program id stays symbolic.""" + g = _graph("ttir/golden_grid_stride_sm80.ttir") + params = {"n_rows": 8, "stride": 64} + _clean(check_graph(g, _bind((4, 1, 1), params, **_all(8 * 64, "x_ptr", "out_ptr")))) + r = check_graph( + g, _bind((4, 1, 1), params, x_ptr=_facts(7 * 64), out_ptr=_facts(8 * 64)) + ) + ((f,),) = [r.findings] + assert f.base_param == "x_ptr" and f.violation_offset >= 7 * 64 + # row = pid + 4 * iter reaches row 7 only + assert f.witness["pid_0"] + 4 * f.witness["iter_loop"] == 7 + + +def test_expanded_loop_carried_tiles_keep_their_lanes(): + """q[i, j, l] = x + j*N + l - i*N + k (N = 4, k the iteration): the + reader's expanded iter_arg tile, one make_range on three dims.""" + g = _graph("reader_ttir/expand_iterarg_3d.ttir") + r = check_graph( + g, _bind((1, 1, 1), {"n": 2}, x_ptr=_facts(1000), out_ptr=_facts(64)) + ) + ((f,),) = [r.findings] + assert f.access_index == 0 and f.violation_offset < 0 + lane = {int(k[-1]): v for k, v in f.witness.items() if k.startswith("arange_")} + assert sorted(lane) == [0, 1, 2] + k = f.witness["iter_loop"] + assert 0 <= k < 2 + assert lane[1] * 4 + lane[2] - lane[0] * 4 + k == f.violation_offset + + +def test_select_of_pointers(): + """``tl.where(offs < n, x + offs, x + 100)`` over 16 lanes.""" + g = _graph("reader_ttir/where_pointer.ttir") + _clean(check_graph(g, _bind(params={"n": 8}, x_ptr=_facts(101)))) + _clean(check_graph(g, _bind(params={"n": 16}, x_ptr=_facts(16)))) + ((f,),) = [check_graph(g, _bind(params={"n": 8}, x_ptr=_facts(100))).findings] + assert f.violation_offset == 100 + + +def test_three_lanes_of_one_make_range(): + g = _graph("reader_ttir/tile3d_shared_arange.ttir") + _clean(check_graph(g, _bind(x_ptr=_facts(64)))) + ((f,),) = [check_graph(g, _bind(x_ptr=_facts(63))).findings] + assert f.violation_offset == 63 + + +def test_aranges_along_one_dim_share_the_lane(): + """tl.arange(1, 65) - tl.arange(0, 64) is 1 on every lane: two + make_range sites along one dim index one lane position, so positions 1 + and 0 of the two are no witness (the real kernel reads x[1] only).""" + g, _ = _module( + """ + %r1 = tt.make_range {end = 65 : i32, start = 1 : i32} : tensor<64xi32> + %r0 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> + %d = arith.subi %r1, %r0 : tensor<64xi32> + %ps = tt.splat %p : !tt.ptr -> tensor<64x!tt.ptr> + %a = tt.addptr %ps, %d : tensor<64x!tt.ptr>, tensor<64xi32> + %v = tt.load %a : tensor<64x!tt.ptr>""" + ) + _clean(check_graph(g, _bind(arg0=_facts(2)))) + ((f,),) = [check_graph(g, _bind(arg0=_facts(1))).findings] + assert f.violation_offset == 1 + # the witness names each arange by its range, at the one lane + assert f.witness["arange_1_65"] == f.witness["arange_0_64"] + 1 + + +def test_a_broadcast_extent_one_arange_keeps_its_position(): + """tl.arange(5, 6) broadcast to 64 lanes plus tl.arange(0, 64): the + extent-1 range stays at its one position, every lane of the other + range stays free.""" + g, _ = _module( + """ + %r5 = tt.make_range {end = 6 : i32, start = 5 : i32} : tensor<1xi32> + %b = tt.broadcast %r5 : tensor<1xi32> -> tensor<64xi32> + %r0 = tt.make_range {end = 64 : i32, start = 0 : i32} : tensor<64xi32> + %o = arith.addi %b, %r0 : tensor<64xi32> + %ps = tt.splat %p : !tt.ptr -> tensor<64x!tt.ptr> + %a = tt.addptr %ps, %o : tensor<64x!tt.ptr>, tensor<64xi32> + %v = tt.load %a : tensor<64x!tt.ptr>""" + ) + _clean(check_graph(g, _bind(arg0=_facts(69)))) + ((f,),) = [check_graph(g, _bind(arg0=_facts(68))).findings] + assert f.violation_offset == 68 + assert (f.witness["arange_5_6"], f.witness["arange_0_64"]) == (5, 63) + + +# ─────────────────── synthetic graphs (ported from #361) ─────────────────── + + +def test_reused_arange_rows_and_cols_are_independent(): + """One make_range on the row and the column dim is two variables: with + offset = row - col, collapsing them makes the offset 0 (never OOB).""" + r = "%shared_range" + offset = Bin("-", Arange(r, 0, 4, dim=0), Arange(r, 0, 4, dim=1)) + ((f,),) = [check_graph(_synthetic(_access(offset)), _bind(p=_facts(64))).findings] + assert f.violation_offset < 0 + + +def _store_loop(lower, step, upper, offset): + return _synthetic( + _access(offset, kind="store", base="out", in_loop=True), + loop=LoopInfo("%loop", lower=lower, upper=upper, step=step), + args=("out",), + ) + + +def test_loop_nonzero_lower_no_false_positive(): + """for k in range(1, n): store(out + (k - 1)) writes 0..n-2: the + iteration model never runs k = 0.""" + g = _store_loop( + Const(1), Const(1), Param("n"), Bin("-", LoopVar("%loop"), Const(1)) + ) + _clean(check_graph(g, _bind(params={"n": 8}, out=_facts(7)))) + + +def test_loop_step_skips_unrun_iterations(): + """for k in range(0, n, 2): store(out + k) writes only even offsets.""" + g = _store_loop(Const(0), Const(2), Param("n"), LoopVar("%loop")) + _clean(check_graph(g, _bind(params={"n": 8}, out=_facts(7)))) + ((f,),) = [check_graph(g, _bind(params={"n": 8}, out=_facts(6))).findings] + assert f.violation_offset == 6 and f.witness["iter_loop"] == 3 + + +def test_descending_loop_refuses_non_positive_step(): + g = _store_loop(Const(10), Const(-1), Const(0), LoopVar("%loop")) + r = check_graph(g, _bind(out=_facts(16))) + assert r.findings == () and r.abstained == ((0, K.NON_POSITIVE_STEP),) + assert r.refusal.kind == "non-positive-step" and "step is -1" in r.refusal.message + + +def test_step_that_depends_on_the_program_id(): + lower, upper = Const(0), Const(8) + ok = _store_loop(lower, Bin("+", Pid(0), Const(1)), upper, LoopVar("%loop")) + _clean(check_graph(ok, _bind((4, 1, 1), out=_facts(8)))) + bad = _store_loop(lower, Bin("-", Pid(0), Const(1)), upper, LoopVar("%loop")) + r = check_graph(bad, _bind((4, 1, 1), out=_facts(8))) + assert r.abstained == ((0, K.NON_POSITIVE_STEP),) + assert "step can be" in r.refusal.message + + +def _flat_load(offset, numel, grid_x): + return check_graph( + _synthetic(_access(offset)), _bind((grid_x, 1, 1), p=_facts(numel)) + ) + + +@pytest.mark.parametrize( + "offset, numel, grid_x, oob", + [ + # pid % 8 stays in [0, 8): the grouped-swizzle `pid % group_size_m` + (Bin("%", Pid(0), Const(8)), 8, 64, None), + (Bin("%", Pid(0), Const(8)), 7, 64, 7), + # min(pid, 9) never leaves [0, 10) + (Bin("min", Pid(0), Const(9)), 10, 1000, None), + (Bin("min", Pid(0), Const(9)), 9, 1000, 9), + # max(pid, 5) pins the floor at 5 + (Bin("max", Pid(0), Const(5)), 6, 4, None), + (Bin("max", Pid(0), Const(5)), 5, 4, 5), + # divsi truncates: (0 - 1) // 2 == 0, not Euclidean -1 + (Bin("//", Bin("-", Const(0), Pid(0)), Const(2)), 4, 2, None), + # remsi keeps the dividend's sign: (0 - 1) % 2 == -1 + (Bin("%", Bin("-", Const(0), Pid(0)), Const(2)), 2, 2, -1), + ], +) +def test_integer_semantics(offset, numel, grid_x, oob): + r = _flat_load(offset, numel, grid_x) + assert r.refusal is None + assert [f.violation_offset for f in r.findings] == ([] if oob is None else [oob]) + + +# ─────────────────── uncertainty: guarded, dropped, observed ─────────────────── + + +def test_guarded_access_unsat_is_still_a_proof(): + g = _synthetic(_access(Pid(0), guarded=True)) + _clean(check_graph(g, _bind((4, 1, 1), p=_facts(4)))) + + +def test_guarded_access_sat_abstains(): + """SAT under a branch condition the reader could not model may be a + branch the launch never takes: abstain, never a witness.""" + g = _synthetic(_access(Bin("-", Pid(0), Const(1)), guarded=True)) + r = check_graph(g, _bind((4, 1, 1), p=_facts(4))) + assert r.findings == () and r.abstained == ((0, K.UNMODELABLE_CONDITION),) + assert r.refusal.kind == "unmodelable-condition" + assert "possible out-of-bounds" in r.refusal.message + assert r.refusal.message.startswith("TTIR line 1: ") + + +def test_exact_finding_kept_alongside_guarded_abstention(): + g = _synthetic( + _access(Bin("-", Pid(0), Const(1)), guarded=True), + _access(Bin("+", Pid(0), Const(100)), line=2), + ) + r = check_graph(g, _bind((4, 1, 1), p=_facts(4))) + assert _found(r) == [("out-of-bounds", 1)] and r.findings[0].line_no == 2 + assert r.findings[0].violation_offset >= 100 + assert r.abstained == ((0, K.UNMODELABLE_CONDITION),) + + +def test_mask_dropped_abstains_and_exact_findings_stay(): + """atomic_fmax's atomics are masked by loaded data (dropped as free).""" + g = _graph("ttir/golden_atomic_fmax_sm80.ttir") + assert [a.mask_dropped for a in g.accesses] == [False, True, True] + r = check_graph( + g, + _bind((4, 1, 1), {"n_elements": 1024}, x_ptr=_facts(2048), out_ptr=_facts(10)), + ) + assert r.findings == () + assert r.abstained == ((1, K.DATA_DEPENDENT_MASK), (2, K.DATA_DEPENDENT_MASK)) + assert ( + r.refusal.kind == "data-dependent-mask" + and r.refusal.line_no == g.accesses[1].line_no + ) + # the exact load's finding is kept (#361 dropped it behind the abstention) + r = check_graph( + g, _bind((4, 1, 1), {"n_elements": 1024}, x_ptr=_facts(10), out_ptr=_facts(10)) + ) + assert _found(r) == [("out-of-bounds", 0)] + assert r.abstained == ((1, K.DATA_DEPENDENT_MASK), (2, K.DATA_DEPENDENT_MASK)) + # UNSAT behind a dropped mask is still a proof + _clean( + check_graph( + g, + _bind( + (4, 1, 1), + {"n_elements": 1024}, + x_ptr=_facts(1024), + out_ptr=_facts(1024), + ), + ) + ) + + +def test_observation_gated_mask_and_path_abstain(): + atomic = AccessEvent("atomic_rmw", "c", Const(0), None, 32, None, 1) + gate = Cmp("slt", Observed(0), Const(4), 32) + g = AccessGraph( + "k", + [FuncArg("c", True, 32), FuncArg("p", True, 32)], + [ + atomic, + _access(Const(10), mask=gate, line=2), + _access(Const(10), path=gate, line=3), + _access(Const(0), mask=gate, line=4), # in bounds: a proof + ], + None, + ) + r = check_graph(g, _bind(c=_facts(1), p=_facts(4))) + assert r.findings == () + assert r.abstained == ((1, K.DATA_DEPENDENT_MASK), (2, K.UNMODELABLE_CONDITION)) + + +@pytest.mark.parametrize( + "name", ["p4_observed_direct", "p4_observed_loop", "p4_observed_delta"] +) +def test_observation_in_address_refuses(name): + """An atomic's old value in an address (directly, in a loop-carried + pointer's offset0 or in its delta) would make any address reachable.""" + g = _graph(f"reader_ttir/{name}.ttir") + r = check_graph(g, _bind((4, 1, 1), {"n": 8}, cnt_ptr=_facts(1), x_ptr=_facts(8))) + assert r.findings == () + assert r.abstained == ((1, K.OBSERVATION_IN_ADDRESS), (2, K.OBSERVATION_IN_ADDRESS)) + assert r.refusal.kind == "observation-in-address" + assert r.refusal.line_no == g.accesses[1].line_no + assert r.refusal.message.endswith( + "the address depends on the value an atomic observed" + ) + + +# ─────────────────────────── bindings and loops ─────────────────────────── + + +def test_missing_bindings_refuse(): + tensors = _all(4096, *ADD_TENSORS) + r = check_graph(_graph(ADD), _bind((4, 1, 1), **tensors)) # no n_elements + assert r.findings == () and r.abstained == tuple( + (i, K.MISSING_BINDING) for i in range(3) + ) + assert "scalar argument 'n_elements' has no launch binding" in r.refusal.message + b = _bind((4, 1, 1), {"n_elements": 4096}, **tensors) + b = LaunchBinding( + b.params, b.tensors, b.constexprs, None, None, b.config, "grid: boom" + ) + r = check_graph(_graph(ADD), b) + assert r.abstained == tuple((i, K.MISSING_BINDING) for i in range(3)) + assert "grid is unknown (unreadable: grid: boom)" in r.refusal.message + + +def test_an_unbound_loop_bound_refuses_the_loop(): + g = _graph("reader_ttir/iv_wrap.ttir") + r = check_graph(g, _bind(params={"lo": 0}, x_ptr=_facts(10))) + assert r.findings == () and r.abstained == ((0, K.MISSING_BINDING),) + assert r.refusal.line_no == g.loop.line_no + assert "scalar argument 'n'" in r.refusal.message + + +def test_unmodeled_values_and_unusable_facts_refuse(): + """Defensive: the reader never lets a DataDep into an address, and + bind_launch never builds such facts.""" + g = _synthetic(_access(Bin("+", Const(0), DataDep("loaded value")))) + r = check_graph(g, _bind(p=_facts(4))) + assert r.abstained == ((0, K.UNMODELED_VALUE),) + assert "an unmodeled value (loaded value)" in r.refusal.message + bad = TensorFacts(PTR, 4, 4, (4,), (), "torch.float32", False) + r = check_graph(_synthetic(_access(Const(0))), _bind(p=bad)) + assert r.abstained == ((0, K.MISSING_BINDING),) + assert "unusable" in r.refusal.message + + +def test_exact_finding_kept_alongside_a_refusal(): + """A tensor without a binding refuses its access only: the other + accesses' exact findings are returned with it.""" + r = _add((5, 1, 1), 10**9, **_all(4096, "x_ptr", "out_ptr")) + assert _found(r) == [("out-of-bounds", 0), ("out-of-bounds", 2)] + assert r.abstained == ((1, K.MISSING_BINDING),) + assert r.refusal == Refusal( + kind="missing-binding", + message=r.refusal.message, + line_no=r.findings[0].line_no + 3, + loc=r.refusal.loc, + ) + assert "pointer argument 'y_ptr' has no tensor binding" in r.refusal.message + + +_LOOP_THEN_TAIL = """ + %c0 = arith.constant 0 : i32 + %c10 = arith.constant 10 : i32 + %c100 = arith.constant 100 : i32 + scf.for %i = %c0 to %c10 step %n : i32 { + %a = tt.addptr %p, %i : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + } + %b = tt.addptr %p, %c100 : !tt.ptr, i32 + tt.store %b, %c0 : !tt.ptr +""" + + +@pytest.mark.parametrize("step", [0, -1]) +def test_non_positive_step_refuses_only_the_loop(step): + g, text = _module(_LOOP_THEN_TAIL) + r = check_graph(g, _bind(params={"arg1": step}, arg0=_facts(10))) + assert r.abstained == ((0, K.NON_POSITIVE_STEP),) + assert r.refusal.line_no == _line(text, "scf.for") + assert _found(r) == [("out-of-bounds", 1)] and r.findings[0].violation_offset == 100 + r = check_graph(g, _bind(params={"arg1": 2}, arg0=_facts(10))) + assert _found(r) == [("out-of-bounds", 1)] and r.abstained == () + + +def test_zero_trip_loop_has_no_footprint(): + """An access in the loop whose offset does not mention the induction + variable still runs only on iterations that run (#361 gave a zero-trip + loop's body a witness).""" + g, _ = _module( + """ + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + %c100 = arith.constant 100 : i32 + scf.for %i = %c0 to %n step %c1 : i32 { + %a = tt.addptr %p, %c100 : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + }""" + ) + assert g.accesses[0].in_loop + _clean(check_graph(g, _bind(params={"arg1": 0}, arg0=_facts(10)))) + _clean(check_graph(g, _bind(params={"arg1": -5}, arg0=_facts(10)))) + ((f,),) = [check_graph(g, _bind(params={"arg1": 1}, arg0=_facts(10))).findings] + assert f.violation_offset == 100 and f.witness["iter_loop"] == 0 + + +def test_solver_unknown_abstains(): + """x^3 + y^3 == z^3 with positive x, y, z (no solution, which Z3 cannot + prove in time) gates an always-OOB offset: unknown, never a proof.""" + x, y, z = Pid(0), Pid(1), Pid(2) + + def cube(t): + return Bin("*", Bin("*", t, t), t) + + def pos(t): + return Cmp("sgt", t, Const(0)) + + fermat = BoolBin( + "and", + BoolBin("and", pos(x), pos(y)), + BoolBin("and", pos(z), Cmp("eq", Bin("+", cube(x), cube(y)), cube(z))), + ) + g = _synthetic(_access(Const(-1), mask=fermat)) + before = z3.get_param("timeout") + r = check_graph(g, _bind((1 << 20,) * 3, p=_facts(4)), timeout_ms=50) + assert z3.get_param("timeout") == before # a per-Solver timeout only + assert r.findings == () and r.abstained == ((0, K.SOLVER_UNKNOWN),) + assert r.refusal.kind == "solver-unknown" + assert "could not decide the out-of-bounds query" in r.refusal.message + + +# ─────────────────────────── view footprint ─────────────────────────── + + +def test_strided_view_gap_is_out_of_bounds(): + """x[::2] has 4096 elements two apart: the odd offsets between them are + outside the view, although every offset is below numel.""" + x = torch.empty(8192)[::2] + tensors = {"x_ptr": tensor_facts(x), **_all(4096, "y_ptr", "out_ptr")} + r = _add((4, 1, 1), 4096, **tensors) + ((f,),) = [r.findings] + assert (f.kind, f.access_index) == ("out-of-bounds", 0) and r.refusal is None + assert f.violation_offset % 2 == 1 and 0 < f.violation_offset < 4096 + assert f.violation_address == x.data_ptr() + f.violation_offset * 4 + assert "shape (4096,) and strides (2,)" in f.detail + + +def _tile2d(stride_m: int, stride_n: int, inp, out) -> CheckResult: + params = {"M": 64, "N": 64, "stride_m": stride_m, "stride_n": stride_n} + return check_graph( + _graph(TILE2D), + _bind((2, 2, 1), params, in_ptr=tensor_facts(inp), out_ptr=tensor_facts(out)), + ) + + +def test_strided_views_indexed_by_their_strides_are_clean(): + # a column slice (strides (128, 1)) and a transpose (a dense permutation) + sliced = torch.empty(64, 128)[:, :64] + _clean(_tile2d(128, 1, sliced, torch.empty(64, 128)[:, :64])) + t = torch.empty(64, 64).t() + _clean(_tile2d(1, 64, t, torch.empty(64, 64).t())) + # indexing the slice as if it were dense reads the gaps + r = _tile2d(64, 1, sliced, torch.empty(64, 64)) + ((f,),) = [r.findings] + assert f.access_index == 0 and 64 <= f.violation_offset % 128 < 128 + + +def test_broadcast_stride0_tensor(): + """A row expanded to 64x64 (strides (0, 1)) has 64 distinct elements: + offsets past them are out of bounds although below numel (4096).""" + row = torch.empty(64).expand(64, 64) + assert tensor_facts(row).numel == 4096 + _clean(_tile2d(0, 1, row, torch.empty(64, 64))) + r = _tile2d(64, 1, row, torch.empty(64, 64)) + ((f,),) = [r.findings] + assert f.access_index == 0 and 64 <= f.violation_offset < 4096 + + +def _views(): + base = torch.empty(64) + return { + "contiguous": base, + "step2": base[::2], + "offset_step3": base[3:40:3], + "column_slice": base.view(8, 8)[:, :5], + "transpose": base.view(8, 8).t(), + "grid_slice": base.view(8, 8)[::2, 1::3], + "broadcast_rows": torch.empty(5).expand(3, 5), + "broadcast_cols": torch.empty(4, 1).expand(4, 6), + "gappy": base.as_strided((3, 3), (5, 1)), + "overlapping": base.as_strided((3, 4), (2, 1)), + "scalar": torch.empty(()), + "empty": torch.empty(0), + "empty_2d": torch.empty(3, 0), + "channels_last": torch.empty(2, 3, 2, 2).to(memory_format=torch.channels_last), + "expand_to_0": torch.empty(1).expand(0), + "size1_odd_stride": base.as_strided((1, 4), (100, 1)), + "mixed_broadcast": base.as_strided((3, 4, 5), (0, 10, 2)), + "duplicate_strides": base.as_strided((3, 3), (3, 3)), + } + + +def _eager_offsets(t: torch.Tensor) -> set[int]: + """The element offsets the eager sanitizer admits for ``t``: the byte + segments SymbolicClient._tensor_physical_addresses dispatches to + (contiguous, storage-contiguous, inner-stride-1 slices, per element).""" + if not t.numel(): + return set() + base, item = t.data_ptr(), t.element_size() + return { + (a - base) // item + for start, end, _ in SymbolicClient._tensor_physical_addresses(None, "t", t) + for a in range(start, end + 1) + if (a - base) % item == 0 + } + + +@pytest.mark.parametrize("name", list(_views())) +def test_view_footprint_matches_the_eager_legal_set(name): + """The element offsets _in_view admits are exactly the ones the eager + sanitizer's own dispatch admits.""" + t = _views()[name] + facts = tensor_facts(t) + eager = _eager_offsets(t) + for e in range(-3, max([70, *eager]) + 8): + solver = z3.Solver() + solver.add(_in_view(z3.IntVal(e), facts)) + assert (solver.check() == z3.sat) == (e in eager), (name, e) + + +def test_empty_tensor(): + """No element is legal: every executed access is out of bounds, and an + access that never executes is none.""" + g = _graph("ttir/golden_cas_sm80.ttir") + empty = TensorFacts(PTR, 4, 0, (0,), (1,), "torch.int32", True) + r = check_graph(g, _bind(lock_ptr=empty, out_ptr=_facts(1))) + ((f,),) = [r.findings] + assert (f.access_index, f.access_kind, f.violation_offset) == (0, "atomic_cas", 0) + assert "the empty tensor" in f.detail + _clean(_add((4, 1, 1), 0, **{n: empty for n in ADD_TENSORS})) + + +def test_reinterpreted_width_checks_every_byte(): + """An f32 pointer over a byte tensor (a width-changing reinterpret): + offsets count f32 elements, the view counts bytes.""" + g = _graph(ADD) + bytes_ = {n: _facts(4096, elem_size=1) for n in ADD_TENSORS} + _clean(check_graph(g, _bind((1, 1, 1), {"n_elements": 1024}, **bytes_))) + short = {**bytes_, "x_ptr": _facts(4095, elem_size=1)} + r = check_graph(g, _bind((1, 1, 1), {"n_elements": 1024}, **short)) + ((f,),) = [r.findings] + assert f.access_index == 0 and f.violation_offset == 1023 + assert f.violation_address == PTR + 1023 * 4 + + +# ─────────────────────────── integer widths ─────────────────────────── + + +def test_i32_wrap_is_integer_overflow_not_out_of_bounds(): + """(pid * S) * S wraps to 0 in i32 for S = 65536: every program stores + x[0]. The unbounded reading's offsets are an overflow, not OOB.""" + g = _graph("reader_ttir/rv_i32_wrap.ttir") + r = check_graph(g, _bind((4, 1, 1), {"S": 65536}, x_ptr=_facts(1))) + ((f,),) = [r.findings] + assert (f.kind, f.access_index) == ("integer-overflow", 0) and r.refusal is None + assert f.violation_offset is None and f.violation_address is None + assert not -(1 << 31) <= f.witness["value"] < 1 << 31 + assert ( + "arith.muli" + in _text("reader_ttir/rv_i32_wrap.ttir").splitlines()[f.line_no - 1] + ) + assert "i32 range" in f.detail + # no wrap for a small S: a clean launch + _clean(check_graph(g, _bind((4, 1, 1), {"S": 2}, x_ptr=_facts(16)))) + + +def test_trunci_is_integer_overflow_not_out_of_bounds(): + """trunc_i32(pid_i64 * 2**32) is 0 for every pid.""" + g = _graph("reader_ttir/rv_trunci_alias.ttir") + r = check_graph(g, _bind((2, 1, 1), x_ptr=_facts(1))) + ((f,),) = [r.findings] + assert (f.kind, f.witness["pid_0"], f.witness["value"]) == ( + "integer-overflow", + 1, + 1 << 32, + ) + assert ( + "arith.trunci" + in _text("reader_ttir/rv_trunci_alias.ttir").splitlines()[f.line_no - 1] + ) + assert "truncation to i32" in f.detail + _clean(check_graph(g, _bind((1, 1, 1), x_ptr=_facts(1)))) + + +def test_loop_increment_overflow(): + """range(lo, n, 1 << 20) wraps its induction variable near INT32_MAX.""" + g = _graph("reader_ttir/iv_wrap.ttir") + n = (1 << 31) - (1 << 19) + r = check_graph(g, _bind(params={"lo": 0, "n": n}, x_ptr=_facts(1 << 31))) + ((f,),) = [r.findings] + assert (f.kind, f.access_index) == ("integer-overflow", 0) + assert "scf.for" in _text("reader_ttir/iv_wrap.ttir").splitlines()[f.line_no - 1] + assert "induction-variable increment" in f.detail + _clean(check_graph(g, _bind(params={"lo": 0, "n": 1000}, x_ptr=_facts(1000)))) + # a zero-trip loop never increments + _clean(check_graph(g, _bind(params={"lo": n, "n": n}, x_ptr=_facts(1)))) + + +def test_unsigned_read_of_a_negative_param(): + """``pid < n.to(tl.uint32)`` reads n unsigned; the model reads the i32 + argument signed (2**32 - 1 and -1 are one i32).""" + g = _graph("reader_ttir/unsigned_index.ttir") + _clean(check_graph(g, _bind((8, 1, 1), {"n": 8}, x_ptr=_facts(3)))) + for n in (-1, (1 << 32) - 1): + r = check_graph(g, _bind((8, 1, 1), {"n": n}, x_ptr=_facts(3))) + ((f,),) = [r.findings] + assert f.kind == "integer-overflow" and f.witness["value"] == -1 + assert ( + "cmpi ult" + in _text("reader_ttir/unsigned_index.ttir").splitlines()[f.line_no - 1] + ) + + +def test_one_finding_per_op_site(): + """The add kernel's three accesses share one mask (pid * 1024 + lane): + its wrap is reported once.""" + r = _add((1 << 22, 1, 1), (1 << 31) - 1, **_all(1 << 32, *ADD_TENSORS)) + assert _found(r) == [("integer-overflow", 0)] and r.refusal is None + + +# A wrap that decides its own role's divisor (with BIG = 2**31 - 1): +# t = pid + BIG, d = where(t > BIG, 0, -3). At pid 1, +# t wraps in i32, so the kernel divides by -3 and r = t // d - BIG // -3 is +# 1431655764 (0 at pid 0); in the unbounded reading d is 0 there. Checking +# the wrap assuming the divisor non-zero, and the divisor assuming no wrap, +# neither query has a model: that was a false proof in every role. +_CIRCULAR = """ + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + %c4 = arith.constant 4 : i32 + %cm3 = arith.constant -3 : i32 + %pid = tt.get_program_id x : i32 + %t = arith.addi %pid, %n : i32 + %gt = arith.cmpi sgt, %t, %n : i32 + %d = arith.select %gt, %c0, %cm3 : i32 + %q = arith.divsi %t, %d : i32 + %b3 = arith.divsi %n, %cm3 : i32 + %r = arith.subi %q, %b3 : i32 +""" +_CIRCULAR_ACCESS = { + # x + r + "offset": """ + %a = tt.addptr %p, %r : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr""", + # if r > 0: x[1] + "path": """ + %pos = arith.cmpi sgt, %r, %c0 : i32 + scf.if %pos { + %a = tt.addptr %p, %c1 : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + }""", + # x + arange(16), masked to lanes < r + 1 + "mask": """ + %lim = arith.addi %r, %c1 : i32 + %offs = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32> + %sl = tt.splat %lim : i32 -> tensor<16xi32> + %m = arith.cmpi slt, %offs, %sl : tensor<16xi32> + %ps = tt.splat %p : !tt.ptr -> tensor<16x!tt.ptr> + %a = tt.addptr %ps, %offs : tensor<16x!tt.ptr>, tensor<16xi32> + %z = arith.constant dense<0> : tensor<16xi32> + tt.store %a, %z, %m : tensor<16x!tt.ptr>""", + # for i in range(min(r + 1, 4)): x[i] + "loop": """ + %lim = arith.addi %r, %c1 : i32 + %hi = arith.minsi %lim, %c4 : i32 + scf.for %i = %c0 to %hi step %c1 : i32 { + %a = tt.addptr %p, %i : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + }""", +} + + +@pytest.mark.parametrize("role", list(_CIRCULAR_ACCESS)) +def test_a_role_never_assumes_its_own_conditions(role): + """Each role is one joint query: the wrap at pid 1 is found, as the + innermost failure (the divisor's zero is only the unbounded reading's).""" + g, text = _module(_CIRCULAR + _CIRCULAR_ACCESS[role]) + big = (1 << 31) - 1 + r = check_graph(g, _bind((2, 1, 1), {"arg1": big}, arg0=_facts(1))) + assert _found(r) == [("integer-overflow", 0)] and r.refusal is None + (f,) = r.findings + assert f.line_no == _line(text, "arith.addi %pid") + assert (f.witness["pid_0"], f.witness["value"]) == (1, 1 << 31) + # pid 0 alone is in bounds, and proved so + _clean(check_graph(g, _bind((1, 1, 1), {"arg1": big}, arg0=_facts(1)))) + + +def test_a_divisor_zero_where_a_sibling_wraps_is_found(): + """The reverse dependence: pid 1 both wraps (pid + BIG) and divides by + zero (d = 1 - pid); either is a finding, never a proof.""" + g, _ = _module( + """ + %c1 = arith.constant 1 : i32 + %pid = tt.get_program_id x : i32 + %t = arith.addi %pid, %n : i32 + %d = arith.subi %c1, %pid : i32 + %q = arith.divsi %pid, %d : i32 + %o = arith.subi %t, %n : i32 + %s = arith.addi %o, %q : i32 + %a = tt.addptr %p, %s : !tt.ptr, i32 + %v = tt.load %a : !tt.ptr""" + ) + r = check_graph(g, _bind((2, 1, 1), {"arg1": (1 << 31) - 1}, arg0=_facts(1))) + (f,) = r.findings + assert f.kind in ("integer-overflow", "division-by-zero") + assert f.witness["pid_0"] == 1 and r.refusal is None + + +_SELECT_ARM = """ + %c0 = arith.constant 0 : i32 + %c10 = arith.constant 10 : i32 + %big = arith.constant 1073741824 : i32 + %pid = tt.get_program_id x : i32 + %is0 = arith.cmpi eq, %pid, %c0 : i32 + %m = arith.muli %pid, %big : i32 +""" + + +def test_a_wrap_in_the_arm_a_select_discards_is_no_finding(): + """off = where(pid == 0, pid * 2**30, pid): pid * 2**30 wraps from pid 2 + on, but only pid 0 reads it.""" + g, _ = _module( + _SELECT_ARM + + """ + %o = arith.select %is0, %m, %pid : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + %v = tt.load %a : !tt.ptr""" + ) + _clean(check_graph(g, _bind((4, 1, 1), arg0=_facts(4)))) + + +def test_a_wrap_in_the_arm_a_select_takes_is_a_finding(): + g, text = _module( + _SELECT_ARM + + """ + %o = arith.select %is0, %pid, %m : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + %v = tt.load %a : !tt.ptr""" + ) + (f,) = check_graph(g, _bind((4, 1, 1), arg0=_facts(1 << 32))).findings + assert (f.kind, f.line_no) == ("integer-overflow", _line(text, "arith.muli")) + assert f.witness["pid_0"] in (2, 3) + + +def test_a_value_also_read_directly_is_checked_unguarded(): + """pid * 2**30 sits in the select's discarded arm of the offset, but the + mask reads it directly: its wrap matters where the mask is computed.""" + g, text = _module( + _SELECT_ARM + + """ + %lt = arith.cmpi slt, %m, %c10 : i32 + %o = arith.select %is0, %m, %pid : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + %v = tt.load %a, %lt : !tt.ptr""" + ) + (f,) = check_graph(g, _bind((4, 1, 1), arg0=_facts(4))).findings + assert (f.kind, f.line_no) == ("integer-overflow", _line(text, "arith.muli")) + + +def test_undefined_divisions_count_in_either_arm(): + """where(n != 0, pid // n, 0) divides by zero in the discarded arm too + (the GPU run faults), and so does INT_MIN // -1.""" + g, text = _module( + """ + %c0 = arith.constant 0 : i32 + %cm1 = arith.constant -1 : i32 + %pid = tt.get_program_id x : i32 + %nz = arith.cmpi ne, %n, %c0 : i32 + %q = arith.divsi %pid, %n : i32 + %o = arith.select %nz, %q, %c0 : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + %v = tt.load %a : !tt.ptr""" + ) + r = check_graph(g, _bind((4, 1, 1), {"arg1": 0}, arg0=_facts(4))) + assert _found(r) == [("division-by-zero", 0)] + _clean(check_graph(g, _bind((4, 1, 1), {"arg1": 2}, arg0=_facts(4)))) + g, text = _module( + """ + %c0 = arith.constant 0 : i32 + %c5 = arith.constant 5 : i32 + %cm1 = arith.constant -1 : i32 + %pid = tt.get_program_id x : i32 + %never = arith.cmpi eq, %pid, %c5 : i32 + %q = arith.divsi %n, %cm1 : i32 + %o = arith.select %never, %q, %c0 : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + %v = tt.load %a : !tt.ptr""" + ) + (f,) = check_graph(g, _bind(params={"arg1": -(1 << 31)}, arg0=_facts(1))).findings + assert (f.kind, f.line_no) == ("integer-overflow", _line(text, "arith.divsi")) + assert f.witness["value"] == 1 << 31 and "the result of '//'" in f.detail + + +# ─────────────────────────── division by zero ─────────────────────────── + + +def test_division_by_zero_in_an_offset(): + g, text = _module( + """ + %pid = tt.get_program_id x : i32 + %q = arith.divsi %pid, %n : i32 + %a = tt.addptr %p, %q : !tt.ptr, i32 + %v = tt.load %a : !tt.ptr""" + ) + r = check_graph(g, _bind((4, 1, 1), {"arg1": 0}, arg0=_facts(2))) + assert _found(r) == [("division-by-zero", 0)] and r.refusal is None + assert r.findings[0].line_no == _line(text, "arith.divsi") + assert "divisor of this division" in r.findings[0].detail + _clean(check_graph(g, _bind((4, 1, 1), {"arg1": 2}, arg0=_facts(2)))) + + +def test_division_by_zero_in_a_mask(): + g, text = _module( + """ + %pid = tt.get_program_id x : i32 + %c1 = arith.constant 1 : i32 + %r = arith.remsi %pid, %n : i32 + %m = arith.cmpi slt, %r, %c1 : i32 + %v = tt.load %p, %m : !tt.ptr""" + ) + r = check_graph(g, _bind((4, 1, 1), {"arg1": 0}, arg0=_facts(1))) + assert _found(r) == [("division-by-zero", 0)] + assert r.findings[0].line_no == _line(text, "arith.remsi") + assert "remainder" in r.findings[0].detail + _clean(check_graph(g, _bind((4, 1, 1), {"arg1": 1}, arg0=_facts(1)))) + + +def test_division_by_zero_in_a_loop_bound(): + g, text = _module( + """ + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + %c8 = arith.constant 8 : i32 + %u = arith.divsi %c8, %n : i32 + scf.for %i = %c0 to %u step %c1 : i32 { + %a = tt.addptr %p, %i : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + }""" + ) + r = check_graph(g, _bind(params={"arg1": 0}, arg0=_facts(4))) + assert _found(r) == [("division-by-zero", 0)] + assert r.findings[0].line_no == _line(text, "arith.divsi") + _clean(check_graph(g, _bind(params={"arg1": 2}, arg0=_facts(4)))) + ((f,),) = [check_graph(g, _bind(params={"arg1": 2}, arg0=_facts(3))).findings] + assert f.kind == "out-of-bounds" and f.violation_offset == 3 + + +# ─────────────────────────── robustness ─────────────────────────── + + +def test_deep_terms_are_lowered_iteratively(): + """kernel_deep_chain's offset nests more than 1000 levels: at Python's + default recursion limit the generated == / hash would raise.""" + g = _graph("ttir/kernel_deep_chain.ttir") + limit = sys.getrecursionlimit() + sys.setrecursionlimit(1000) + try: + clean = check_graph(g, _bind((4, 1, 1), {"s": 1}, out_ptr=_facts(1 << 20))) + wrap = check_graph(g, _bind((4, 1, 1), {"s": 3}, out_ptr=_facts(1 << 20))) + finally: + sys.setrecursionlimit(limit) + _clean(clean) + assert _found(wrap) == [("integer-overflow", 0)] + + +def test_a_launch_key_ignores_only_the_addresses(): + """Launches whose bindings share a launch_key get one result, the + findings' addresses moved to each launch's tensors; nothing else in a + result (a withheld finding's message included) holds an address.""" + g = _graph(ADD) + + def at(ptr): + return _bind( + (5, 1, 1), + {"n_elements": 10**9}, + **{n: _facts(4096, data_ptr=ptr) for n in ADD_TENSORS}, + ) + + first, moved = at(PTR), at(PTR + 0x10000) + assert launch_key(first) == launch_key(moved) + result = check_graph(g, first) + again = readdressed(result, g, moved) + assert again == check_graph(g, moved) != result + assert [f.violation_address for f in again.findings] == [ + PTR + 0x10000 + f.violation_offset * 4 for f in result.findings + ] + assert launch_key(at(PTR)) != launch_key(_bind((5, 1, 1), {"n_elements": 1})) + guarded = _synthetic(_access(Bin("-", Pid(0), Const(1)), guarded=True)) + messages = { + check_graph(guarded, _bind((4, 1, 1), p=_facts(4, data_ptr=ptr))).refusal + for ptr in (PTR, PTR + 0x10000) + } + assert len(messages) == 1 and "0x" not in messages.pop().message + + +@pytest.mark.parametrize("timeout_ms", [0, -1, 2.5, True]) +def test_a_timeout_must_be_positive(timeout_ms): + # Z3 reads 0 and below as no timeout at all. + with pytest.raises(ValueError, match="positive int"): + check_graph(_graph(ADD), _bind(), timeout_ms=timeout_ms) + + +class _InterruptedSolver: + """A Solver whose query Z3's Ctrl+C handler cancelled.""" + + def __init__(self, ctx=None): + pass + + def set(self, **kwargs): + pass + + def add(self, *formulas): + pass + + def check(self): + return z3.unknown + + def reason_unknown(self): + return "interrupted from keyboard" + + +def test_a_ctrl_c_during_a_query_is_raised(monkeypatch): + """Z3 turns a Ctrl+C into an unknown; it is the user's interrupt, never + a solver-unknown abstention.""" + monkeypatch.setattr(oob, "Solver", _InterruptedSolver) + with pytest.raises(KeyboardInterrupt): + _add((4, 1, 1), 4096, **_all(4096, *ADD_TENSORS)) + + +_SIGINT_SCRIPT = """ +import os, signal, threading, time +from types import MappingProxyType +from tilelens.clients.sanitizer.compiled.oob import check_graph +from tilelens.ir.launch import LaunchBinding, TensorFacts +from tilelens.ir.ttir_reader import ( + AccessEvent, AccessGraph, Bin, BoolBin, Cmp, Const, FuncArg, Pid, +) + +def cube(t): + return Bin("*", Bin("*", t, t), t) + +x, y, z = Pid(0), Pid(1), Pid(2) +pos = [Cmp("sgt", t, Const(0)) for t in (x, y, z)] +fermat = BoolBin("and", BoolBin("and", pos[0], pos[1]), BoolBin( + "and", pos[2], Cmp("eq", Bin("+", cube(x), cube(y)), cube(z)))) +graph = AccessGraph("k", [FuncArg("p", True, 32)], + [AccessEvent("load", "p", Const(-1), fermat, 32, None, 1)], None) +facts = TensorFacts(4096, 4, 4, (4,), (1,), "torch.float32", True) +grid = (1 << 20,) * 3 +binding = LaunchBinding(MappingProxyType({}), MappingProxyType({"p": facts}), + MappingProxyType({}), grid, grid, MappingProxyType({})) +threading.Timer(1.5, os.kill, (os.getpid(), signal.SIGINT)).start() +start = time.monotonic() +try: + result = check_graph(graph, binding, timeout_ms=60_000) +except KeyboardInterrupt as exc: + print("interrupted", repr(exc), round(time.monotonic() - start)) +else: + print("returned", result.abstained) +""" + + +def test_a_real_ctrl_c_stops_a_hard_query(): + """A SIGINT 1.5 s into a query Z3 cannot decide in a minute (the Fermat + mask of test_solver_unknown_abstains): Z3 cancels it, and the check + raises KeyboardInterrupt at once instead of abstaining and going on.""" + proc = subprocess.run( + [sys.executable, "-c", _SIGINT_SCRIPT], + capture_output=True, + text=True, + cwd=REPO, + timeout=120, + ) + assert proc.returncode == 0, proc.stderr + assert proc.stdout.startswith( + "interrupted KeyboardInterrupt('Z3 query" + ), proc.stdout + assert int(proc.stdout.split()[-1]) < 30 + + +_THREADS_SCRIPT = """ +import sys, threading +from pathlib import Path +from types import MappingProxyType +from tilelens.clients.sanitizer.compiled.oob import check_graph +from tilelens.ir.launch import LaunchBinding, TensorFacts +from tilelens.ir.ttir_reader import parse_ttir + +graph = parse_ttir(Path(sys.argv[1]).read_text()) +facts = TensorFacts(4096, 4, 4096, (4096,), (1,), "torch.float32", True) +grid = (5, 1, 1) +binding = LaunchBinding( + MappingProxyType({"n_elements": 10**6}), + MappingProxyType(dict.fromkeys(("x_ptr", "y_ptr", "out_ptr"), facts)), + MappingProxyType({}), grid, grid, MappingProxyType({})) +errors, kinds = [], set() + +def worker(): + try: + for _ in range(25): + kinds.add(tuple(f.kind for f in check_graph(graph, binding).findings)) + except Exception as exc: + errors.append(repr(exc)) + +threads = [threading.Thread(target=worker) for _ in range(4)] +for t in threads: + t.start() +for t in threads: + t.join() +print(errors[:3], sorted(kinds)) +""" + + +def _corpus_file(tmp_path: Path, name: str) -> Path: + """Corpus text ``name`` in a file under ``tmp_path``, for a subprocess.""" + path = tmp_path / Path(name).name + path.write_text(_text(name), encoding="utf-8") + return path + + +def test_checks_on_several_host_threads_at_once(tmp_path): + """Two traced kernels finalized on different host threads check at the + same time: each check has a Z3 context of its own (one context shared + by threads failed with Z3 errors or segfaulted). In a subprocess, so a + crash is a failure here and not the test run's end.""" + proc = subprocess.run( + [sys.executable, "-c", _THREADS_SCRIPT, str(_corpus_file(tmp_path, ADD))], + capture_output=True, + text=True, + cwd=REPO, + timeout=600, + ) + assert proc.returncode == 0, proc.stderr[-2000:] + oob_kinds = ("out-of-bounds",) * 3 + assert proc.stdout.strip() == f"[] [{oob_kinds!r}]" + + +def test_results_are_plain_data(): + r = _add((5, 1, 1), 10**9, **_all(4096, "x_ptr", "out_ptr")) + again = pickle.loads(pickle.dumps(r)) + assert again == r and isinstance(again.findings[0], Finding) + # a Refusal holds its kind as a plain str; the abstentions keep the enum + assert not isinstance(r.refusal.kind, SanitizerKind) + assert ( + r.refusal.kind == "missing-binding" and r.abstained[0][1] is K.MISSING_BINDING + ) + with pytest.raises(TypeError): + hash(r.findings[0]) + # the sanitizer's kinds are its own, apart from the reader's + assert not {k.value for k in SanitizerKind} & {k.value for k in TTIRKind} + assert f"{K.SOLVER_UNKNOWN}" == str(K.SOLVER_UNKNOWN) == "solver-unknown" diff --git a/tests/unit/test_wrapper.py b/tests/unit/test_wrapper.py index 12c478964..326fb4dc4 100644 --- a/tests/unit/test_wrapper.py +++ b/tests/unit/test_wrapper.py @@ -1,5 +1,7 @@ +import os import subprocess import sys +from pathlib import Path import pytest from unittest.mock import MagicMock, patch @@ -7,15 +9,20 @@ import tilelens from tilelens.core.config import config as cfg from tilelens.core.trace import TraceInterface +from tilelens.clients.sanitizer.compiled import CompiledSanitizer from tilelens.wrapper import ( + COMPILE_NOTE, create_patched_jit, create_patched_autotune, sanitizer_wrapper, + compiled_sanitizer_wrapper, profiler_wrapper, apply_sanitizer, apply_profiler, ) +REPO = Path(__file__).resolve().parents[2] + @pytest.fixture def _isolate_cli_active(): @@ -60,6 +67,24 @@ def test_sanitizer_wrapper_accepts_frontend(): assert result == "wrapped_kernel" +def test_compiled_sanitizer_wrapper_traces_with_the_compiled_sanitizer(monkeypatch): + monkeypatch.setattr(cfg, "enable_sanitizer", True) + mock_kernel = MagicMock() + mock_kernel.__name__ = "test_kernel" + + with patch("tilelens.wrapper.tilelens.trace") as mock_trace: + mock_decorator = MagicMock(return_value="wrapped_kernel") + mock_trace.return_value = mock_decorator + + result = compiled_sanitizer_wrapper(mock_kernel, frontend="gluon") + + client = mock_trace.call_args.kwargs["client"] + assert isinstance(client, CompiledSanitizer) and client.abort_on_error + assert mock_trace.call_args.kwargs["frontend"] == "gluon" + mock_decorator.assert_called_once_with(mock_kernel) + assert result == "wrapped_kernel" + + def test_profiler_wrapper_applies_trace(): mock_kernel = MagicMock() mock_kernel.__name__ = "test_kernel" @@ -256,3 +281,60 @@ def test_wrapper_imports_without_pytest(): code = "import sys; sys.modules['pytest'] = None; import tilelens.wrapper" proc = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True) assert proc.returncode == 0, proc.stderr + + +# ======== tile-sanitizer --compile =========== + +_CLIENTS_SCRIPT = """\ +import sys +import triton + + +@triton.jit +def kernel(x_ptr): + pass + + +print(sorted(kernel.client_manager.clients), sys.argv[1:]) +""" + + +def _run_cli(argv): + """Run apply_sanitizer in a subprocess (it patches triton.jit for good) + as the command ``argv[0]`` with ``argv[1:]``.""" + code = ( + f"import sys; sys.argv = {argv!r}; " + "from tilelens.wrapper import apply_sanitizer; apply_sanitizer()" + ) + env = {k: v for k, v in os.environ.items() if k != "TRITON_INTERPRET"} + return subprocess.run( + [sys.executable, "-c", code], capture_output=True, text=True, cwd=REPO, env=env + ) + + +@pytest.mark.parametrize( + "command, before, after, clients", + [ + ("tile-sanitizer", ["--compile"], [], ["compiled_sanitizer"]), + ("triton-sanitizer", ["--compile"], ["--compile"], ["compiled_sanitizer"]), + # After the script name the flag is the script's own argument. + ("tile-sanitizer", [], ["--compile"], ["sanitizer"]), + ], +) +def test_the_compile_flag_before_the_script_selects_the_compiled_sanitizer( + tmp_path, command, before, after, clients +): + script = tmp_path / "script.py" + script.write_text(_CLIENTS_SCRIPT) + proc = _run_cli([command, *before, str(script), *after]) + assert proc.returncode == 0, proc.stderr + assert proc.stdout.strip() == f"{clients} {after}" + # Under the flag the user is told, once, that kernels do not run. + assert proc.stderr.count(COMPILE_NOTE) == (1 if before else 0) + + +def test_the_compile_flag_without_a_script_prints_the_usage(): + proc = _run_cli(["tile-sanitizer", "--compile"]) + assert proc.returncode == 1 + assert "Usage: tile-sanitizer [--compile] [args...]" in proc.stdout + assert "kernels are not run" in proc.stdout diff --git a/tilelens/clients/__init__.py b/tilelens/clients/__init__.py index d4810590a..2ae7ba322 100644 --- a/tilelens/clients/__init__.py +++ b/tilelens/clients/__init__.py @@ -10,7 +10,15 @@ "OpTypeCounts": ("tilelens.clients.profiler.data", "OpTypeCounts"), "RaceDetector": ("tilelens.clients.race_detector.race_detector", "RaceDetector"), "Sanitizer": ("tilelens.clients.sanitizer.sanitizer", "Sanitizer"), + "CompiledSanitizer": ( + "tilelens.clients.sanitizer.compiled.client", + "CompiledSanitizer", + ), "OutOfBoundsRecord": ("tilelens.clients.sanitizer.data", "OutOfBoundsRecord"), + "CompiledSanitizerRecord": ( + "tilelens.clients.sanitizer.data", + "CompiledSanitizerRecord", + ), "SymbolicExpr": ("tilelens.clients.symbolic_engine", "SymbolicExpr"), "SymbolicClient": ("tilelens.clients.symbolic_engine", "SymbolicClient"), "RangeWrapper": ("tilelens.clients.symbolic_engine", "RangeWrapper"), diff --git a/tilelens/clients/sanitizer/compiled/__init__.py b/tilelens/clients/sanitizer/compiled/__init__.py new file mode 100644 index 000000000..e1ca894d7 --- /dev/null +++ b/tilelens/clients/sanitizer/compiled/__init__.py @@ -0,0 +1,42 @@ +"""The compiled-mode sanitizer (``Sanitizer(compile=True)``): out-of-bounds, +integer-overflow and division-by-zero checks over a kernel's compiled TTIR, +instantiated per launch with its LaunchBinding. + +``oob`` is the evaluator: an ``AccessGraph`` from ``tilelens.ir.ttir_reader`` +and a launch's binding in, a ``CheckResult`` out. ``client`` is the trace +client, ``CompiledSanitizer``, which runs it on every config a launch +compiled. + +Exports resolve on first access, so importing this package imports neither +Z3 nor Triton. +""" + +from __future__ import annotations + +from importlib import import_module +from typing import Any + + +_EXPORTS: dict[str, tuple[str, str]] = { + "CompiledSanitizer": ( + "tilelens.clients.sanitizer.compiled.client", + "CompiledSanitizer", + ), + "CheckResult": ("tilelens.clients.sanitizer.compiled.oob", "CheckResult"), + "Finding": ("tilelens.clients.sanitizer.compiled.oob", "Finding"), + "SanitizerKind": ("tilelens.clients.sanitizer.compiled.oob", "SanitizerKind"), + "check_graph": ("tilelens.clients.sanitizer.compiled.oob", "check_graph"), +} + +__all__ = list(_EXPORTS) + + +def __getattr__(name: str) -> Any: + try: + module_name, attr_name = _EXPORTS[name] + except KeyError as exc: + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") from exc + + value = getattr(import_module(module_name), attr_name) + globals()[name] = value + return value diff --git a/tilelens/clients/sanitizer/compiled/client.py b/tilelens/clients/sanitizer/compiled/client.py new file mode 100644 index 000000000..8cb6416be --- /dev/null +++ b/tilelens/clients/sanitizer/compiled/client.py @@ -0,0 +1,773 @@ +"""The compiled sanitizer client (``Sanitizer(compile=True)``). + +An ``IRClient`` that reads each compiled config's TTIR and checks the +launch against it with ``oob.check_graph``, without running the kernel +(``LAUNCH = "skip"``): output tensors are left untouched, and addresses +are the launch's tensor addresses. The TTIR is compiled on the host for the +client's target (``target=``, else ``TILELENS_IR_TARGET``, else +``cuda:89``), so no GPU is needed and a result never depends on the machine: +CPU tensors are checked like device ones. + +Every config the launch compiled is checked, autotune benchmark configs +included, each against the bindings the core delivered for it, and +gets its own ``ConfigVerdict``: configs that compile to one kernel but bind +other runtime arguments or grids are checked and reported apart. + +A config that failed to compile never fails the launch: the target is +this client's choice, not the machine's. What the failure means (see +_compiles_nowhere): + +- an error no target compiles past (a failing ``tl.static_assert``, a Python + construct Triton's code generator never accepts), from a compile that + never asked Triton's driver anything (nor did an earlier compile of the + kernel for the target, whose answer the kernel's code may keep): the + config never launches anywhere (Triton's autotuner drops it too), so it + is only noted; +- any other compile error may be the target's (``num_ctas > 1`` below + cuda:90, an fp8 type the target lacks, 16-bit descriptor atomic min/max + without native TMA, ...): the config may launch on the user's GPU and was + not checked, so it is unsupported, kind ``"compile-failed"``, with a + refusal naming the target and how to name another; +- a call that does not bind the kernel's parameters (a missing, extra or + misnamed argument) is no compile failure: the core raises the JIT + binder's error from the launch, as the untraced call raises it on any GPU, + so it never reaches this client from a traced launch; one + delivered anyway (an event built outside the core) is unsupported, kind + ``"compile-failed"``, naming no other target, which would not help; +- a host compile that could not run at all (Triton's compile API, or a + compile that asked for a device: ``HostCompileUnavailable``; or one + refused while an interpreted traced launch has the language patched: + ``LanguagePatchedError``) says nothing about the kernel: unsupported, + kind ``"host-compile-unavailable"``. + +Each refusal and note names where the kernel failed (``file:line`` of the +kernel's code the error points at, else of its ``def``) and the innermost +error in one line; the whole error stays on the ArtifactLog's +CompileFailure. A launch no config of which compiled is unsupported, with +the first failed config's refusal (``"compile-failed"`` when every config +was only noted; its notes are printed with it). +The launch's status is the union: ``"violations"`` if any config has a +finding, else ``"unsupported"`` if any config was not (fully) checked, else +``"ok"``. Nothing is ``"ok"`` before a check says so. + +What ``"ok"`` covers (``IRVerdict.scope == "launch"``): this launch, with +the scalar arguments, grid and tensors it was called with, as compiled for +the target (a kernel's TTIR can differ between targets, e.g. for tensor +descriptors below sm90). Since no kernel +runs, its outputs are never written: a later launch whose arguments the +host computes from them (a count, an offset, a size read back) is checked +with the values the untraced program would not have used, so a program's +launches can each be "ok" while the untraced program goes out of bounds. +In a trace shared with an eager client that runs the interpreter, the +interpreted run happens before this analysis, unchecked by it. + +A check is remembered: a later launch of the same compiled kernel with the +same scalar arguments, grid, and tensor shapes, strides and element sizes +(only the addresses differ, as in a training loop) reuses it, with the +findings' addresses moved to the new tensors. + +Each finding becomes a ``CompiledSanitizerRecord``; the launch's +``IRVerdict`` follows them in ``Launch.records``. With +``abort_on_error`` every finding is printed as this client finalizes the +launch, which then raises ``SystemExit(1)``: the core still finalizes the +trace's other clients and exits after them, but the launch keeps none of +this client's records and is not added to ``tilelens.launches``. On a host +thread other than the main one the exit ends only that thread (the +process's exit status is unchanged). Without ``abort_on_error`` findings +are printed only under ``TILELENS_VERBOSE=1``. The parts of a launch that +were not checked are printed alongside, each op once per client and +compiled kernel (a kernel launched in a loop would repeat them). +""" + +from __future__ import annotations + +import sys +from collections import OrderedDict +from collections.abc import Hashable, Mapping +from dataclasses import replace +from typing import Any, ClassVar + +from ....core.client import LanguagePatchedError, LaunchCall +from ....core.config import config as cfg +from ....core.data import Load, Store +from ....core.host_compile import ( + bind_failed, + format_ir_target, + host_compile_unavailable, + parse_ir_target, + target_queried, + unknown_options, +) +from ....ir.capture import ( + ArtifactLog, + CompiledSpecialization, + CompileFailure, + ParseCache, +) +from ....ir.client import IRClient +from ....ir.launch import LaunchBinding, TensorFacts, tensor_facts +from ....ir.verdict import ConfigVerdict, IRVerdict, Refusal, SourceLocation +from ....utils.traceback_utils import location_to_traceback_info +from ..data import CompiledSanitizerRecord +from ..sanitizer import Sanitizer +from .oob import ( + CheckResult, + Finding, + SanitizerKind, + check_graph, + describe_site, + launch_key, + readdressed, +) + +# Checks remembered per client (see the module docstring), least recently +# used dropped first. +_REMEMBERED_CHECKS = 256 + + +class CompiledSanitizer(IRClient): + """Out-of-bounds, integer-width and division-by-zero checks over a + kernel's compiled TTIR. + + ``records`` (the last launch's findings), ``last_verdict`` and + ``last_status`` (its status, None before a launch is finalized) are the + compatibility view of the last launch. + """ + + NAME = "compiled_sanitizer" + IR_STAGES: ClassVar[frozenset[str]] = frozenset({"ttir"}) + LAUNCH = "skip" + + def __init__( + self, + abort_on_error: bool = True, + timeout_ms: int = 10_000, + target: Any = None, + ) -> None: + """``target``: what the kernels are compiled for, e.g. ``"cuda:89"``, + ``"cuda:90"``, ``"hip:gfx942"`` or a triton ``GPUTarget`` (see + tilelens.core.host_compile.parse_ir_target); None for the configured + default (``TILELENS_IR_TARGET``, else ``"cuda:89"``). A spec that + names no target raises ValueError here.""" + super().__init__() + # Parsed now, so a bad spec fails at construction, not at a launch. + self.ir_target = None if target is None else parse_ir_target(target) + if ( + isinstance(timeout_ms, bool) + or not isinstance(timeout_ms, int) + or timeout_ms <= 0 + ): + # Z3 reads a timeout of 0 or below as none at all. + raise ValueError(f"timeout_ms must be a positive int, not {timeout_ms!r}") + self.abort_on_error = abort_on_error + # Bounds each Z3 query; one that times out is "solver-unknown". + self.timeout_ms = timeout_ms + # Kept across launches: a TTIR text is read once, and a refusal keeps + # its kind on every later hit. + self.parses = ParseCache() + self.records: list[CompiledSanitizerRecord] = [] + # Remembered checks: (id(graph), timeout, launch_key) -> (graph, + # result); holding the graph keeps its id unique. + self._checks: OrderedDict[Hashable, tuple[Any, CheckResult]] = OrderedDict() + # The not-checked parts already printed (see print_unchecked). + self._printed: set[Hashable] = set() + + @property + def last_status(self) -> str | None: + verdict = self.last_verdict + return None if verdict is None else verdict.status + + # ── launch lifecycle ───────────────────────────────────────────── + + def begin_launch(self, call: LaunchCall) -> None: + super().begin_launch(call) + self.records = [] + + def finalize(self) -> list: + out = super().finalize() + verdict = self.last_verdict + assert verdict is not None + if self.abort_on_error or cfg.verbose: + for record in self.records: + print_compiled_record(record) + print_unchecked(verdict, self._printed) + if self.abort_on_error and self.records: + sys.exit(1) + return out + + # ── the analysis ───────────────────────────────────────────────── + + def analyze_launch(self, log: ArtifactLog) -> tuple[list, IRVerdict]: + # Each failed config (see the module docstring): unsupported, unless + # it compiles for no target, which is only noted. + failed: list[ConfigVerdict] = [] + notes: list[str] = [] + for failure in log.failures: + if _compiles_nowhere(failure.error): + notes.append(_nowhere_note(failure)) + else: + failed.append( + ConfigVerdict( + None, + _saveable(failure.config), + "unsupported", + _failure_refusal(failure), + ) + ) + specializations = log.specializations + if not specializations: + if failed: + refusal = failed[0].refusal + elif log.failures: + refusal = _none_compiled(log) + else: + refusal = Refusal( + SanitizerKind.NO_COMPILED_KERNEL, _nothing_compiled(log) + ) + return [], IRVerdict( + self.NAME, + "unsupported", + refusal=refusal, + per_config=tuple(failed), + notes=tuple(notes), + ) + records: list[CompiledSanitizerRecord] = [] + per_config: list[ConfigVerdict] = [] + for spec in specializations: + for found, verdict in self._check_specialization(spec): + records += found + per_config.append(verdict) + per_config += failed + statuses = {verdict.status for verdict in per_config} + status = next(s for s in ("violations", "unsupported", "ok") if s in statuses) + # The first refusal: under "violations" too, where it says the + # findings may not be all there is. + refusal = next((c.refusal for c in per_config if c.refusal is not None), None) + self.records = records + return records, IRVerdict( + self.NAME, + status, + # A proof or a finding holds for this launch's arguments only. + scope=None if status == "unsupported" else "launch", + refusal=refusal, + per_config=tuple(per_config), + notes=tuple(notes), + ) + + def _check_specialization( + self, spec: CompiledSpecialization + ) -> list[tuple[list[CompiledSanitizerRecord], ConfigVerdict]]: + """One (records, ConfigVerdict) per config among ``spec``'s + bindings, in the order first seen (see _configs_of): each config is + checked against its own bindings only.""" + configs = _configs_of(spec) + try: + graph = self._graph_of(spec) + except Exception as exc: + # A bug in this kernel's analysis; the other kernels still count. + graph = Refusal(SanitizerKind.INTERNAL_ERROR, _describe(exc)) + if isinstance(graph, Refusal): + return [ + ([], ConfigVerdict(spec.specialization, config, "unsupported", graph)) + for config, _ in configs + ] + return [ + self._check_config(spec.specialization, graph, config, bindings) + for config, bindings in configs + ] + + def _graph_of(self, spec: CompiledSpecialization) -> Any: + """The access graph read from ``spec``'s TTIR, or the Refusal + saying why there is none.""" + text = spec.artifacts.stages.get("ttir") + if not isinstance(text, str): + why = spec.artifacts.error or "the compiled kernel holds no TTIR" + return Refusal(SanitizerKind.NO_TTIR, why) + outcome = self.parses.get(text) + if outcome.refusal is not None: + # The reader's kind and fields, the message prefixed with the + # location like the sanitizer's own refusals. + refused = Refusal.from_exception(outcome.refusal) + where = describe_site(refused.line_no, refused.loc) + return replace(refused, message=f"{where}: {refused.message}") + if outcome.error is not None: + return Refusal( + SanitizerKind.INTERNAL_ERROR, + f"the TTIR reader failed: {outcome.error}", + ) + return outcome.graph + + def _check_config( + self, + specialization: Hashable, + graph: Any, + config: dict[str, Any], + bindings: list[LaunchBinding], + ) -> tuple[list[CompiledSanitizerRecord], ConfigVerdict]: + if not bindings: + # Nothing to check against: never "ok" (the core delivers every + # compiled kernel with a binding; another producer might not). + return [], ConfigVerdict( + specialization, + config, + "unsupported", + Refusal( + SanitizerKind.INTERNAL_ERROR, + "no binding was delivered for this kernel, so nothing was checked", + ), + ) + try: + records: list[CompiledSanitizerRecord] = [] + refusal: Refusal | None = None + for binding in bindings: + result = self._check(graph, binding) + records += [ + _record(finding, graph.kernel_name, binding, config) + for finding in result.findings + ] + if refusal is None: + refusal = result.refusal + except Exception as exc: + # A bug in this config's analysis; the other configs still count. + refusal = Refusal(SanitizerKind.INTERNAL_ERROR, _describe(exc)) + return [], ConfigVerdict(specialization, config, "unsupported", refusal) + if records: + status = "violations" + else: + status = "ok" if refusal is None else "unsupported" + return records, ConfigVerdict( + specialization, config, status, refusal, n_reports=len(records) + ) + + def _check(self, graph: Any, binding: LaunchBinding) -> CheckResult: + """check_graph, remembered (see the module docstring).""" + key = (id(graph), self.timeout_ms, launch_key(binding)) + hit = self._checks.get(key) + if hit is not None and hit[0] is graph: + self._checks.move_to_end(key) + return readdressed(hit[1], graph, binding) + result = check_graph(graph, binding, timeout_ms=self.timeout_ms) + self._checks[key] = (graph, result) + if len(self._checks) > _REMEMBERED_CHECKS: + self._checks.popitem(last=False) + return result + + def on_analysis_error(self, exc: Exception) -> IRVerdict: + self.records = [] + return IRVerdict( + self.NAME, + "unsupported", + refusal=Refusal(SanitizerKind.INTERNAL_ERROR, _describe(exc)), + ) + + +def _describe(exc: BaseException) -> str: + return f"{type(exc).__name__}: {exc}" + + +def _nothing_compiled(log: ArtifactLog) -> str: + if log.call is not None and not log.call.capture: + return ( + "no compiled kernel to check: the launch has no JITFunction " + "(TRITON_INTERPRET=1, an InterpretedFunction runner, Gluon or NKI); " + "the eager Sanitizer() checks such launches" + ) + return "no config of the launch was compiled, so nothing was checked" + + +# How to check a kernel for another target, as a refusal or note says it. +_NAME_A_TARGET = "Sanitizer(compile=True, target=...), or TILELENS_IR_TARGET" + + +def _compiles_nowhere(error: BaseException | None) -> bool: + """Whether a config's host compile error says the config compiles for + no target, so it never launches anywhere (a note), rather than perhaps + only not for this client's target (unsupported, compile-failed). + + The rule, by exception type: the error the kernel's code raised (below + the CompilationErrors Triton's code generator wraps a called @jit + helper's or builtin's error in, following ``__cause__``) is a failing + ``tl.static_assert`` (CompileTimeAssertionFailure, which Triton's + autotuner drops a config for too) or a Python construct the code + generator never accepts (UnsupportedLanguageConstruct), and the compile + never asked Triton's driver anything (tilelens.core.host_compile. + target_queried: nor did an earlier compile of the kernel for the target): + a static_assert on ``tl.target_info``, on a device query the kernel's + code caught, or in a branch taken for the target only, is the target's. + Every other error may be the target's: Triton raises a target's + refusals (``num_ctas > 1`` below sm90, an fp8 type the target lacks, + 16-bit descriptor atomic min/max without native TMA, a dot shape the + target's MMA lacks, a PTXAS failure) as plain ValueError / + AssertionError / TypeError / RuntimeError, often wrapped in a + CompilationError like an error of the kernel's own code, so none of + them is taken to hold for every target. So is an unknown error (None). + A call that does not bind (bind_failed) never launches either, but the + untraced program raises for it (so does a traced launch: only an + event built outside the core delivers one), so it is no note; a host + compile that could not run (HostCompileUnavailable, LanguagePatchedError) + is no compile error at all. + """ + if error is None or host_compile_unavailable(error) is not None: + return False + from triton.compiler.errors import ( + CompilationError, + CompileTimeAssertionFailure, + UnsupportedLanguageConstruct, + ) + + target_free = (CompileTimeAssertionFailure, UnsupportedLanguageConstruct) + seen: set[int] = set() + link: BaseException = error + while ( + isinstance(link, CompilationError) + and not isinstance(link, target_free) + and link.__cause__ is not None + and id(link) not in seen + ): + seen.add(id(link)) + link = link.__cause__ + return isinstance(link, target_free) and not target_queried(error) + + +def _for_target(failure: CompileFailure) -> str: + if failure.target is None: + return "" + return f" for {format_ir_target(failure.target)}" + + +def _error_chain(error: BaseException) -> list[BaseException]: + """``error`` and the errors behind it, outermost first, as Triton's code + generator nests them: a CompilationError raised from (``__cause__``), + or ``from None`` while handling (the suppressed ``__context__``), the + error of a called @jit helper, a builtin or the kernel's own code. Only + CompilationErrors are followed.""" + from triton.compiler.errors import CompilationError + + chain = [error] + link: BaseException | None = error + while isinstance(link, CompilationError): + link = link.__cause__ or ( + link.__context__ if link.__suppress_context__ else None + ) + if link is None or any(link is seen for seen in chain): + break + chain.append(link) + return chain + + +def _summary(error: BaseException) -> str: + """One line for a compile error: the innermost error's type and the + first line of its message (a CompilationError's own message, without + the source excerpt it formats around it).""" + from triton.compiler.errors import CompilationError + + inner = _error_chain(error)[-1] + if isinstance(inner, CompilationError): + text = str(getattr(inner, "error_message", None) or "") + else: + text = str(inner) + first = next((line.strip() for line in text.splitlines() if line.strip()), "") + return f"{type(inner).__name__}: {first}" if first else type(inner).__name__ + + +def _kernel_site(jit_fn: Any, error: BaseException | None = None) -> Any: + """Where in ``jit_fn``'s source file ``error`` was raised: the line of + the kernel's own code the code generator names for it (for an error in + a called @jit helper, the call), else the kernel's ``def`` line; None + for a JITFunction without source (e.g. a stand-in).""" + fn = getattr(jit_fn, "fn", None) + code = getattr(fn, "__code__", None) + start = getattr(jit_fn, "starting_line_number", None) + raw_src = getattr(jit_fn, "raw_src", None) + if code is None or not isinstance(start, int) or not raw_src: + return None + # The line of ``def`` after the decorators, as Triton counts it. + offset = next( + (i for i, line in enumerate(raw_src) if line.strip().startswith("def ")), 0 + ) + def_line = start + offset + src = getattr(jit_fn, "src", None) + for link in [] if error is None else _error_chain(error): + lineno = getattr(getattr(link, "node", None), "lineno", None) + if src is not None and getattr(link, "src", None) == src and lineno: + # The code generator's lines count from the ``def`` line. + return SourceLocation(code.co_filename, def_line + lineno - 1) + return SourceLocation(code.co_filename, def_line) + + +def _at(site: Any) -> str: + return "" if site is None else f"{site.file}:{site.line}: " + + +def _failure_refusal(failure: CompileFailure) -> Refusal: + """A failed config's refusal: why its host compile failed (see + _compiles_nowhere), where, and how to check it for another target where + another target may compile it.""" + error = failure.error + site = _kernel_site(failure.jit_fn, error) + unavailable = None if error is None else host_compile_unavailable(error) + if unavailable is not None: + return Refusal( + SanitizerKind.HOST_COMPILE_UNAVAILABLE, + f"{_at(site)}{unavailable}", + loc=site, + ) + if isinstance(error, LanguagePatchedError): + # No compile ran, so it says nothing about the kernel or the target. + return Refusal( + SanitizerKind.HOST_COMPILE_UNAVAILABLE, f"{_at(site)}{error}", loc=site + ) + if error is not None and bind_failed(error): + # The call's own error: the untraced call raises it on any GPU, and + # a traced launch does too; only an event built outside the + # core gets here. + return Refusal( + SanitizerKind.COMPILE_FAILED, + f"{_at(site)}the call does not bind to the kernel's parameters " + f"({_summary(error)}), so it was not checked; Triton raises this " + "error for the call whatever the GPU", + loc=site, + ) + why = "" if error is None else f" ({_summary(error)})" + names = unknown_options(error) + if names: + # No kernel code failed: the call names options the target's backend + # lacks. Only a target of a backend that has them can compile it. + listed = ", ".join(repr(name) for name in names) + return Refusal( + SanitizerKind.COMPILE_FAILED, + f"{_at(site)}it failed to compile{_for_target(failure)}{why}: the " + f"call passes {listed}, neither a parameter of the kernel nor a " + f"compile option{_for_target(failure)}, so it was not checked; " + "Triton raises this for the call on every GPU whose backend lacks " + "that option (a misspelled option: on every GPU); if it is another " + "backend's option, name a target of that backend to check it " + f"({_NAME_A_TARGET})", + loc=site, + ) + return Refusal( + SanitizerKind.COMPILE_FAILED, + f"{_at(site)}it failed to compile{_for_target(failure)}{why}, so it was " + "not checked; a kernel can compile for one target and fail for another, " + "so it may launch on a GPU of another kind: to check it, name a target " + f"it compiles for ({_NAME_A_TARGET})", + loc=site, + ) + + +def _nowhere_note(failure: CompileFailure) -> str: + config = f"config {_saveable(failure.config)}" if failure.config else "the kernel" + error = failure.error + why = "" if error is None else f" ({_summary(error)})" + return ( + f"{config} was not checked: {_at(_kernel_site(failure.jit_fn, error))}it " + f"failed to compile{_for_target(failure)}{why}, an error of its own code " + "whatever the target, so it never launches" + ) + + +def _none_compiled(log: ArtifactLog) -> Refusal: + targets = sorted( + {format_ir_target(f.target) for f in log.failures if f.target is not None} + ) + where = f" for {', '.join(targets)}" if targets else "" + site = _kernel_site(log.failures[0].jit_fn) if log.failures else None + return Refusal( + SanitizerKind.COMPILE_FAILED, + f"{_at(site)}no config of the launch compiled{where}, so nothing was " + "checked: each failed with an error of its own code whatever the target " + "(see the notes)", + loc=site, + ) + + +_PLAIN = (str, int, float, bool, type(None)) + + +def _tensor_of(value: Any) -> TensorFacts | None: + """``value``'s tensor facts, None if it is no (readable) tensor.""" + if not hasattr(value, "data_ptr"): + return None + try: + return tensor_facts(value) + except Exception: + return None + + +def _saveable(config: Mapping[str, Any]) -> dict[str, Any]: + """The config kwargs, each value a saved trace cannot hold put as text: + a tensor (e.g. a heuristic's view) described by its facts, never its + data, anything else (e.g. a heuristic's tl.dtype) as its repr.""" + + def describe(value: Any) -> Any: + if isinstance(value, _PLAIN): + return value + facts = _tensor_of(value) + if facts is None: + return repr(value) + return ( + f"" + ) + + return {name: describe(value) for name, value in config.items()} + + +def _config_key(config: Mapping[str, Any]) -> Hashable: + """What tells two bindings' config kwargs apart: a plain value by type + and value, a tuple item by item, a tensor by its data_ptr, shape, + strides and dtype (never its data), any other hashable value by type + and equality, anything else by identity (the bindings hold the values + while the keys are compared).""" + + def key(value: Any) -> Hashable: + if isinstance(value, float): + return (float, value.hex()) # -0.0 apart from 0.0, NaN equal + if isinstance(value, _PLAIN): + return (type(value), value) + if isinstance(value, tuple): + return (type(value), tuple(key(item) for item in value)) + facts = _tensor_of(value) + if facts is not None: + return ("tensor", facts.data_ptr, facts.shape, facts.strides, facts.dtype) + try: + hash(value) + except Exception: + return ("id", id(value)) + return (type(value), value) + + return tuple(sorted((name, key(value)) for name, value in config.items())) + + +def _configs_of( + spec: CompiledSpecialization, +) -> list[tuple[dict[str, Any], list[LaunchBinding]]]: + """``spec``'s bindings grouped by their config kwargs (see _config_key), + each group with its saveable config, in the order first seen, so the + first group's config is ``spec.config``. Configs that compile to one + kernel (e.g. differing only in a runtime int kwarg) share a + specialization but not their bindings.""" + groups: dict[Hashable, tuple[dict[str, Any], list[LaunchBinding]]] = {} + for binding in spec.bindings: + key = _config_key(binding.config) + group = groups.get(key) + if group is None: + groups[key] = (_saveable(binding.config), [binding]) + else: + group[1].append(binding) + return list(groups.values()) or [(_saveable(spec.config), [])] + + +def _record( + finding: Finding, kernel_name: str, binding: LaunchBinding, config: dict[str, Any] +) -> CompiledSanitizerRecord: + loc = finding.loc + tracebacks = ( + [] + if loc is None + else [location_to_traceback_info((loc.file, loc.line, kernel_name))] + ) + detail = finding.detail + if loc is None and finding.line_no is not None: + detail = f"{detail} (TTIR line {finding.line_no})" + return CompiledSanitizerRecord( + kind=finding.kind, + # An atomic reads and writes; it is reported as the write. + op_type=Load if finding.access_kind == "load" else Store, + tensor_name=finding.base_param, + tensor_facts=binding.tensors.get(finding.base_param), + witness=dict(finding.witness), + config=config, + user_code_tracebacks=tracebacks, + violation_offset=finding.violation_offset, + violation_address=finding.violation_address, + detail=detail, + ) + + +# ─────────────────────────── reporting ─────────────────────────── + +_TITLES = { + "out-of-bounds": "Out-Of-Bounds Access Detected", + "integer-overflow": "Integer Width Overflow Detected", + "division-by-zero": "Division By Zero Detected", +} + + +def print_compiled_record(record: CompiledSanitizerRecord) -> None: + """Print one compiled sanitizer finding, in the eager report's layout.""" + rule = "=" * 60 + print(rule) + print(f"{_TITLES[record.kind]:^60}".rstrip()) + print(f"{'(compiled sanitizer)':^60}".rstrip()) + print(rule) + print(f"Operation: {record.op_type.__name__}") + print(f"Tensor Arg: {record.tensor_name}") + facts = record.tensor_facts + if facts is not None: + print( + f"Tensor Info: dtype={facts.dtype}, shape={facts.shape}, " + f"strides={facts.strides}, contiguous={facts.contiguous}" + ) + print(f"Tensor base memory address: {facts.data_ptr:#x}") + if record.config: + print(f"Config: {record.config}") + for tb in record.user_code_tracebacks: + print(f"File: {tb.filename}, Line: {tb.lineno}, in {tb.func_name}") + print(f" Code: {tb.line_of_code.strip()}") + print("-" * 60) + if record.violation_address is not None: + print( + f"Invalid access detected at address: {record.violation_address:#x} " + f"(element offset {record.violation_offset})" + ) + witness = ", ".join(f"{name}={value}" for name, value in record.witness.items()) + print(f"Witness: {witness}") + if record.detail: + print(f"Detail: {record.detail}") + print(rule) + + +def print_unchecked(verdict: IRVerdict, printed: set[Hashable]) -> None: + """Print one line per part of a launch the compiled sanitizer did not + check (each refused config, else the launch's own refusal) that is not + in ``printed`` yet, and add it there: once per specialization, kind and + op, whatever launch-specific numbers the message holds; a config that + compiled to nothing (a failed compile) once per config and message. + The launch's own refusal is followed by its notes, each once: the + refusal of a launch no config of which compiled says why in them.""" + refused = [ + (c.specialization, c.config, c.refusal) + for c in verdict.per_config + if c.refusal is not None + ] + notes: tuple[str, ...] = () + if not refused and verdict.refusal is not None: + refused = [(None, {}, verdict.refusal)] + notes = verdict.notes + for specialization, config, refusal in refused: + ident: Hashable = specialization + if specialization is None: + ident = ( + tuple(sorted((name, repr(value)) for name, value in config.items())), + refusal.message, + ) + key = (ident, refusal.kind, refusal.line_no, refusal.loc) + if key in printed: + continue + printed.add(key) + where = f" (config {config})" if config else "" + print( + f"[CompiledSanitizer] not checked{where}: " + f"{refusal.kind}: {refusal.message}" + ) + for note in notes: + if ("note", note) in printed: + continue + printed.add(("note", note)) + print(f"[CompiledSanitizer] note: {note}") + + +# A sanitizer mode like the eager one: isinstance(..., Sanitizer) holds, so +# trace()'s ENABLE_SANITIZER=0 escape hatch covers a CompiledSanitizer() too. +Sanitizer.register(CompiledSanitizer) diff --git a/tilelens/clients/sanitizer/compiled/oob.py b/tilelens/clients/sanitizer/compiled/oob.py new file mode 100644 index 000000000..1a8fdae61 --- /dev/null +++ b/tilelens/clients/sanitizer/compiled/oob.py @@ -0,0 +1,1248 @@ +"""The compiled sanitizer's per-launch checks over an AccessGraph. + +Given a kernel's ``AccessGraph`` (``tilelens.ir.ttir_reader``) and one +launch's ``LaunchBinding``, each access becomes a few Z3 queries over the +launch's free variables (program ids, the access's lane positions, the +loop's iteration index), with its scalar arguments and grid substituted as +constants: + +* ``out-of-bounds``: the access executes (its loop runs, its path and + mask hold) at an element offset outside the view footprint of its + tensor, the eager sanitizer's legal set: the offsets + ``sum(i_d * stride_d)`` with ``0 <= i_d < size_d``, a stride-0 dim + counting once, relative to the view's ``data_ptr``; +* ``integer-overflow``: a width obligation the access depends on can + fail, so the IR's fixed-width arithmetic is not the unbounded reading; +* ``division-by-zero``: a ``//`` or ``%`` the access reaches (from its + offset, mask or path, or from its loop's bounds) can divide by zero. + +The unbounded reading of a term is exact only where the obligations below +it hold and its divisors are non-zero. The reader's roles order that +dependence: the loop's bounds, then its increment (only when the loop +runs), then the access's path, its mask (under the path) and its offset +(under path and mask). Each role is discharged by ONE query, "some +obligation or division of this role fails", under the safety of the roles +before it and never under its own (a wrap can decide a sibling divisor, and +the reverse); the out-of-bounds query assumes every role's safety. A +wrapped offset is therefore an integer overflow, never an out-of-bounds +witness that only the unbounded reading has. + +A value in an arm of a ``Select`` (``tl.where``) matters only where the +select picks that arm, so the width obligations of the access's path, mask +and offset are discharged under the conditions of the arms their terms are +read through: a wrap the select discards is harmless (``arith`` ops wrap, +with defined results). Divisions are not: the IR computes both arms, and a +zero divisor (or ``INT_MIN // -1``) in the discarded one is still +undefined, so they are checked wherever the access evaluates them. (The +loop's own obligations are always checked unguarded.) + +Lanes: every tensor one access's terms combine has the access's shape +(TTIR broadcasts explicitly, and only size-1 dims), so all aranges along +one dim index the same position there: an arange's value is its start plus +the lane position of its dim (one position per dim and extent; an extent-1 +arange broadcast along a longer dim stays at position 0). + +Findings are exact: a SAT model is a concrete launch state, and the op a +role reports is the innermost one failing in the model (so one wrap is not +reported at the ops computed from it). A model that may be unreachable is +withheld and the access *abstains* instead (as in #361): a +branch condition the reader could not model (``guarded``) or one reading an +atomic observation, a mask it dropped (``mask_dropped``) or one reading an +observation. Z3's ``unknown`` abstains too, each query on a fresh +Solver with its own timeout, so whether a hard query ends in +``solver-unknown`` can depend on the machine's load. An access the model +cannot check at all is refused (an unbound argument, a non-positive loop +step, an address that depends on an atomic observation). Abstentions never +suppress another access's findings: a ``CheckResult`` carries both, and +only one with neither is a proof for the launch. + +Each check builds its Z3 terms in a Z3 context of its own, so checks may +run on several host threads at once (one Z3 context is not thread-safe). A +Ctrl+C that Z3 caught during a query is raised as ``KeyboardInterrupt``, +never taken for an unknown. + +Terms can be deeper than Python's recursion limit and their generated +``==`` / ``hash`` recurse, so lowering is iterative and memoized by term +identity. +""" + +from __future__ import annotations + +from collections.abc import Hashable, Iterable, Iterator, Mapping, Sequence +from dataclasses import dataclass, replace +from enum import Enum +from typing import Any, Literal + +from z3 import ( + And, + ArithRef, + BoolRef, + BoolVal, + Context, + Exists, + If, + Implies, + Int, + IntVal, + ModelRef, + Or, + Solver, + Sum, + is_bool, + is_false, + is_int_value, + sat, + simplify, + unknown, + unsat, +) +from z3 import Not as Z3Not + +from ....ir.launch import LaunchBinding, TensorFacts +from ....ir.ttir_reader import ( + AccessEvent, + AccessGraph, + Arange, + Bin, + BoolBin, + Cmp, + Const, + DataDep, + IntCast, + IterArgOffset, + LoopInfo, + LoopVar, + Not, + NumPrograms, + Observed, + Param, + Pid, + Select, + WidthObligation, + mentions_observed, + width_obligations, +) + +# The one import site of the source-location type findings carry. +from ....ir.verdict import Refusal, SourceLocation +from ..data import CompiledFindingKind + + +class SanitizerKind(str, Enum): + """Why the compiled sanitizer left (part of) a launch undecided. The + sanitizer's own kinds, apart from the reader's ``TTIRKind``.""" + + # A kernel argument the access (or its loop) reads has no launch + # binding, or the launch grid is unknown (never left unconstrained). + MISSING_BINDING = "missing-binding" + # The loop step is not positive (or may not be) for this launch. + NON_POSITIVE_STEP = "non-positive-step" + # Z3 answered unknown (a timeout included). + SOLVER_UNKNOWN = "solver-unknown" + # A witness sits under a branch condition that is not modeled (loaded + # data, or an atomic observation read as a free value). + UNMODELABLE_CONDITION = "unmodelable-condition" + # A witness sits behind a mask that is not modeled (dropped as loaded + # data, or reading an atomic observation). + DATA_DEPENDENT_MASK = "data-dependent-mask" + # The address depends on the value an atomic observed. + OBSERVATION_IN_ADDRESS = "observation-in-address" + # A term the reader marked unmodelable (DataDep) reached a lowered + # address, mask, path or loop bound; the reader never lets it. + UNMODELED_VALUE = "unmodeled-value" + # The client's own (compiled/client.py): the launch compiled nothing to + # read (no JITFunction, or none of its configs was delivered); a config + # failed to compile for the IR target with an error that need not hold + # for another target (it may launch on the user's GPU), or its call + # does not bind the kernel's parameters (only in an event built outside + # the core, which raises such a call), or no config compiled at all; a + # compiled kernel holds no TTIR; the host compile could not run for a + # config (the installed Triton's compile API, a compile that asked for a + # device, or one refused while an interpreted launch has the language + # patched: no error of the kernel's, so the config may well launch); the + # analysis itself failed (a bug, contained). + NO_COMPILED_KERNEL = "no-compiled-kernel" + COMPILE_FAILED = "compile-failed" + NO_TTIR = "no-ttir" + HOST_COMPILE_UNAVAILABLE = "host-compile-unavailable" + INTERNAL_ERROR = "internal-error" + + def __str__(self) -> str: + return self.value + + def __format__(self, spec: str) -> str: + return format(self.value, spec) + + +@dataclass(frozen=True) +class Finding: + """One exact finding: its witness is a concrete state of the launch.""" + + kind: CompiledFindingKind + access_index: int # into graph.accesses + access_kind: str # "load" | "store" | "atomic_rmw" | "atomic_cas" + base_param: str + # The op the finding is about: the access (out-of-bounds), the op or + # cast whose width does not hold (integer-overflow), the division + # (division-by-zero); the access's when that op's site is unknown. + line_no: int | None + loc: SourceLocation | None + # The free variables by name: pid_; arange__, the + # value of that tl.arange at the witness lane (with _d for a dim + # of a 2-D or wider tile); iter_loop, the loop's 0-based iteration. An + # integer-overflow adds "value", the out-of-range value. + witness: Mapping[str, int] + # out-of-bounds only: the element offset (in the accessed pointer's + # elements) from the view's data_ptr, and its byte address. + violation_offset: int | None = None + violation_address: int | None = None + detail: str = "" + + __hash__ = None # type: ignore[assignment] + + def __post_init__(self) -> None: + object.__setattr__(self, "witness", dict(self.witness)) + + +@dataclass(frozen=True) +class CheckResult: + """What ``check_graph`` decided for one launch: every exact finding, + and every undecided (access index, kind), whose first one is + ``refusal``. Only a result with neither proves the launch in bounds.""" + + findings: tuple[Finding, ...] + refusal: Refusal | None + abstained: tuple[tuple[int, SanitizerKind], ...] + + __hash__ = None # type: ignore[assignment] + + +def check_graph( + graph: AccessGraph, binding: LaunchBinding, *, timeout_ms: int = 10_000 +) -> CheckResult: + """Check every access of ``graph`` under ``binding`` (see the module + docstring); ``timeout_ms`` (a positive int) bounds each Z3 query. + + A limit of the model is returned as an abstention, never raised; an + exception means a bug (e.g. a graph that breaks the reader's + invariants), or is the KeyboardInterrupt of a Ctrl+C.""" + _check_timeout(timeout_ms) + return _Check(graph, binding, timeout_ms).run() + + +def launch_key(binding: LaunchBinding) -> Hashable: + """Everything ``check_graph`` reads of ``binding`` but the tensors' + addresses: for one graph (and timeout), bindings with one key get the + same CheckResult up to the findings' ``violation_address`` (see + ``readdressed``); no message holds an address.""" + return ( + tuple(sorted(binding.params.items())), + binding.grid, + tuple( + sorted( + (name, (f.elem_size, f.numel, f.shape, f.strides, f.contiguous)) + for name, f in binding.tensors.items() + ) + ), + binding.error, + ) + + +def readdressed( + result: CheckResult, graph: AccessGraph, binding: LaunchBinding +) -> CheckResult: + """``result``, the check of ``graph`` under a binding with the + ``launch_key`` of ``binding``, with its findings' byte addresses in + ``binding``'s tensors.""" + findings = tuple( + finding + if finding.violation_offset is None + else replace( + finding, + violation_address=binding.tensors[finding.base_param].data_ptr + + finding.violation_offset + * _pointee_bytes(graph.accesses[finding.access_index].elem_bits), + ) + for finding in result.findings + ) + return CheckResult(findings, result.refusal, result.abstained) + + +def describe_site(line_no: int | None, loc: Any) -> str: + """Where a refused or reported op is: ``file:line`` of its user-source + location, else its TTIR line.""" + if loc is not None: + return f"{loc.file}:{loc.line}" + return f"TTIR line {line_no}" if line_no is not None else "the kernel" + + +def _check_timeout(timeout_ms: Any) -> None: + # Z3 reads a timeout of 0 or below as none at all. + if ( + isinstance(timeout_ms, bool) + or not isinstance(timeout_ms, int) + or timeout_ms <= 0 + ): + raise ValueError(f"timeout_ms must be a positive int, not {timeout_ms!r}") + + +# ─────────────────────────── lowering ─────────────────────────── + + +class _Refused(Exception): + """An access (or the loop) the model cannot check: becomes an + abstention; never escapes check_graph.""" + + def __init__( + self, + kind: SanitizerKind, + message: str, + line_no: int | None = None, + loc: Any = None, + ) -> None: + super().__init__(message) + self.kind = kind + self.message = message + self.line_no = line_no + self.loc = loc + + def at(self, line_no: int | None, loc: Any) -> _Refused: + if self.line_no is None and self.loc is None: + self.line_no, self.loc = line_no, loc + return self + + +def _as_bool(e: Any) -> BoolRef: + """An i1 value in a boolean position: i1 constants (e.g. the dense + mask of an unmasked atomic, Const(1)) lower to Int.""" + return e if is_bool(e) else e != 0 + + +def _as_int(e: Any) -> ArithRef: + """An i1 value in an integer position (an extui of a compare, ...).""" + return If(e, IntVal(1, e.ctx), IntVal(0, e.ctx)) if is_bool(e) else e + + +def _trunc_div(a: ArithRef, b: ArithRef) -> ArithRef: + """arith.divsi rounds toward zero, but Z3's Int ``/`` is Euclidean + (floor for a positive divisor): they disagree on negative dividends. + Divide the magnitudes, where the two agree, and re-apply the sign.""" + aa = If(a >= 0, a, -a) + ab = If(b >= 0, b, -b) + q = aa / ab + return If((a >= 0) == (b >= 0), q, -q) + + +def _bin(op: str, a: ArithRef, b: ArithRef) -> ArithRef: + # The unsigned ops read their operands unsigned; their width + # obligations make both non-negative, where they equal the signed ones. + if op == "+": + return a + b + if op == "-": + return a - b + if op == "*": + return a * b + if op in ("//", "u//"): + return _trunc_div(a, b) + if op in ("%", "u%"): + # arith.remsi: the remainder carries the dividend's sign + return a - b * _trunc_div(a, b) + if op in ("min", "umin"): + return If(a <= b, a, b) + if op in ("max", "umax"): + return If(a >= b, a, b) + raise ValueError(f"unknown integer op {op!r}") + + +# Unsigned predicates read their operands unsigned; their width obligations +# make both non-negative, where they equal the signed ones. +_SIGNED_PRED = {"ult": "slt", "ule": "sle", "ugt": "sgt", "uge": "sge"} + + +def _cmp(pred: str, a: Any, b: Any) -> BoolRef: + if is_bool(a) or is_bool(b): # i1 operands, as 0/1 + a, b = _as_int(a), _as_int(b) + table = { + "slt": a < b, "sle": a <= b, "sgt": a > b, + "sge": a >= b, "eq": a == b, "ne": a != b, + } # fmt: skip + try: + return table[_SIGNED_PRED.get(pred, pred)] + except KeyError: + raise ValueError(f"unknown cmpi predicate {pred!r}") from None + + +def _kids(t: object, graph: AccessGraph) -> tuple: + """The terms ``t`` is computed from; a loop-carried pointer's offset + is computed from its IterArgInfo's ``offset0`` and ``delta``.""" + if isinstance(t, (Bin, Cmp, BoolBin)): + return (t.a, t.b) + if isinstance(t, Select): + return (t.cond, t.t, t.f) + if isinstance(t, Not): + return (t.a,) + if isinstance(t, IntCast): + return (t.x,) + if isinstance(t, IterArgOffset): + info = graph.iter_args[t.arg_id] + return (info.offset0, info.delta) + return () + + +def _walk( + roots: Iterable[object], graph: AccessGraph, seen: set[int] +) -> Iterator[object]: + """The nodes reachable from ``roots`` in pre-order, skipping (and + adding to) ``seen`` by identity. Iterative.""" + stack = [r for r in reversed(list(roots)) if r is not None] + while stack: + t = stack.pop() + if id(t) in seen: + continue + seen.add(id(t)) + yield t + stack.extend(reversed(_kids(t, graph))) + + +_DIVISIONS = frozenset({"//", "%", "u//", "u%"}) + + +def _divisions( + roots: Iterable[object], graph: AccessGraph, seen: set[int] +) -> list[Bin]: + return [ + n + for n in _walk(roots, graph, seen) + if isinstance(n, Bin) and n.op in _DIVISIONS + ] + + +def _signed(value: int, bits: int) -> int: + """The IR's signed reading of an argument of ``bits`` (an i1 stays 0/1, + the boolean model).""" + if bits <= 1: + return value + half = 1 << (bits - 1) + return (value + half) % (1 << bits) - half + + +def _lane_name(t: Arange) -> str: + # Named like the eager engine's arange variables; a tile's dim added + # (a 1-D range has dim -1). + name = f"arange_{t.start}_{t.end}" + return name if t.dim < 0 else f"{name}_d{t.dim}" + + +class _Env: + """The Z3 variables of one family of queries (one access, or the + loop's own checks), their range premises, and the lowering memo.""" + + def __init__( + self, + graph: AccessGraph, + binding: LaunchBinding, + grid: tuple[int, int, int], + ctx: Context, + ) -> None: + self.graph = graph + self.binding = binding + self.grid = grid + self.ctx = ctx + # Range premises of the variables created so far. + self.premises: list[BoolRef] = [] + self.pids = tuple(Int(f"pid_{axis}", ctx) for axis in range(3)) + for pid, size in zip(self.pids, grid): + self.premises += [pid >= 0, pid < size] + # (dim, extent) -> the lane's position along that dim (see lane). + self.positions: dict[tuple[int, int], ArithRef] = {} + # witness name -> the arange's value at the lane + self.lanes: dict[str, ArithRef] = {} + self.iteration: ArithRef | None = None + self._observed: dict[int, ArithRef] = {} + # id(term) -> (term, lowered): holding the term keeps its id unique. + self._memo: dict[int, tuple[object, Any]] = {} + + # ── leaves ── + + def param(self, name: str) -> int: + try: + value = self.binding.params[name] + except KeyError: + raise _Refused( + SanitizerKind.MISSING_BINDING, + f"scalar argument {name!r} has no launch binding" + + _binding_error(self.binding), + ) from None + arg = self.graph.arg(name) + return _signed(value, arg.int_bits if arg is not None else 0) + + def lane(self, t: Arange) -> ArithRef: + """``t``'s value at the access's lane: its start plus the lane's + position along its dim, one position per dim and extent (see the + module docstring).""" + extent = t.end - t.start + key = (t.dim, extent) + pos = self.positions.get(key) + if pos is None: + pos = Int(f"lane_d{t.dim}_n{extent}", self.ctx) + self.positions[key] = pos + self.premises += [pos >= 0, pos < extent] + value = pos + t.start + self.lanes.setdefault(_lane_name(t), value) + return value + + def observed(self, index: int) -> ArithRef: + """An atomic observation: a free value of the atomic's width.""" + v = self._observed.get(index) + if v is None: + v = Int(f"observed_{index}", self.ctx) + self._observed[index] = v + bits = self.graph.accesses[index].elem_bits + if bits > 1: + half = 1 << (bits - 1) + self.premises += [v >= -half, v < half] + return v + + def bounds(self) -> tuple[ArithRef, ArithRef, ArithRef]: + loop = self._loop() + return self.value(loop.lower), self.value(loop.upper), self.value(loop.step) + + def loop_iteration(self) -> ArithRef: + """The loop's 0-based iteration index ``k``, with its premise ``k >= + 0 and lower + k*step < upper``: only iterations that run, and none + when the launch's trip count is zero.""" + if self.iteration is None: + loop = self._loop() + lower, upper, step = self.bounds() + k = Int(f"iter_{loop.loop_ssa.strip('%')}", self.ctx) + self.premises += [k >= 0, lower + k * step < upper] + self.iteration = k + return self.iteration + + def _loop(self) -> LoopInfo: + loop = self.graph.loop + if loop is None: + raise ValueError( + f"kernel {self.graph.kernel_name!r}: a loop term without a loop" + ) + return loop + + # ── terms ── + + def value(self, term: object) -> ArithRef: + return _as_int(self.lower(term)) + + def cond(self, term: object) -> BoolRef: + return _as_bool(self.lower(term)) + + def lower(self, root: object) -> Any: + """``root`` as a Z3 expression (Int, or Bool for a compare).""" + memo = self._memo + stack: list[tuple[object, bool]] = [(root, False)] + while stack: + t, ready = stack.pop() + if id(t) in memo: + continue + kids = _kids(t, self.graph) + if kids and not ready: + stack.append((t, True)) + stack.extend((k, False) for k in reversed(kids) if id(k) not in memo) + continue + memo[id(t)] = (t, self._apply(t, [memo[id(k)][1] for k in kids])) + return memo[id(root)][1] + + def _apply(self, t: object, kids: Sequence[Any]) -> Any: + if isinstance(t, Const): + return IntVal(t.value, self.ctx) + if isinstance(t, Param): + return IntVal(self.param(t.name), self.ctx) + if isinstance(t, Pid): + return self.pids[t.axis] + if isinstance(t, NumPrograms): + return IntVal(self.grid[t.axis], self.ctx) + if isinstance(t, Arange): + return self.lane(t) + if isinstance(t, LoopVar): + lower, _upper, step = self.bounds() + return lower + self.loop_iteration() * step + if isinstance(t, IterArgOffset): + return _as_int(kids[0]) + self.loop_iteration() * _as_int(kids[1]) + if isinstance(t, Bin): + return _bin(t.op, _as_int(kids[0]), _as_int(kids[1])) + if isinstance(t, Cmp): + return _cmp(t.pred, kids[0], kids[1]) + if isinstance(t, BoolBin): + a, b = _as_bool(kids[0]), _as_bool(kids[1]) + return And(a, b) if t.op == "and" else Or(a, b) + if isinstance(t, Select): + a, b = kids[1], kids[2] + if is_bool(a) != is_bool(b): + a, b = _as_int(a), _as_int(b) + return If(_as_bool(kids[0]), a, b) + if isinstance(t, Not): + return Z3Not(_as_bool(kids[0])) + if isinstance(t, IntCast): + # Its value is the operand's while its width obligation holds. + return _as_int(kids[0]) + if isinstance(t, Observed): + return self.observed(t.access_index) + if isinstance(t, DataDep): + raise _Refused( + SanitizerKind.UNMODELED_VALUE, f"an unmodeled value ({t.why})" + ) + raise TypeError(f"unknown term {type(t).__name__}") + + +class _Dag: + """The nodes one family of queries reads from ``roots``: each one's + rank (children before parents) and, when built with the family's env + (whose memo holds the lowered roots), its guard: the condition under + which its value reaches a root through the arms of Selects (a node + without one reaches a root directly). Iterative, keyed by identity.""" + + def __init__( + self, roots: Sequence[object], graph: AccessGraph, env: _Env | None = None + ) -> None: + self.graph = graph + self.order: dict[int, int] = {} + nodes: list[object] = [] + stack = [(r, False) for r in reversed(roots) if r is not None] + while stack: + t, done = stack.pop() + if id(t) in self.order: + continue + if done: + self.order[id(t)] = len(nodes) + nodes.append(t) + continue + stack.append((t, True)) + stack.extend( + (k, False) for k in reversed(_kids(t, graph)) if id(k) not in self.order + ) + self.guards: dict[int, BoolRef] = {} + if env is None: + return + direct = {id(r) for r in roots if r is not None} + shares: dict[int, list[BoolRef]] = {} + for t in reversed(nodes): # every parent before its children + parts = shares.pop(id(t), []) + guard = None + if id(t) not in direct: + guard = parts[0] if len(parts) == 1 else Or(parts) + self.guards[id(t)] = guard + edges: list[tuple[object, BoolRef | None]] + if isinstance(t, Select): + c = env.cond(t.cond) + edges = [(t.cond, guard)] + for arm, taken in ((t.t, c), (t.f, Z3Not(c))): + edges.append((arm, taken if guard is None else And(guard, taken))) + else: + edges = [(k, guard) for k in _kids(t, graph)] + for kid, kid_guard in edges: + if kid_guard is None: + direct.add(id(kid)) + elif id(kid) not in direct: + shares.setdefault(id(kid), []).append(kid_guard) + + def rank(self, term: object) -> float: + """Children first, and a Select's condition before its arms; a term + the walk does not hold (an obligation's own, e.g. a remainder's + quotient) right after its children.""" + index = self.order.get(id(term)) + if index is not None: + return index + kids = [ + self.order[id(k)] for k in _kids(term, self.graph) if id(k) in self.order + ] + return max(kids) + 0.5 if kids else len(self.order) + + +# ─────────────────────────── view footprint ─────────────────────────── + + +def _in_view(e: ArithRef, facts: TensorFacts) -> BoolRef: + """``e`` (an element index from the view's data_ptr) is the offset of one + of the view's elements: ``sum(i_d * stride_d)``, ``0 <= i_d < size_d``.""" + if facts.numel == 0: + return BoolVal(False, e.ctx) + if facts.contiguous: + return And(e >= 0, e < facts.numel) + # Size-1 dims add nothing and stride-0 dims alias one element: drop both. + # (stride, size), largest stride first + dims = sorted( + ((st, sz) for sz, st in zip(facts.shape, facts.strides) if sz != 1 and st != 0), + reverse=True, + ) + if any(st < 0 for st, _ in dims): + return _any_index(e, dims) + extent = 1 # dense (a permutation of a contiguous layout): an interval + for st, sz in reversed(dims): + if st != extent: + break + extent *= sz + else: + return And(e >= 0, e < extent) + # Without overlap (each stride exceeds the reach of the smaller ones), + # the indices are the greedy quotients: quantifier-free and exact. + reach = 0 + for st, sz in reversed(dims): + if st <= reach: + return _any_index(e, dims) + reach += (sz - 1) * st + conds = [e >= 0] + rest = e + for st, sz in dims: + conds.append(rest / st < sz) + rest = rest % st + conds.append(rest == 0) + return And(conds) + + +def _any_index(e: ArithRef, dims: list[tuple[int, int]]) -> BoolRef: + """The stride equation itself, for overlapping or negative strides (as + a quantifier, Z3 may answer unknown).""" + idx = [Int(f"view_index_{d}", e.ctx) for d in range(len(dims))] + body = [i >= 0 for i in idx] + [i < sz for i, (_, sz) in zip(idx, dims)] + body.append(e == Sum([i * st for i, (st, _) in zip(idx, dims)])) + return Exists(idx, And(body)) + + +def _legal(offset: ArithRef, facts: TensorFacts, width: int) -> BoolRef: + """An access of ``width`` bytes at element ``offset`` touches only bytes + of the view's elements.""" + if width == facts.elem_size: + return _in_view(offset, facts) + # A pointer whose element width differs from the tensor's (a + # reinterpreting view): every byte it touches must be in an element. + first = offset * width + return And( + [ + And(first + j >= 0, _in_view((first + j) / facts.elem_size, facts)) + for j in range(width) + ] + ) + + +def _footprint(facts: TensorFacts) -> str: + if facts.numel == 0: + return "the empty tensor" + if facts.contiguous: + return f"the tensor's {facts.numel} elements" + return f"the view of shape {facts.shape} and strides {facts.strides}" + + +def _unusable(facts: TensorFacts) -> str | None: + if facts.elem_size <= 0 or facts.numel < 0: + return f"element size {facts.elem_size}, numel {facts.numel}" + if len(facts.shape) != len(facts.strides): + return f"shape {facts.shape} with strides {facts.strides}" + return None + + +# ─────────────────────────── the checks ─────────────────────────── + + +def _fits(value: ArithRef, bits: int, signed: bool) -> BoolRef: + if signed: + half = 1 << (bits - 1) + return And(value >= -half, value < half) + return And(value >= 0, value < (1 << bits)) + + +def _undefined_when_wide(ob: WidthObligation) -> bool: + """A signed quotient that does not fit (``INT_MIN // -1``, a + remainder's quotient included) is undefined in the IR, not a wrap: like + a zero divisor, it counts in either arm of a Select.""" + return isinstance(ob.term, Bin) and ob.term.op == "//" + + +def _is_increment(ob: WidthObligation, loop: LoopInfo, bound_ids: set[int]) -> bool: + """The loop's increment obligation (``upper - 1 + step``, which + width_obligations builds from the bound nodes themselves): the one loop + obligation that holds only when the loop runs.""" + t = ob.term + return ( + id(t) not in bound_ids + and isinstance(t, Bin) + and t.op == "+" + and t.b is loop.step + and isinstance(t.a, Bin) + and t.a.op == "-" + and t.a.a is loop.upper + ) + + +def _location(loc: Any) -> SourceLocation | None: + if loc is None or isinstance(loc, SourceLocation): + return loc + return SourceLocation(loc.file, loc.line, getattr(loc, "col", None)) + + +def _binding_error(binding: LaunchBinding) -> str: + return f" (unreadable: {binding.error})" if binding.error else "" + + +def _pointee_bytes(elem_bits: int) -> int: + return max(1, (elem_bits + 7) // 8) + + +_WITHHELD = { + SanitizerKind.UNMODELABLE_CONDITION: "under a branch condition that is not " + "modeled (loaded data, or an atomic observation read as a free value)", + SanitizerKind.DATA_DEPENDENT_MASK: "behind a mask that is not modeled " + "(loaded data, or an atomic observation read as a free value)", +} + +# What Z3 answers for a query its Ctrl+C handler cancelled (Python's own +# handler does not run during the query). +_INTERRUPTED = "interrupted from keyboard" + + +@dataclass(frozen=True, eq=False) +class _Condition: + """One width obligation or division of a role: ``safe`` holds where it + does not fail; the lowest ``rank`` is the innermost failure.""" + + kind: Literal["integer-overflow", "division-by-zero"] + site: WidthObligation | Bin + safe: BoolRef + rank: tuple + + +class _Check: + def __init__( + self, graph: AccessGraph, binding: LaunchBinding, timeout_ms: int + ) -> None: + self.graph = graph + self.binding = binding + self.timeout_ms = timeout_ms + # This check's own Z3 context (see the module docstring). + self.ctx = Context() + self.findings: list[Finding] = [] + self.abstained: list[tuple[int, SanitizerKind]] = [] + self.refusal: Refusal | None = None + # (finding kind, site) already reported: one finding per op site. + self.reported: set[tuple[str, object]] = set() + self._alive: list[object] = [] # reported sites keyed by id + # The loop's obligations and divisions, shared by its accesses. + self.loop_obs: tuple[WidthObligation, ...] = () + self.loop_divs: tuple[Bin, ...] = () + self.loop_ids: set[int] = set() + + def run(self) -> CheckResult: + graph, grid = self.graph, self.binding.grid + if grid is None: + for i, access in enumerate(graph.accesses): + self.abstain( + i, + _Refused( + SanitizerKind.MISSING_BINDING, + "the launch grid is unknown" + _binding_error(self.binding), + access.line_no, + access.loc, + ), + ) + return self.result() + in_loop = [i for i, a in enumerate(graph.accesses) if a.in_loop] + loop_refusal = None + if graph.loop is not None and in_loop: + loop_refusal = self.check_loop(in_loop[0], grid) + for i, access in enumerate(graph.accesses): + if access.in_loop and loop_refusal is not None: + self.abstain(i, loop_refusal) + continue + try: + self.check_access(i, access, grid) + except _Refused as r: + self.abstain(i, r.at(access.line_no, access.loc)) + return self.result() + + def result(self) -> CheckResult: + return CheckResult(tuple(self.findings), self.refusal, tuple(self.abstained)) + + # ── the loop, once for all of its accesses ── + + def check_loop(self, first: int, grid: tuple[int, int, int]) -> _Refused | None: + """Check the loop's step, bounds and increment; findings go to the + loop's ``first`` access. A refusal refuses every access in the loop. + No Select guards here: the loop's accesses assume its obligations + unguarded.""" + graph = self.graph + loop = graph.loop + assert loop is not None + env = _Env(graph, self.binding, grid, self.ctx) + try: + _lower, _upper, step = env.bounds() + refusal = self.step_refusal(env, step) + except _Refused as r: + refusal = r + if refusal is not None: + return refusal.at(loop.line_no, loop.loc) + roots = (loop.lower, loop.upper, loop.step) + dag = _Dag(roots, graph) + self.loop_ids = {id(n) for n in _walk(roots, graph, set())} + self.loop_divs = tuple(_divisions(roots, graph, set())) + self.loop_obs = tuple( + ob + for ob in width_obligations(graph, graph.accesses[first]) + if ob.role == "loop" + ) + increment = [ + ob for ob in self.loop_obs if _is_increment(ob, loop, self.loop_ids) + ] + bounds = [ + ob for ob in self.loop_obs if not _is_increment(ob, loop, self.loop_ids) + ] + # The bounds are computed whether or not the loop runs; the + # increment only matters when it does. + assumed = self.role(env, first, bounds, self.loop_divs, [], None, dag) + env.loop_iteration() + self.role(env, first, increment, (), assumed, None, dag) + return None + + def step_refusal(self, env: _Env, step: ArithRef) -> _Refused | None: + s = simplify(step) + if is_int_value(s): + if s.as_long() > 0: + return None + return _Refused( + SanitizerKind.NON_POSITIVE_STEP, + f"the loop step is {s.as_long()}; only positive steps are modeled", + ) + status, model, reason = self.solve([*env.premises, step <= 0]) + if status == unsat: + return None + if status == sat: + assert model is not None + return _Refused( + SanitizerKind.NON_POSITIVE_STEP, + f"the loop step can be {model.eval(step, model_completion=True)} " + f"(at {self.witness(env, model)}); only positive steps are modeled", + ) + return _Refused( + SanitizerKind.SOLVER_UNKNOWN, + f"Z3 could not decide whether the loop step is positive ({reason})", + ) + + # ── one access ── + + def check_access( + self, index: int, access: AccessEvent, grid: tuple[int, int, int] + ) -> None: + graph = self.graph + if mentions_observed(access.offset, graph): + # A free observation would make any address reachable. + raise _Refused( + SanitizerKind.OBSERVATION_IN_ADDRESS, + "the address depends on the value an atomic observed", + ) + facts = self.binding.tensors.get(access.base_param) + if facts is None: + raise _Refused( + SanitizerKind.MISSING_BINDING, + f"pointer argument {access.base_param!r} has no tensor binding" + + _binding_error(self.binding), + ) + problem = _unusable(facts) + if problem is not None: + raise _Refused( + SanitizerKind.MISSING_BINDING, + f"the tensor facts of {access.base_param!r} are unusable ({problem})", + ) + env = _Env(graph, self.binding, grid, self.ctx) + assumed: list[BoolRef] = [] + if access.in_loop: + # Even an offset without the induction variable executes only + # on iterations that run: none for a zero-trip loop. + env.loop_iteration() + assumed += [ + _fits(env.value(ob.term), ob.bits, ob.signed) for ob in self.loop_obs + ] + assumed += [env.value(d.b) != 0 for d in self.loop_divs] + # Lower every root first: a refusal leaves no partial result. + offset = env.value(access.offset) + path = env.cond(access.path) if access.path is not None else None + mask = env.cond(access.mask) if access.mask is not None else None + # One DAG over the three roles: a node's guard covers every root it + # reaches, so a node one role reads through a Select arm and another + # directly is checked wherever it is read. + dag = _Dag((access.path, access.mask, access.offset), graph, env) + obs: dict[str, list[WidthObligation]] = {"path": [], "mask": [], "offset": []} + for ob in width_obligations(graph, access): + if ob.role != "loop": + obs[ob.role].append(ob) + seen = set(self.loop_ids) if access.in_loop else set() + divs = { + role: _divisions((root,), graph, seen) + for role, root in ( + ("path", access.path), + ("mask", access.mask), + ("offset", access.offset), + ) + } + observed_path = access.path is not None and mentions_observed( + access.path, graph + ) + observed_mask = access.mask is not None and mentions_observed( + access.mask, graph + ) + + def uncertain(role: str) -> SanitizerKind | None: + if access.guarded or observed_path: + return SanitizerKind.UNMODELABLE_CONDITION + if role != "path" and observed_mask: + return SanitizerKind.DATA_DEPENDENT_MASK + if role == "offset" and access.mask_dropped: + return SanitizerKind.DATA_DEPENDENT_MASK + return None + + for role, condition in (("path", path), ("mask", mask), ("offset", None)): + assumed = self.role( + env, index, obs[role], divs[role], assumed, uncertain(role), dag + ) + if condition is not None: + assumed.append(condition) + self.out_of_bounds( + env, index, access, facts, offset, assumed, uncertain("offset") + ) + + def role( + self, + env: _Env, + index: int, + obs: Sequence[WidthObligation], + divs: Sequence[Bin], + assumed: list[BoolRef], + uncertain: SanitizerKind | None, + dag: _Dag, + ) -> list[BoolRef]: + """Discharge one role's width obligations and divisions under + ``assumed`` (the earlier roles' safety) in one joint query: some + condition of the role fails. None of the role's own conditions is + assumed while they are checked, only the sites reported already (so + a wrap another access reported does not resurface at the ops + computed from it). At most one finding of each kind: the model's + innermost failure, then a query for the other kind with that site + assumed. Returns ``assumed`` plus the role's safety.""" + conds: list[_Condition] = [] + for i, ob in enumerate(obs): + fits = _fits(env.value(ob.term), ob.bits, ob.signed) + guard = None if _undefined_when_wide(ob) else dag.guards.get(id(ob.term)) + conds.append( + _Condition( + "integer-overflow", + ob, + fits if guard is None else Implies(guard, fits), + # Children first (a select's condition before its arms), + # and on one term the op's own result before the reads + # of the ops that read it. + (dag.rank(ob.term), -i), + ) + ) + for d in divs: + nonzero = env.value(d.b) != 0 + conds.append(_Condition("division-by-zero", d, nonzero, (dag.rank(d), 0))) + held = [c.safe for c in conds if self.is_reported(c.kind, c.site)] + pending = [c for c in conds if not self.is_reported(c.kind, c.site)] + while pending: + status, model, reason = self.solve( + [ + *env.premises, + *assumed, + *held, + Or([Z3Not(c.safe) for c in pending]), + ] + ) + if status == unknown: + self.unknown(index, "an integer-overflow or division-by-zero", reason) + if status != sat: + break + assert model is not None + failed = [c for c in pending if is_false(model.eval(c.safe, True))] + chosen = min(failed, key=lambda c: c.rank) if failed else pending[0] + self.found(chosen, env, model, index, uncertain) + held.append(chosen.safe) + pending = [c for c in pending if c.kind != chosen.kind] + return [*assumed, *(c.safe for c in conds)] + + def out_of_bounds( + self, + env: _Env, + index: int, + access: AccessEvent, + facts: TensorFacts, + offset: ArithRef, + assumed: list[BoolRef], + uncertain: SanitizerKind | None, + ) -> None: + width = _pointee_bytes(access.elem_bits) + status, model, reason = self.solve( + [*env.premises, *assumed, Z3Not(_legal(offset, facts, width))] + ) + if status == unknown: + self.unknown(index, "the out-of-bounds", reason) + if status != sat: + return + assert model is not None + off = model.eval(offset, True).as_long() + # No address in the text: it is the same for every launch with the + # same launch_key (the address is violation_address). + detail = ( + f"{access.kind} of {access.base_param!r} at element offset {off} " + f"is outside {_footprint(facts)}" + ) + if uncertain is not None: + self.withhold(index, uncertain, "out-of-bounds", detail) + return + self.findings.append( + Finding( + kind="out-of-bounds", + access_index=index, + access_kind=access.kind, + base_param=access.base_param, + line_no=access.line_no, + loc=_location(access.loc), + witness=self.witness(env, model), + violation_offset=off, + violation_address=facts.data_ptr + off * width, + detail=detail, + ) + ) + + # ── results ── + + def found( + self, + cond: _Condition, + env: _Env, + model: ModelRef, + index: int, + uncertain: SanitizerKind | None, + ) -> None: + site = cond.site + extra: dict[str, int] = {} + if isinstance(site, WidthObligation): + value = model.eval(env.value(site.term), True).as_long() + detail = self.overflow_detail(site, value) + extra["value"] = value + else: + what = "remainder" if site.op in ("%", "u%") else "division" + detail = f"the divisor of this {what} ({site.op!r}) can be 0" + if uncertain is not None: + self.withhold(index, uncertain, cond.kind, detail) + return + self.reported.add(self.site_key(cond.kind, site)) + self._alive.append(site) + access = self.graph.accesses[index] + line_no, loc = site.line_no, site.loc + if line_no is None: + line_no, loc = access.line_no, access.loc + self.findings.append( + Finding( + kind=cond.kind, + access_index=index, + access_kind=access.kind, + base_param=access.base_param, + line_no=line_no, + loc=_location(loc), + witness={**self.witness(env, model), **extra}, + detail=detail, + ) + ) + + @staticmethod + def site_key(kind: str, site: WidthObligation | Bin) -> tuple[str, object]: + """One finding per op: by its TTIR line, else by the term's identity + (``found`` keeps a reported term alive, so its id stays unique).""" + if site.line_no is not None: + return kind, site.line_no + return kind, id(site.term if isinstance(site, WidthObligation) else site) + + def is_reported(self, kind: str, site: WidthObligation | Bin) -> bool: + return self.site_key(kind, site) in self.reported + + def withhold( + self, index: int, kind: SanitizerKind, finding: str, detail: str + ) -> None: + self.abstain( + index, + _Refused( + kind, + f"possible {finding} {_WITHHELD[kind]}, so the witness may be " + f"unreachable: {detail}", + ), + ) + + def unknown(self, index: int, query: str, reason: str | None) -> None: + self.abstain( + index, + _Refused( + SanitizerKind.SOLVER_UNKNOWN, + f"Z3 could not decide {query} query ({reason})", + ), + ) + + def abstain(self, index: int, refused: _Refused) -> None: + access = self.graph.accesses[index] + refused.at(access.line_no, access.loc) + entry = (index, refused.kind) + if entry not in self.abstained: + self.abstained.append(entry) + if self.refusal is None: + self.refusal = Refusal( + kind=refused.kind.value, + message=f"{describe_site(refused.line_no, refused.loc)}: " + f"{refused.message}", + line_no=refused.line_no, + loc=_location(refused.loc), + ) + + def overflow_detail(self, ob: WidthObligation, value: int) -> str: + if ob.signed: + half = 1 << (ob.bits - 1) + bounds = f"i{ob.bits} range [{-half}, {half})" + else: + bounds = f"unsigned i{ob.bits} range [0, {1 << ob.bits})" + # The obligation's origin (see width_obligations): an op's result, a + # trunci operand (the other signed ones), or an unsigned read. + loop = self.graph.loop + if loop is not None and _is_increment(ob, loop, self.loop_ids): + what = "the loop's induction-variable increment" + elif isinstance(ob.term, Bin) and ob.signed and ob.term.bits == ob.bits: + what = f"the result of {ob.term.op!r}" + elif ob.signed: + what = f"the operand of a truncation to i{ob.bits}" + else: + what = "a value read as unsigned" + return ( + f"{what} can be {value}, outside the {bounds}: the IR's " + "fixed-width arithmetic differs from the unbounded reading" + ) + + def witness(self, env: _Env, model: ModelRef) -> dict[str, int]: + def val(v: ArithRef) -> int: + return model.eval(v, model_completion=True).as_long() + + out = {f"pid_{axis}": val(pid) for axis, pid in enumerate(env.pids)} + for name, lane in env.lanes.items(): + out[name] = val(lane) + if env.iteration is not None: + out[str(env.iteration)] = val(env.iteration) + return out + + def solve(self, formulas: list[Any]) -> tuple[Any, ModelRef | None, str | None]: + """One query on a fresh Solver with its own timeout (never the + process-global z3.set_param). A query Z3 cancelled for a Ctrl+C + raises KeyboardInterrupt.""" + solver = Solver(ctx=self.ctx) + solver.set(timeout=self.timeout_ms) + solver.add(*formulas) + status = solver.check() + if status == sat: + return status, solver.model(), None + if status == unknown: + reason = solver.reason_unknown() + if reason == _INTERRUPTED: + raise KeyboardInterrupt(f"Z3 query cancelled ({reason})") + return status, None, reason + return status, None, None diff --git a/tilelens/clients/sanitizer/data.py b/tilelens/clients/sanitizer/data.py index 61b68e57c..161f6bb7f 100644 --- a/tilelens/clients/sanitizer/data.py +++ b/tilelens/clients/sanitizer/data.py @@ -1,11 +1,13 @@ from ...core.data import Store, Load +import operator import numpy as np from numpy.typing import NDArray from dataclasses import dataclass -from typing import Any +from typing import Any, Literal, get_args import torch import z3 +from ...ir.launch import TensorFacts from ...utils.traceback_utils import TracebackInfo @@ -54,3 +56,54 @@ class OutOfBoundsRecordZ3(OutOfBoundsRecord): violation_address: int symbolic_expr: Any = None # Optional symbolic expression tree tensor_name: str | None = None + + +CompiledFindingKind = Literal["out-of-bounds", "integer-overflow", "division-by-zero"] + + +@dataclass +class CompiledSanitizerRecord: + """One finding of the compiled sanitizer (``Sanitizer(compile=True)``) + on one access, with a witness: an out-of-bounds address (outside the + view's footprint), an address/mask/path term that can overflow its + declared integer width, or one that can divide by zero. + + Plain data: the tensor is described by its launch-time facts, never held, + since records outlive the launch and must not pin device memory. + """ + + kind: CompiledFindingKind + # The accessing op; an atomic (a read and a write) is reported as Store. + op_type: type[Store | Load] + # The kernel parameter of the accessed tensor, and its facts at launch; + # None when the launch bound no tensor to it (a finding in the loop's + # bounds is attributed to the loop's first access, whatever it reads). + tensor_name: str + tensor_facts: TensorFacts | None + # The free variables' values in the witness (program ids, arange lanes, + # loop iterations), by name. + witness: dict[str, int] + # The config kwargs of the config (binding) the finding is in. + config: dict[str, Any] + user_code_tracebacks: list[TracebackInfo] + # For "out-of-bounds": the element offset from the view's data_ptr and + # its byte address; None for the other kinds. + violation_offset: int | None = None + violation_address: int | None = None + detail: str | None = None + + def __post_init__(self) -> None: + if self.kind not in get_args(CompiledFindingKind): + raise ValueError(f"unknown compiled sanitizer finding: {self.kind!r}") + if self.op_type not in (Load, Store): + raise TypeError(f"op_type must be Load or Store, not {self.op_type!r}") + # Ints and plain containers only, so a saved trace holds the record. + self.witness = { + name: operator.index(value) for name, value in self.witness.items() + } + self.config = dict(self.config) + self.user_code_tracebacks = list(self.user_code_tracebacks) + for name in ("violation_offset", "violation_address"): + value = getattr(self, name) + if value is not None: + setattr(self, name, operator.index(value)) diff --git a/tilelens/clients/sanitizer/sanitizer.py b/tilelens/clients/sanitizer/sanitizer.py index 730d57040..d6b55232a 100644 --- a/tilelens/clients/sanitizer/sanitizer.py +++ b/tilelens/clients/sanitizer/sanitizer.py @@ -57,8 +57,10 @@ class Sanitizer(Client): """ - Factory class that returns the concrete sanitizer implementation - based on the value of ``cfg.enable_sanitizer``. + Factory class that returns the concrete sanitizer implementation: + ``Sanitizer()`` the eager one, ``Sanitizer(compile=True)`` the compiled + one (``CompiledSanitizer``, a virtual subclass); both are the + ``NullSanitizer`` while ``cfg.enable_sanitizer`` is off. """ NAME = "sanitizer" @@ -67,13 +69,23 @@ class Sanitizer(Client): def __new__(cls: type[SanitizerT], *args: Any, **kwargs: Any) -> SanitizerT: if cls is Sanitizer: + # The disable flag wins over the mode: trace() leaves a kernel + # traced with a NullSanitizer untraced, compile=True or not. + if kwargs.pop("compile", False) and cfg.enable_sanitizer: + from .compiled.client import CompiledSanitizer + + # Only a virtual Sanitizer subclass, so Python does not call + # its __init__ after __new__: call it here, without ``compile``. + compiled = object.__new__(CompiledSanitizer) + CompiledSanitizer.__init__(compiled, *args, **kwargs) + return cast(SanitizerT, compiled) target_cls = cast( type["Sanitizer"], SymbolicSanitizer if cfg.enable_sanitizer else NullSanitizer, ) - obj = object.__new__(target_cls) - cast(Any, target_cls).__init__(obj, *args, **kwargs) - return cast(SanitizerT, obj) + # A Sanitizer subclass: Python calls its __init__ once this + # returns, with the call's own arguments, ``compile`` included. + return cast(SanitizerT, object.__new__(target_cls)) return cast(SanitizerT, object.__new__(cls)) def __init__(self, abort_on_error: bool = True, *args, **kwargs): @@ -154,7 +166,14 @@ def __hash__(self) -> int: class SymbolicSanitizer(Sanitizer, SymbolicClient): - def __init__(self, abort_on_error: bool = True): + # ``compile`` is the factory's mode switch: Sanitizer(compile=False) + # reaches this __init__ with it. The eager sanitizer is never compiled. + def __init__(self, abort_on_error: bool = True, *, compile: bool = False): + if compile: + raise TypeError( + "SymbolicSanitizer is the eager sanitizer; the compiled one is " + "Sanitizer(compile=True)" + ) super().__init__(abort_on_error=abort_on_error) self.records: list[OutOfBoundsRecordZ3] = [] self.cache_args: list[Any] = [] diff --git a/tilelens/wrapper.py b/tilelens/wrapper.py index d2759da85..3321c5b75 100644 --- a/tilelens/wrapper.py +++ b/tilelens/wrapper.py @@ -14,6 +14,17 @@ PROFILER_COMMAND = "tile-profiler" RACE_DETECTOR_COMMAND = "tile-race" +# tile-sanitizer's flag for the compiled sanitizer, given before the script +# name. +COMPILE_FLAG = "--compile" +# Printed (to stderr) once when the flag is given: the kernels do not run. +COMPILE_NOTE = ( + f"[{COMPILE_FLAG}] kernel launches are compiled and checked, not run: " + "their outputs are never written, so the script sees them unchanged, and " + "a launch whose arguments it computes from them is checked with those " + "values" +) + # Former Triton-Viz command names, still installed as aliases. LEGACY_COMMANDS = { SANITIZER_COMMAND: "triton-sanitizer", @@ -36,6 +47,16 @@ def sanitizer_wrapper(kernel, *, frontend: str = "triton"): return tracer(kernel) +def compiled_sanitizer_wrapper(kernel, *, frontend: str = "triton"): + # Checks each launch against the kernel's compiled TTIR instead of + # interpreting it; the kernel does not run (Sanitizer(compile=True)). + tracer = tilelens.trace( + client=Sanitizer(compile=True, abort_on_error=True), + frontend=frontend, + ) + return tracer(kernel) + + def profiler_wrapper(kernel, *, frontend: str = "triton"): tracer = tilelens.trace(client=Profiler(), frontend=frontend) return tracer(kernel) @@ -78,9 +99,11 @@ def _decorator(f): return _patched_autotune -def _apply_wrapper(wrapper_func, command_name, usage_msg): +def _apply_wrapper(wrapper_func, command_name, usage_msg, compile_wrapper=None): """ Generic function to apply a wrapper to triton.jit and run the user script. + A command with a ``compile_wrapper`` uses it instead when its first + argument is COMPILE_FLAG. """ legacy_command = LEGACY_COMMANDS[command_name] if os.path.basename(sys.argv[0]) not in (command_name, legacy_command): @@ -91,6 +114,11 @@ def _apply_wrapper(wrapper_func, command_name, usage_msg): cfg.cli_active = True + if compile_wrapper is not None and sys.argv[1:2] == [COMPILE_FLAG]: + del sys.argv[1] + wrapper_func = compile_wrapper + print(COMPILE_NOTE, file=sys.stderr) + # Patch Triton kernels with the Triton frontend. _patched_jit = create_patched_jit( wrapper_func, @@ -149,12 +177,21 @@ def _apply_wrapper(wrapper_func, command_name, usage_msg): def apply_sanitizer(): """ - Apply the sanitizer wrapper to triton.jit and run the user script. + Apply the sanitizer wrapper to triton.jit and run the user script; with + ``--compile`` before the script, the compiled sanitizer's. """ _apply_wrapper( sanitizer_wrapper, SANITIZER_COMMAND, - f"Usage: {SANITIZER_COMMAND} [args...]", + f"Usage: {SANITIZER_COMMAND} [{COMPILE_FLAG}] [args...]\n" + f" {COMPILE_FLAG} check each launch against the compiled kernel " + "(Sanitizer(compile=True)) instead of interpreting it; kernels are not " + "run, so their outputs are never written. An 'ok' holds for the " + "arguments each launch was called with. Kernels are compiled on the " + "host (no GPU needed) for TILELENS_IR_TARGET, by default cuda:89 " + "(e.g. cuda:90, hip:gfx942); a config that fails to compile for it is " + "reported as not checked, and the script goes on.", + compile_wrapper=compiled_sanitizer_wrapper, )