From eafc0a0e392455428e067477c168b285c1aad996 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Wed, 30 Sep 2026 07:07:38 -0400 Subject: [PATCH] [REFACTOR] Share the Term to Z3 lowering between compiled clients The compiled sanitizer's evaluator becomes tilelens.ir.lowering, the one reading of the TTIR reader's term algebra as Z3 terms, so the compiled race detector can sit on the same semantics instead of a second copy. - lowering.py is mechanism only: operator semantics (truncating division, unsigned ops and predicates as their signed twins under the width obligations, IntCast as its operand, i1 coercion, loop terms) live there; every leaf (scalar arguments, pids, grid, lanes, the loop iteration, observations, unmodeled values) comes from the client's TermLeaves, which also raises the client's own refusals. - It never creates a Z3 context: constants use the leaves' context and every other term its operands', so a client keeps all of its terms in one context (per check for the sanitizer). - The walk is iterative and memoized by term identity, so terms deeper than the recursion limit lower. - The sanitizer keeps its policy (obligation findings, refusals, solving, timeouts). Its leaves reach the lowering through a weak proxy and a refused loop keeps a never-raised copy of its refusal, so a check's Z3 context is freed when the check returns instead of by the cyclic GC on another thread (concurrent checks could hang or crash). - Its walks stop at the induction variable, as before, so a loop bound read only through a Select arm keeps that arm's guard. No verdict, finding or witness changes: the sanitizer's results match the previous evaluator on the golden texts and the differential corpus on Triton 3.6 and 3.8. Hand-built graphs outside the reader's invariants (e.g. an iter arg naming another loop) now raise instead of being misread. --- tests/unit/ir/test_lowering.py | 417 +++++++++++++++++++++ tests/unit/sanitizer_compiled/test_oob.py | 102 +++++ tests/unit/test_ir_version_gate.py | 1 + tilelens/clients/sanitizer/compiled/oob.py | 249 ++++-------- tilelens/ir/lowering.py | 291 ++++++++++++++ 5 files changed, 889 insertions(+), 171 deletions(-) create mode 100644 tests/unit/ir/test_lowering.py create mode 100644 tilelens/ir/lowering.py diff --git a/tests/unit/ir/test_lowering.py b/tests/unit/ir/test_lowering.py new file mode 100644 index 00000000..be1b6f83 --- /dev/null +++ b/tests/unit/ir/test_lowering.py @@ -0,0 +1,417 @@ +"""tilelens.ir.lowering: the Term -> Z3 lowering shared by the compiled-mode +clients (D13). + +CPU only: terms are built by hand over small AccessGraphs and lowered with +a test TermLeaves whose leaves are free Z3 variables (or constants), then +checked by Z3 equivalence or by evaluation; kernel_deep_chain comes from +the golden the installed release reads. The compiled sanitizer's own +results on top of the lowering are pinned in tests/unit/sanitizer_compiled/. +""" + +from __future__ import annotations + +import sys + +import pytest +import z3 + +from tilelens.ir.lowering import Lowerer, children, fold +from tilelens.ir.ttir_reader import ( + AccessGraph, + Arange, + Bin, + BoolBin, + Cmp, + Const, + DataDep, + IntCast, + IterArgInfo, + IterArgOffset, + LoopInfo, + LoopVar, + Not, + NumPrograms, + Observed, + Param, + Pid, + PtrValue, + Select, + parse_ttir, +) + +from . import _goldens as G + +LOOP = LoopInfo("%loop", "%i", lower=Param("lo"), upper=Param("hi"), step=Param("st")) + + +class _Refusal(Exception): + """A client's own refusal, raised from a leaf.""" + + +class _Leaves: + """Every leaf a free Int of ``ctx`` named after it (a Param in + ``params`` the constant instead); records the names it made.""" + + def __init__(self, ctx: z3.Context | None = None, params: dict | None = None): + self.ctx = ctx + self.params = params or {} + self.made: list[str] = [] + + def _var(self, name: str) -> z3.ArithRef: + self.made.append(name) + return z3.Int(name, self.ctx) + + def param(self, t): + if t.name in self.params: + return z3.IntVal(self.params[t.name], self.ctx) + return self._var(t.name) + + def pid(self, t): + return self._var(f"pid_{t.axis}") + + def num_programs(self, t): + return self._var(f"grid_{t.axis}") + + def arange(self, t): + return self._var(f"arange_{t.start}_{t.end}_d{t.dim}") + + def iteration(self, loop_ssa): + return self._var(f"k{loop_ssa}") + + def observed(self, t): + return self._var(f"observed_{t.access_index}") + + def data_dep(self, t): + raise _Refusal(t.why) + + +def _graph(loop: LoopInfo | None = None, iter_args=()) -> AccessGraph: + return AccessGraph("k", (), (), loop, iter_args) + + +def _lower(term, graph: AccessGraph | None = None, leaves: _Leaves | None = None): + return Lowerer(graph or _graph(LOOP), leaves or _Leaves()).lower(term) + + +def _v(name: str, ctx: z3.Context | None = None) -> z3.ArithRef: + return z3.Int(name, ctx) + + +def _proved(claim) -> bool: + solver = z3.Solver(ctx=claim.ctx) + solver.add(z3.Not(claim)) + return solver.check() == z3.unsat + + +def _eval(term) -> int | bool: + e = z3.simplify(_lower(term)) + if z3.is_bool(e): + assert z3.is_true(e) or z3.is_false(e), e + return z3.is_true(e) + return e.as_long() + + +X, Y = Param("x"), Param("y") + + +# ─────────────────────────── the Z3 context ─────────────────────────── + + +def _every_kind() -> tuple[object, AccessGraph]: + """One term reaching every kind of the algebra but DataDep.""" + graph = _graph(LOOP, (IterArgInfo(0, "p", Param("o0"), Const(4), "%loop"),)) + lane = Bin("+", Arange("%r", 0, 16, dim=0), IntCast("extsi", 32, 64, Pid(1))) + moved = Bin("*", Bin("//", IterArgOffset(0), NumPrograms(0)), LoopVar("%loop")) + cond = BoolBin("or", Cmp("ult", X, Const(8)), Not(Cmp("eq", Observed(2), Y))) + return Select(cond, Bin("umin", lane, moved), Bin("%", X, Const(3))), graph + + +@pytest.mark.parametrize("own_context", [False, True]) +def test_terms_are_made_in_the_leaves_context(own_context, monkeypatch): + """ctx=None lowers into Z3's main context, a given Context into that + Context (a constant included), and the lowering creates none.""" + ctx = z3.Context() if own_context else None + expected = ctx if own_context else z3.main_ctx() + term, graph = _every_kind() + + def no_context(*args, **kwargs): + raise AssertionError("the lowering created a z3.Context") + + monkeypatch.setattr(z3.Context, "__init__", no_context) + lowerer = Lowerer(graph, _Leaves(ctx)) + lowered = [lowerer.lower(term), lowerer.value(term), lowerer.cond(term)] + lowered += [lowerer.lower(Const(7)), lowerer.cond(Const(1))] + lowered += [lowerer.value(Cmp("slt", X, Y)), lowerer.cond(Not(Const(0)))] + monkeypatch.undo() + assert all(e.ctx is expected for e in lowered) + assert z3.is_int(lowered[0]) and z3.is_bool(lowered[2]) + + +def test_main_context_terms_take_a_main_context_substitution(): + """A client that renames variables with z3.substitute over pairs of the + main context gets its rename on terms lowered with ctx=None.""" + lowered = _lower(Bin("*", Pid(1), Const(3))) + renamed = z3.substitute(lowered, (_v("pid_1"), _v("pid_1_copy"))) + assert _proved(renamed == _v("pid_1_copy") * 3) + + +# ─────────────────────────── iteration and the memo ─────────────────────────── + + +def test_a_chain_deeper_than_the_recursion_limit_lowers(): + graph = _graph() + chain = Pid(0) + for _ in range(5000): + chain = Bin("+", chain, Const(1), 32) + limit = sys.getrecursionlimit() + sys.setrecursionlimit(1000) + try: + lowered = Lowerer(graph, _Leaves()).value(chain) + finally: + sys.setrecursionlimit(limit) + assert _proved(lowered == _v("pid_0") + 5000) + + +def test_kernel_deep_chain_lowers_at_the_default_recursion_limit(): + """kernel_deep_chain's offset nests more than 1000 levels (N = 600 in + generate_ttir.py: off = off * s + pid, 600 times from off = pid), where + the generated == / hash raise at Python's default limit.""" + text = G.texts("ttir")["kernel_deep_chain.ttir"].read_text(encoding="utf-8") + graph = parse_ttir(text) + (store,) = graph.accesses + applied: list[object] = [] + limit = sys.getrecursionlimit() + sys.setrecursionlimit(1000) + try: + offset = Lowerer(graph, _Leaves(params={"s": 1})).value(store.offset) + fold(store.offset, graph, lambda t, _: applied.append(t), {}) + finally: + sys.setrecursionlimit(limit) + assert sum(isinstance(t, Bin) for t in applied) > 1000 + assert _proved(offset == 601 * _v("pid_0")) + + +def test_each_term_is_lowered_once_by_identity(): + """The memo is keyed by identity: one Param object read twice is one + leaf call, two equal Param objects are two; a Lowerer keeps its memo + across calls.""" + leaves = _Leaves() + lowerer = Lowerer(_graph(), leaves) + n = Param("n") + shared = Bin("+", n, n) + lowerer.lower(shared) + lowerer.lower(Bin("*", shared, n)) + assert leaves.made == ["n"] + lowerer.lower(Bin("+", Param("n"), Param("n"))) + assert leaves.made == ["n", "n", "n"] + + +def test_fold_applies_children_first_and_each_term_once(): + graph = _graph(LOOP, (IterArgInfo(0, "p", Const(2), Const(3)),)) + two = Const(2) + root = Bin("+", Bin("*", two, two), IterArgOffset(0)) + order: list[object] = [] + + def apply(t, values): + order.append(t) + if isinstance(t, Const): + return t.value + if isinstance(t, IterArgOffset): + return values[0] + 10 * values[1] # iteration 10 + return values[0] * values[1] if t.op == "*" else values[0] + values[1] + + memo: dict = {} + assert fold(root, graph, apply, memo) == 2 * 2 + (2 + 10 * 3) + ids = [id(t) for t in order] + assert len(ids) == 6 and ids.count(id(two)) == 1 + assert ids.index(id(two)) < ids.index(id(root.a)) < ids.index(id(root)) + # the memo holds every term: a second fold applies nothing + assert fold(root, graph, apply, memo) == 36 and len(order) == 6 + + +def test_children_is_the_value_dependency_relation(): + info = IterArgInfo(0, "p", Param("o0"), Param("d")) + graph = _graph(LOOP, (info,)) + less = Cmp("slt", X, Y) + + def kids(t) -> list[int]: + return [id(k) for k in children(t, graph)] + + assert ( + kids(Bin("+", X, Y)) + == kids(less) + == kids(BoolBin("or", X, Y)) + == [id(X), id(Y)] + ) + assert kids(Select(less, X, Y)) == [id(less), id(X), id(Y)] + assert kids(Not(less)) == kids(IntCast("trunci", 64, 32, less)) == [id(less)] + assert kids(IterArgOffset(0)) == [id(info.offset0), id(info.delta)] + assert kids(LoopVar("%loop")) == [id(LOOP.lower), id(LOOP.step)] + # a DataDep's keep is not its value + assert kids(DataDep(keep=less)) == [] + for leaf in (Const(1), Pid(0), NumPrograms(1), Arange("%r", 0, 4), X, Observed(0)): + assert kids(leaf) == [] + + +# ─────────────────────────── operator semantics ─────────────────────────── + + +@pytest.mark.parametrize( + "unsigned, signed", [("u//", "//"), ("u%", "%"), ("umin", "min"), ("umax", "max")] +) +def test_unsigned_ops_read_as_their_signed_twins(unsigned, signed): + lowerer = Lowerer(_graph(), _Leaves()) + assert _proved( + lowerer.lower(Bin(unsigned, X, Y)) == lowerer.lower(Bin(signed, X, Y)) + ) + + +@pytest.mark.parametrize( + "unsigned, signed", [("ult", "slt"), ("ule", "sle"), ("ugt", "sgt"), ("uge", "sge")] +) +def test_unsigned_predicates_read_as_their_signed_twins(unsigned, signed): + lowerer = Lowerer(_graph(), _Leaves()) + got, twin = lowerer.lower(Cmp(unsigned, X, Y)), lowerer.lower(Cmp(signed, X, Y)) + assert z3.is_bool(got) and _proved(got == twin) + + +@pytest.mark.parametrize( + "pred, expected", + [ + ("slt", lambda x, y: x < y), + ("sle", lambda x, y: x <= y), + ("sgt", lambda x, y: x > y), + ("sge", lambda x, y: x >= y), + ("eq", lambda x, y: x == y), + ("ne", lambda x, y: x != y), + ], +) +def test_predicates(pred, expected): + assert _proved(_lower(Cmp(pred, X, Y)) == expected(_v("x"), _v("y"))) + + +@pytest.mark.parametrize("kind", ["trunci", "extsi", "extui"]) +def test_an_int_cast_reads_as_its_operand(kind): + assert z3.eq(_lower(IntCast(kind, 64, 32, X)), _v("x")) + # of an i1: the compare's 0/1 + cast = _lower(IntCast(kind, 1, 32, Cmp("slt", X, Y))) + assert z3.is_int(cast) and _proved(cast == z3.If(_v("x") < _v("y"), 1, 0)) + + +def test_i1_values_are_bool_or_0_1_by_position(): + x, y = _v("x"), _v("y") + lowerer = Lowerer(_graph(), _Leaves()) + less = Cmp("slt", X, Y) + # an integer position: 0/1 + assert _proved(lowerer.lower(Bin("+", less, Const(1))) == z3.If(x < y, 1, 0) + 1) + assert _proved(lowerer.value(less) == z3.If(x < y, 1, 0)) + # a boolean position: an i1 constant (dense) and an Int are != 0 + assert _proved(lowerer.lower(BoolBin("and", Const(1), less)) == (x < y)) + assert _proved(lowerer.lower(BoolBin("or", X, Const(0))) == (x != 0)) + assert _proved(lowerer.lower(Not(Const(0)))) + assert _proved(lowerer.cond(X) == (x != 0)) + # a compare of an i1 with an integer reads the i1 as 0/1 + assert _proved(lowerer.lower(Cmp("eq", less, Const(1))) == (x < y)) + # a Select: its condition is boolean; Bool arms stay Bool + both = lowerer.lower(Select(X, less, Cmp("eq", X, Y))) + assert z3.is_bool(both) and _proved(both == z3.If(x != 0, x < y, x == y)) + # arms of different sorts are Int + mixed = lowerer.lower(Select(less, Cmp("eq", X, Const(0)), Const(5))) + assert z3.is_int(mixed) + assert _proved(mixed == z3.If(x < y, z3.If(x == 0, 1, 0), 5)) + + +@pytest.mark.parametrize( + "op, a, b, expected", + [ + ("min", 3, -2, -2), + ("min", -2, 3, -2), + ("max", 3, -2, 3), + ("max", -2, 3, 3), + ("umin", 4, 9, 4), + ("umax", 4, 9, 9), + ("+", 7, -9, -2), + ("-", 7, -9, 16), + ("*", 7, -9, -63), + ], +) +def test_arithmetic(op, a, b, expected): + assert _eval(Bin(op, Const(a), Const(b), 32)) == expected + + +@pytest.mark.parametrize( + "a, b, quotient, remainder", + [ + (7, 2, 3, 1), + (-7, 2, -3, -1), + (7, -2, -3, 1), + (-7, -2, 3, -1), + (6, 3, 2, 0), + (-6, 3, -2, 0), + (0, -5, 0, 0), + (-1, 5, 0, -1), + ], +) +def test_division_truncates_toward_zero(a, b, quotient, remainder): + """arith.divsi / remsi (and divui / remui on the non-negative operands + their obligations leave): the remainder has the dividend's sign.""" + assert _eval(Bin("//", Const(a), Const(b), 32)) == quotient + assert _eval(Bin("%", Const(a), Const(b), 32)) == remainder + if a >= 0 and b >= 0: + assert _eval(Bin("u//", Const(a), Const(b), 32)) == quotient + assert _eval(Bin("u%", Const(a), Const(b), 32)) == remainder + + +def test_loop_terms_read_the_clients_iteration(): + """LoopVar is lower + k * step, IterArgOffset offset0 + k * delta, with + k the leaves' iteration of that loop: an IterArgInfo without a loop_ssa + is the graph's loop's.""" + graph = _graph( + LOOP, + ( + IterArgInfo(0, "p", Param("o0"), Param("d0")), + IterArgInfo(1, "p", Param("o1"), Param("d1"), "%other"), + ), + ) + leaves = _Leaves() + lowerer = Lowerer(graph, leaves) + k, other = _v("k%loop"), _v("k%other") + assert _proved(lowerer.lower(LoopVar("%loop")) == _v("lo") + k * _v("st")) + assert _proved(lowerer.lower(IterArgOffset(0)) == _v("o0") + k * _v("d0")) + assert _proved(lowerer.lower(IterArgOffset(1)) == _v("o1") + other * _v("d1")) + assert "hi" not in leaves.made # the upper bound is not the variable's value + + +# ─────────────────────────── leaves and bugs ─────────────────────────── + + +def test_a_leaf_refusal_propagates_unchanged(): + leaves = _Leaves() + with pytest.raises(_Refusal, match="loaded value"): + Lowerer(_graph(), leaves).lower(BoolBin("and", X, DataDep("loaded value"))) + # a DataDep's keep is not lowered + with pytest.raises(_Refusal): + Lowerer(_graph(), leaves).cond(DataDep(keep=Cmp("slt", Param("kept"), Y))) + assert "kept" not in leaves.made + + +@pytest.mark.parametrize( + "term, error, match", + [ + (object(), TypeError, "unknown term object"), + (PtrValue("p", Const(0)), TypeError, "unknown term PtrValue"), + (Bin("**", X, Y), ValueError, "unknown integer op"), + (Cmp("olt", X, Y), ValueError, "unknown cmpi predicate"), + (BoolBin("xor", X, Y), ValueError, "unknown boolean op"), + ], +) +def test_a_term_outside_the_algebra_is_a_bug(term, error, match): + with pytest.raises(error, match=match): + _lower(term) + + +@pytest.mark.parametrize("term", [LoopVar("%loop"), IterArgOffset(0)]) +def test_a_loop_term_without_a_loop_is_a_bug(term): + graph = _graph(None, (IterArgInfo(0, "p", Const(0), Const(1)),)) + with pytest.raises(ValueError, match="without a loop"): + Lowerer(graph, _Leaves()).lower(term) diff --git a/tests/unit/sanitizer_compiled/test_oob.py b/tests/unit/sanitizer_compiled/test_oob.py index ef67db83..b27bc83c 100644 --- a/tests/unit/sanitizer_compiled/test_oob.py +++ b/tests/unit/sanitizer_compiled/test_oob.py @@ -9,6 +9,7 @@ from __future__ import annotations +import gc import pickle import subprocess import sys @@ -947,6 +948,67 @@ def test_a_value_also_read_directly_is_checked_unguarded(): assert (f.kind, f.line_no) == ("integer-overflow", _line(text, "arith.muli")) +# A loop bound (the step; the lower bound) wraps in the arm of a select on +# %m, taken where m < 100; the loop reads the bound itself. +_BOUND_IN_ARM = { + "step truncated": ( + """ + %c0 = arith.constant 0 : i32 + %c100 = arith.constant 100 : i32 + %c1000 = arith.constant 1000 : i32 + %small = arith.cmpi slt, %m, %c100 : i32 + %t = arith.trunci %n : i32 to i8 + %e = arith.extsi %t : i8 to i32 + %sel = arith.select %small, %e, %c0 : i32 + scf.for %i = %c0 to %c1000 step %n : i32 { + %o = arith.addi %i, %sel : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + }""", + 300, + "arith.trunci", + ), + "lower bound read unsigned": ( + """ + %c0 = arith.constant 0 : i32 + %c1 = arith.constant 1 : i32 + %c8 = arith.constant 8 : i32 + %c64 = arith.constant 64 : i32 + %c100 = arith.constant 100 : i32 + %hi = arith.addi %n, %c8 : i32 + %small = arith.cmpi slt, %m, %c100 : i32 + %u = arith.minui %n, %c8 : i32 + %sel = arith.select %small, %u, %c0 : i32 + scf.for %i = %n to %hi step %c1 : i32 { + %j = arith.addi %i, %c64 : i32 + %o = arith.addi %j, %sel : i32 + %a = tt.addptr %p, %o : !tt.ptr, i32 + tt.store %a, %c0 : !tt.ptr + }""", + -10, + "arith.minui", + ), +} + + +@pytest.mark.parametrize("case", list(_BOUND_IN_ARM)) +def test_a_loop_bound_wrapping_in_a_discarded_arm_is_no_finding(case): + """H14 for a loop bound: the loop reads the bound itself, but that read + is the loop's (checked with the loop's obligations), so the access + reads the wrapping op only through the arm: its wrap matters only + where the select takes the arm.""" + body, n, op = _BOUND_IN_ARM[case] + g, text = _module(body, "%p: !tt.ptr, %n: i32, %m: i32") + + def at(m): + return check_graph(g, _bind(params={"arg1": n, "arg2": m}, arg0=_facts(1300))) + + _clean(at(500)) + (f,) = at(5).findings + assert (f.kind, f.line_no) == ("integer-overflow", _line(text, op)) + assert f.witness["value"] == n + + def test_undefined_divisions_count_in_either_arm(): """where(n != 0, pid // n, 0) divides by zero in the discarded arm too (the audit's p12: the GPU run faults), and so does INT_MIN // -1.""" @@ -1168,6 +1230,46 @@ def test_a_real_ctrl_c_stops_a_hard_query(): assert int(proc.stdout.split()[-1]) < 30 +@pytest.mark.parametrize( + "params", + [ + # a loop (its iteration lowers the bounds) and a finding's witness + _MATMUL_PARAMS, + # the loop refused: its bound reads K, which has no binding + {name: v for name, v in _MATMUL_PARAMS.items() if name != "K"}, + ], + ids=["finding", "refused-loop"], +) +def test_a_check_leaves_no_z3_object_to_the_cyclic_gc(params): + """A check's Z3 context and terms are freed by reference counting when + check_graph returns: the cyclic GC would free them later, on whichever + host thread it runs, inside that thread's own Z3 call (concurrent + checks hung or crashed).""" + tensors = _all(128 * 128, "a_ptr", "b_ptr", "c_ptr", elem_size=2) + graph = _graph(MATMUL) + binding = _bind((3, 2, 1), params, **tensors) + + def z3_objects() -> int: + # type(): isinstance would read __class__, which some objects warn on + z3_types = (z3.Context, z3.AstRef) + return sum(issubclass(type(o), z3_types) for o in gc.get_objects()) + + check_graph(graph, binding) + gc.collect() + gc.disable() + try: + before = z3_objects() + result = check_graph(graph, binding) + after = z3_objects() + finally: + gc.enable() + if "K" in params: + assert result.findings + else: + assert K.MISSING_BINDING in {kind for _, kind in result.abstained} + assert after == before + + _THREADS_SCRIPT = """ import sys, threading from pathlib import Path diff --git a/tests/unit/test_ir_version_gate.py b/tests/unit/test_ir_version_gate.py index 8eba829a..b7dfebb5 100644 --- a/tests/unit/test_ir_version_gate.py +++ b/tests/unit/test_ir_version_gate.py @@ -391,6 +391,7 @@ def test_the_ir_mode_modules_are_the_d29_list(tests_conftest): ir_mode = [ "unit/ir/test_host_compile.py", "unit/ir/test_ir_capture.py", + "unit/ir/test_lowering.py", "unit/ir/test_mlir_walk.py", "unit/ir/test_ttir_reader.py", "unit/ir/test_verdict_io.py", diff --git a/tilelens/clients/sanitizer/compiled/oob.py b/tilelens/clients/sanitizer/compiled/oob.py index 360399ab..a49986f2 100644 --- a/tilelens/clients/sanitizer/compiled/oob.py +++ b/tilelens/clients/sanitizer/compiled/oob.py @@ -61,17 +61,20 @@ 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. +Terms are lowered by the shared ``tilelens.ir.lowering`` (the reader's +operator semantics) with ``_Env`` as its leaves: the launch's constants, and +the free variables above with their range premises. Terms can be deeper +than Python's recursion limit and their generated ``==`` / ``hash`` +recurse, so the walks over them are iterative and keyed by term identity. """ from __future__ import annotations +import weakref from collections.abc import Hashable, Iterable, Iterator, Mapping, Sequence from dataclasses import dataclass, replace from enum import Enum -from typing import Any, Literal +from typing import Any, Literal, NoReturn from z3 import ( And, @@ -80,7 +83,6 @@ BoolVal, Context, Exists, - If, Implies, Int, IntVal, @@ -88,7 +90,6 @@ Or, Solver, Sum, - is_bool, is_false, is_int_value, sat, @@ -99,20 +100,15 @@ from z3 import Not as Z3Not from ....ir.launch import LaunchBinding, TensorFacts +from ....ir.lowering import Lowerer, children from ....ir.ttir_reader import ( AccessEvent, AccessGraph, Arange, Bin, - BoolBin, - Cmp, - Const, DataDep, - IntCast, - IterArgOffset, LoopInfo, LoopVar, - Not, NumPrograms, Observed, Param, @@ -312,81 +308,13 @@ def at(self, line_no: int | None, loc: Any) -> _Refused: 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 _operands(t: object, graph: AccessGraph) -> tuple: + """The nodes the checks' walks descend to from ``t``: the lowering's + ``children``, but none of a LoopVar. The loop's bounds belong to the + loop's family (checked there unguarded, assumed by every access in the + loop, kept out of its accesses' divisions by ``loop_ids``), so a bound + node an access reads through a Select arm keeps that arm's guard.""" + return () if isinstance(t, LoopVar) else children(t, graph) def _walk( @@ -401,7 +329,7 @@ def _walk( continue seen.add(id(t)) yield t - stack.extend(reversed(_kids(t, graph))) + stack.extend(reversed(_operands(t, graph))) _DIVISIONS = frozenset({"//", "%", "u//", "u%"}) @@ -435,7 +363,8 @@ def _lane_name(t: Arange) -> str: 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.""" + loop's own checks), their range premises, and the family's lowering: + the env is its leaves (``tilelens.ir.lowering.TermLeaves``).""" def __init__( self, @@ -453,30 +382,41 @@ def __init__( 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). + # (dim, extent) -> the lane's position along that dim (see arange). 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.k: ArithRef | None = None # the loop's iteration index, once used self._observed: dict[int, ArithRef] = {} - # id(term) -> (term, lowered): holding the term keeps its id unique. - self._memo: dict[int, tuple[object, Any]] = {} + # The lowering reaches its leaves through a proxy: a strong + # reference back would be a cycle, keeping the check's Z3 context + # and terms alive until the cyclic GC frees them, on whichever host + # thread it runs, inside that thread's own Z3 call. + self.lowering = Lowerer(graph, weakref.proxy(self)) # ── leaves ── - def param(self, name: str) -> int: - try: - value = self.binding.params[name] - except KeyError: + def param(self, t: Param) -> ArithRef: + if t.name not in self.binding.params: + # Raised outside an except block: a context exception's traceback + # would reach the frames holding this env (see check_loop). raise _Refused( SanitizerKind.MISSING_BINDING, - f"scalar argument {name!r} has no launch binding" + f"scalar argument {t.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) + ) + value = self.binding.params[t.name] + arg = self.graph.arg(t.name) + bits = arg.int_bits if arg is not None else 0 + return IntVal(_signed(value, bits), self.ctx) + + def pid(self, t: Pid) -> ArithRef: + return self.pids[t.axis] + + def num_programs(self, t: NumPrograms) -> ArithRef: + return IntVal(self.grid[t.axis], self.ctx) - def lane(self, t: Arange) -> ArithRef: + def arange(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).""" @@ -491,8 +431,20 @@ def lane(self, t: Arange) -> ArithRef: self.lanes.setdefault(_lane_name(t), value) return value - def observed(self, index: int) -> ArithRef: + def iteration(self, loop_ssa: str) -> ArithRef: + """``k`` of the graph's loop (it has at most one): see + ``loop_iteration``.""" + loop = self._loop() + if loop_ssa != loop.loop_ssa: + raise ValueError( + f"kernel {self.graph.kernel_name!r}: loop {loop_ssa!r} is not " + f"the graph's loop {loop.loop_ssa!r}" + ) + return self.loop_iteration() + + def observed(self, t: Observed) -> ArithRef: """An atomic observation: a free value of the atomic's width.""" + index = t.access_index v = self._observed.get(index) if v is None: v = Int(f"observed_{index}", self.ctx) @@ -503,6 +455,11 @@ def observed(self, index: int) -> ArithRef: self.premises += [v >= -half, v < half] return v + def data_dep(self, t: DataDep) -> NoReturn: + raise _Refused(SanitizerKind.UNMODELED_VALUE, f"an unmodeled value ({t.why})") + + # ── the loop ── + def bounds(self) -> tuple[ArithRef, ArithRef, ArithRef]: loop = self._loop() return self.value(loop.lower), self.value(loop.upper), self.value(loop.step) @@ -511,13 +468,13 @@ 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: + if self.k 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 + self.k = k + return self.k def _loop(self) -> LoopInfo: loop = self.graph.loop @@ -530,67 +487,10 @@ def _loop(self) -> LoopInfo: # ── terms ── def value(self, term: object) -> ArithRef: - return _as_int(self.lower(term)) + return self.lowering.value(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__}") + return self.lowering.cond(term) class _Dag: @@ -617,7 +517,9 @@ def __init__( continue stack.append((t, True)) stack.extend( - (k, False) for k in reversed(_kids(t, graph)) if id(k) not in self.order + (k, False) + for k in reversed(_operands(t, graph)) + if id(k) not in self.order ) self.guards: dict[int, BoolRef] = {} if env is None: @@ -637,7 +539,7 @@ def __init__( 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)] + edges = [(k, guard) for k in _operands(t, graph)] for kid, kid_guard in edges: if kid_guard is None: direct.add(id(kid)) @@ -652,7 +554,9 @@ def rank(self, term: object) -> float: 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 + self.order[id(k)] + for k in _operands(term, self.graph) + if id(k) in self.order ] return max(kids) + 0.5 if kids else len(self.order) @@ -875,7 +779,10 @@ def check_loop(self, first: int, grid: tuple[int, int, int]) -> _Refused | None: _lower, _upper, step = env.bounds() refusal = self.step_refusal(env, step) except _Refused as r: - refusal = r + # A copy that was never raised: keeping ``r`` would keep its + # traceback, which reaches this frame and so ``env``, a cycle + # leaving the check's Z3 context to the cyclic GC (see _Env). + refusal = _Refused(r.kind, r.message, r.line_no, r.loc) if refusal is not None: return refusal.at(loop.line_no, loop.loc) roots = (loop.lower, loop.upper, loop.step) @@ -1226,8 +1133,8 @@ def val(v: ArithRef) -> int: 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) + if env.k is not None: + out[str(env.k)] = val(env.k) return out def solve(self, formulas: list[Any]) -> tuple[Any, ModelRef | None, str | None]: diff --git a/tilelens/ir/lowering.py b/tilelens/ir/lowering.py new file mode 100644 index 00000000..00311a80 --- /dev/null +++ b/tilelens/ir/lowering.py @@ -0,0 +1,291 @@ +"""Term -> Z3 lowering shared by the compiled-mode clients (D13). + +The one reading of the TTIR reader's term algebra (``ttir_reader``'s +``Term``) as Z3 integers and booleans, for every client that queries an +``AccessGraph`` with Z3. Mechanism only: the operator semantics live here, +every leaf's meaning comes from the client. A client's :class:`TermLeaves` +says what a scalar argument, a program id, the grid, a lane, a loop's +iteration, an atomic observation and an unmodeled value are, and a client +that cannot model a leaf raises its own exception from it, which propagates +unchanged. The lowering itself raises only on a malformed graph, a bug and +never a limit of the model: ``TypeError`` or ``ValueError`` for a term or op +outside the algebra, or for a ``LoopVar`` (or an ``IterArgOffset`` whose +IterArgInfo names no loop) in a graph without a loop, and ``IndexError`` for +an ``arg_id`` outside ``iter_args``. + +Operator semantics, the reader's integer model (see ``ttir_reader``): + +* ``+ - *`` are unbounded Int arithmetic; ``//`` and ``%`` truncate toward + zero (``arith.divsi`` / ``remsi``: the remainder takes the dividend's + sign); ``min`` / ``max`` pick an operand; +* the unsigned ops (``u//``, ``u%``, ``umin``, ``umax``) and predicates + (``ult``, ...) read as their signed twins, and an ``IntCast`` as its + operand: exact only where the access's ``width_obligations`` hold, which + the client discharges (a zero divisor is the client's concern too); +* an i1 value is a Bool in a boolean position (a mask or path, a Select's + condition, ``and`` / ``or`` / negation) and 0/1 in an integer one, and a + Select whose arms differ in sort is an Int; +* ``LoopVar`` is ``lower + k * step`` and ``IterArgOffset`` is ``offset0 + + k * delta``, with ``k`` the client's iteration of that loop. + +Z3 context: the lowering never creates one. Constants are made in +``leaves.ctx`` (None: Z3's main context) and every other term in its +operands' context, so all of a client's terms share the context of its +leaves (the compiled sanitizer's is a Context per check). + +Terms can be deeper than Python's recursion limit and their generated +``==`` / ``hash`` recurse, so the walk is iterative and memoized by term +identity. +""" + +from __future__ import annotations + +from collections.abc import Callable, Sequence +from typing import Any, Protocol + +from z3 import And, ArithRef, BoolRef, Context, If, IntVal, Or, is_bool +from z3 import Not as Z3Not + +from .ttir_reader import ( + AccessGraph, + Arange, + Bin, + BoolBin, + Cmp, + Const, + DataDep, + IntCast, + IterArgOffset, + LoopInfo, + LoopVar, + Not, + NumPrograms, + Observed, + Param, + Pid, + Select, +) + + +class TermLeaves(Protocol): + """A client's meaning of every leaf of the algebra (there are no + defaults): each method returns a Z3 Int in ``ctx``, or raises the + client's own refusal.""" + + @property + def ctx(self) -> Context | None: + """The Z3 context of the client's terms (None: the main one).""" + + def param(self, t: Param) -> ArithRef: + """A scalar kernel argument.""" + + def pid(self, t: Pid) -> ArithRef: + """The program id along ``t.axis``.""" + + def num_programs(self, t: NumPrograms) -> ArithRef: + """The grid size along ``t.axis``.""" + + def arange(self, t: Arange) -> ArithRef: + """``t``'s value at the lane a query reads (the reader's contract: + key a lane by ``(dim, end - start)``, see ``Arange``).""" + + def iteration(self, loop_ssa: str) -> ArithRef: + """The iteration index ``k`` of the loop ``loop_ssa``.""" + + def observed(self, t: Observed) -> ArithRef: + """The old value the atomic at ``t.access_index`` observed.""" + + def data_dep(self, t: DataDep) -> ArithRef: + """A value the reader could not model (``t.why`` says which).""" + + +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 _divrem_ir(op: str, a: ArithRef, b: ArithRef) -> ArithRef: + """The IR's quotient (``op`` ``//`` or ``u//``) or remainder (``%`` or + ``u%``): truncating, so the remainder has the dividend's sign.""" + q = _trunc_div(a, b) + return q if op in ("//", "u//") else a - b * 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//", "%", "u%"): + return _divrem_ir(op, 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_TWIN = {"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) + p = _SIGNED_TWIN.get(pred, pred) + if p == "slt": + return a < b + if p == "sle": + return a <= b + if p == "sgt": + return a > b + if p == "sge": + return a >= b + if p == "eq": + return a == b + if p == "ne": + return a != b + raise ValueError(f"unknown cmpi predicate {pred!r}") + + +def _loop(graph: AccessGraph, t: object) -> LoopInfo: + if graph.loop is None: + raise ValueError( + f"kernel {graph.kernel_name!r}: a loop term ({type(t).__name__}) " + "without a loop" + ) + return graph.loop + + +def children(t: object, graph: AccessGraph) -> tuple: + """The terms ``t``'s value is computed from: its operands; for a + loop-carried pointer's offset, its IterArgInfo's ``offset0`` and + ``delta``; for the induction variable, the loop's ``lower`` and + ``step``. A DataDep has none: its ``keep`` is not its value (the + reader's ``observed_indices`` is the relation that reaches it).""" + 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) + if isinstance(t, LoopVar): + loop = _loop(graph, t) + return (loop.lower, loop.step) + return () + + +def fold( + root: object, + graph: AccessGraph, + apply: Callable[[object, Sequence[Any]], Any], + memo: dict[int, tuple[object, Any]], +) -> Any: + """``apply(t, the values of children(t))`` at ``root``, children first, + each term once for ``memo`` (id(term) -> (term, value): holding the + term keeps its id unique). Iterative post-order.""" + stack: list[tuple[object, bool]] = [(root, False)] + while stack: + t, ready = stack.pop() + if id(t) in memo: + continue + kids = children(t, 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, apply(t, [memo[id(k)][1] for k in kids])) + return memo[id(root)][1] + + +class Lowerer: + """The lowering of one family of a client's queries: bound to one graph + and one leaves object (so one Z3 context), with one memo.""" + + def __init__(self, graph: AccessGraph, leaves: TermLeaves) -> None: + self.graph = graph + self.leaves = leaves + self._memo: dict[int, tuple[object, Any]] = {} + + def lower(self, term: object) -> Any: + """``term`` as a Z3 expression: an Int, or a Bool for a compare, + ``and`` / ``or`` / negation, and a Select of Bools.""" + return fold(term, self.graph, self._apply, self._memo) + + 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 _apply(self, t: object, kids: Sequence[Any]) -> Any: + leaves = self.leaves + if isinstance(t, Const): + return IntVal(t.value, leaves.ctx) + if isinstance(t, Param): + return leaves.param(t) + if isinstance(t, Pid): + return leaves.pid(t) + if isinstance(t, NumPrograms): + return leaves.num_programs(t) + if isinstance(t, Arange): + return leaves.arange(t) + if isinstance(t, LoopVar): + k = leaves.iteration(t.loop_ssa) + return as_int(kids[0]) + k * as_int(kids[1]) + if isinstance(t, IterArgOffset): + loop_ssa = self.graph.iter_args[t.arg_id].loop_ssa + k = leaves.iteration(loop_ssa or _loop(self.graph, t).loop_ssa) + return as_int(kids[0]) + k * 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]) + if t.op == "and": + return And(a, b) + if t.op == "or": + return Or(a, b) + raise ValueError(f"unknown boolean op {t.op!r}") + 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 leaves.observed(t) + if isinstance(t, DataDep): + return leaves.data_dep(t) + raise TypeError(f"unknown term {type(t).__name__}")