diff --git a/tests/unit/ir/test_lowering.py b/tests/unit/ir/test_lowering.py new file mode 100644 index 00000000..b9af247e --- /dev/null +++ b/tests/unit/ir/test_lowering.py @@ -0,0 +1,412 @@ +"""tilelens.ir.lowering: the Term -> Z3 lowering shared by the compiled-mode +clients. + +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's TTIR +comes from the IR tests' corpus (ttir_corpus.py, compiled at test time). +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 ttir_corpus + +LOOP = LoopInfo("%loop", 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 + ttir_kernels.py: off = off * s + pid, 600 times from off = pid), where + the generated == / hash raise at Python's default limit.""" + graph = parse_ttir(ttir_corpus.text("ttir/kernel_deep_chain.ttir")) + (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)] + for leaf in (Const(1), Pid(0), NumPrograms(1), Arange("%r", 0, 4), X, Observed(0)): + assert kids(leaf) == [] + # a value the reader could not model is a leaf too + assert kids(DataDep()) == [] + + +# ─────────────────────────── 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(): + with pytest.raises(_Refusal, match="loaded value"): + Lowerer(_graph(), _Leaves()).lower(BoolBin("and", X, DataDep("loaded value"))) + + +@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("oeq", 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 ac2cf24c..b336ab39 100644 --- a/tests/unit/sanitizer_compiled/test_oob.py +++ b/tests/unit/sanitizer_compiled/test_oob.py @@ -10,6 +10,7 @@ from __future__ import annotations +import gc import pickle import subprocess import sys @@ -939,6 +940,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): + """A wrap in a select's untaken arm, now in 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 GPU run faults), and so does INT_MIN // -1.""" @@ -1160,6 +1222,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/tilelens/clients/sanitizer/compiled/oob.py b/tilelens/clients/sanitizer/compiled/oob.py index 1a8fdae6..fc7745f5 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..fc2915ef --- /dev/null +++ b/tilelens/ir/lowering.py @@ -0,0 +1,290 @@ +"""Term -> Z3 lowering shared by the compiled-mode clients. + +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 leaf (a DataDep included) has none.""" + 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__}")