From b5fc25bafe43845a4b049a255800b0b4b5366c93 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Sat, 3 Oct 2026 21:51:52 -0400 Subject: [PATCH] [FEAT] Add a compiled mode to the sanitizer Sanitizer(compile=True), or tile-sanitizer --compile, checks a launch statically on the host-compiled TTIR instead of interpreting it. The kernel is not launched, so no GPU is needed and outputs are not written. Like the rest of IR mode it needs Triton 3.8: on another release every launch is unsupported (host-compile-unavailable) and the program goes on. - Every autotune config is checked with Z3 against the launch's scalar arguments, grid and tensor layouts; a launch is ok only if every config is proven in bounds. - Findings: out-of-bounds (against the view's element footprint, gaps of strided views included), integer-overflow on address, mask, branch and loop values, and division-by-zero, each with a witness. - Anything that cannot be modeled is reported as unsupported with a typed refusal; a solver unknown or timeout never counts as a proof. A kernel that fails to compile for the target is unsupported and the program continues; a call that does not bind raises as untraced. - Each launch records an IRVerdict with per-config verdicts in Launch.records. The tests read their kernels' TTIR from the test-time corpus (tests/unit/ir/ttir_corpus.py) rather than committed goldens. --- README.md | 91 + tests/end_to_end/test_compiled_sanitizer.py | 1875 +++++++++++++++++ tests/end_to_end/test_host_compile.py | 318 +++ .../end_to_end/test_ir_lifecycle_compiled.py | 836 ++++++++ tests/unit/ir/test_verdict_io.py | 305 +++ tests/unit/sanitizer_compiled/__init__.py | 0 tests/unit/sanitizer_compiled/test_client.py | 1053 +++++++++ tests/unit/sanitizer_compiled/test_oob.py | 1233 +++++++++++ tests/unit/test_wrapper.py | 82 + tilelens/clients/__init__.py | 8 + .../clients/sanitizer/compiled/__init__.py | 42 + tilelens/clients/sanitizer/compiled/client.py | 773 +++++++ tilelens/clients/sanitizer/compiled/oob.py | 1248 +++++++++++ tilelens/clients/sanitizer/data.py | 55 +- tilelens/clients/sanitizer/sanitizer.py | 31 +- tilelens/wrapper.py | 43 +- 16 files changed, 7983 insertions(+), 10 deletions(-) create mode 100644 tests/end_to_end/test_compiled_sanitizer.py create mode 100644 tests/end_to_end/test_host_compile.py create mode 100644 tests/end_to_end/test_ir_lifecycle_compiled.py create mode 100644 tests/unit/ir/test_verdict_io.py create mode 100644 tests/unit/sanitizer_compiled/__init__.py create mode 100644 tests/unit/sanitizer_compiled/test_client.py create mode 100644 tests/unit/sanitizer_compiled/test_oob.py create mode 100644 tilelens/clients/sanitizer/compiled/__init__.py create mode 100644 tilelens/clients/sanitizer/compiled/client.py create mode 100644 tilelens/clients/sanitizer/compiled/oob.py 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, )