From ae02a1e232dd23435273717bcb21ef6f52d64894 Mon Sep 17 00:00:00 2001 From: Sieds Lykles <93992551+S-Lykles@users.noreply.github.com> Date: Wed, 28 May 2025 22:28:37 +0200 Subject: [PATCH] [bounty] Z3 symbolic fuzzer [pr] (#10514) * First version, caught a bug? * Nicely print failure to reproduce * Remove that * Put the assert back * Change fuzzing to use testing_unit so it has z3 * Test key to match * Add rule * Add test * Add test for edge case 0 * Merge patterns * update comment * consistent whitespace * whitespace * add condition * add test * update comment * use Variable * fuzzer using z3_renderer * Cleaned up printing and debugging * working new fuzzer * change some comments and printing * more formatting * fuzz failures in seperate file * fix fstring * more tests * naming * remove added line * remove comment * print number of skipped expressions * use self.assertEqual --------- Co-authored-by: chenyu --- .github/workflows/test.yml | 4 +- test/external/fuzz_symbolic.py | 140 +++++++++++++----------- test/unit/test_symbolic_failures.py | 161 ++++++++++++++++++++++++++++ 3 files changed, 240 insertions(+), 65 deletions(-) create mode 100644 test/unit/test_symbolic_failures.py diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 7be737a5c4..8bb7044053 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -381,8 +381,8 @@ jobs: - name: Setup Environment uses: ./.github/actions/setup-tinygrad with: - key: fuzzing-minimal - deps: testing_minimal + key: fuzzing-unit + deps: testing_unit - name: Fuzz Test symbolic run: python test/external/fuzz_symbolic.py - name: Fuzz Test fast idiv diff --git a/test/external/fuzz_symbolic.py b/test/external/fuzz_symbolic.py index 8d2242a0c8..41970a932e 100644 --- a/test/external/fuzz_symbolic.py +++ b/test/external/fuzz_symbolic.py @@ -1,78 +1,92 @@ -import itertools -import random +import random, operator +import z3 from tinygrad import Variable, dtypes -from tinygrad.uop.ops import UOp -from tinygrad.helpers import DEBUG -random.seed(42) +from tinygrad.uop.ops import UOp, graph_rewrite +from tinygrad.uop.spec import z3_renderer +from tinygrad.helpers import DEBUG, Context -def add_v(expr, rng=None): - if rng is None: rng = random.randint(0,2) - return expr + v[rng], rng +seed = random.randint(0, 100) +print(f"Seed: {seed}") +random.seed(seed) -def div(expr, rng=None): - if rng is None: rng = random.randint(1,9) - return expr // rng, rng +unary_ops = [lambda a:a+random.randint(-4, 4), lambda a: a*random.randint(-4, 4), + lambda a: a//random.randint(1, 9), lambda a: a%random.randint(1, 9), + lambda a:a.maximum(random.randint(-10, 10)), lambda a:a.minimum(random.randint(-10, 10))] +binary_ops = [lambda a,b: a+b, lambda a,b: a*b, lambda a,b:a.maximum(b), lambda a,b:a.minimum(b)] +comp_ops = [operator.lt, operator.le, operator.gt, operator.ge] -def mul(expr, rng=None): - if rng is None: rng = random.randint(-4,4) - return expr * rng, rng +def random_or_sub_expression_int(depth, expr): + sub_expr = random.choice([e for e in expr.toposort() if e.dtype is not dtypes.bool]) + return random.choice([random_int_expr(depth-1), sub_expr]) -def mod(expr, rng=None): - if rng is None: rng = random.randint(1,9) - return expr % rng, rng +def random_int_expr(depth=10): + if depth <= 0: return random.choice(v) + expr1 = random_int_expr(depth-1) -def add_num(expr, rng=None): - if rng is None: rng = random.randint(-4,4) - return expr + rng, rng + # we give more weight to arithmatic ops than to minimum and maximum + ops = [ + lambda: random.choices(unary_ops, weights=[4, 4, 4, 4, 1, 1])[0](expr1), + # for the second operand its either another random exprssion or some subexpression of the first operand + lambda: random.choices(binary_ops, [8, 1, 1, 1])[0](expr1, random_or_sub_expression_int(depth-1, expr1)), + lambda: random_bool_expr(3, random_or_sub_expression_int(depth-1, expr1)).where(expr1, random_or_sub_expression_int(depth-1, expr1)), + ] + # we give weight proportional to the amount of ops in each branch + return random.choices(ops, weights=[6, 4, 1])[0]() -def lt(expr, rng=None): - if rng is None: rng = random.randint(-4,4) - return expr < rng, rng +def random_bool_expr(depth=10, expr1=None): + if depth == 0: return True + if expr1 is None: expr1 = random_int_expr(depth-1) + expr2 = random.choice([random_or_sub_expression_int(depth-1, expr1), UOp.const(dtypes.int, random.randint(-10, 10))]) + return random.choice(comp_ops)(expr1, expr2) -def ge(expr, rng=None): - if rng is None: rng = random.randint(-4,4) - return expr >= rng, rng - -def le(expr, rng=None): - if rng is None: rng = random.randint(-4,4) - return expr <= rng, rng - -def gt(expr, rng=None): - if rng is None: rng = random.randint(-4,4) - return expr > rng, rng - -# NOTE: you have to replace these for this test to pass -from tinygrad.uop.ops import python_alu, Ops -python_alu[Ops.MOD] = lambda x,y: x%y -python_alu[Ops.IDIV] = lambda x,y: x//y if __name__ == "__main__": - ops = [add_v, div, mul, add_num, mod] - for _ in range(1000): + skipped = 0 + for i in range(10000): + if i % 1000 == 0: + print(f"Running test {i}") upper_bounds = [*list(range(1, 10)), 16, 32, 64, 128, 256] u1 = Variable("v1", 0, random.choice(upper_bounds)) u2 = Variable("v2", 0, random.choice(upper_bounds)) u3 = Variable("v3", 0, random.choice(upper_bounds)) v = [u1,u2,u3] - tape = [random.choice(ops) for _ in range(random.randint(2, 30))] - # 10% of the time, add one of lt, le, gt, ge - if random.random() < 0.1: tape.append(random.choice([lt, le, gt, ge])) - expr = UOp.const(dtypes.int, 0) - rngs = [] - for t in tape: - expr, rng = t(expr) - if DEBUG >= 1: print(t.__name__, rng) - rngs.append(rng) - if DEBUG >=1: print(expr) - space = list(itertools.product(range(u1.vmin, u1.vmax+1), range(u2.vmin, u2.vmax+1), range(u3.vmin, u3.vmax+1))) - volume = len(space) - for (v1, v2, v3) in random.sample(space, min(100, volume)): - v = [v1,v2,v3] - rn = 0 - for t,r in zip(tape, rngs): rn, _ = t(rn, r) - num = eval(expr.render(simplify=False)) - if num != rn: - unsimplified_num = eval(expr.render(simplify=False)) - assert unsimplified_num == rn, "UNSIMPLIFIED MISMATCH!" - assert num == rn, f"mismatched {expr.render()} at {v1=} {v2=} {v3=} = {num} != {rn}\n{expr.render(simplify=False)}" - if DEBUG >= 1: print(f"matched {expr.render()} at {v1=} {v2=} {v3=} = {num} == {rn}") + expr = random_int_expr(6) + + with Context(CORRECT_DIVMOD_FOLDING=1): + simplified_expr = expr.simplify() + + solver = z3.Solver() + solver.set(timeout=5000) # some expressions take very long verify, but its very unlikely they actually return sat + z3_sink = graph_rewrite(expr.sink(simplified_expr, u1, u2, u3), z3_renderer, ctx=(solver, {})) + z3_expr, z3_simplified_expr = z3_sink.src[0].arg, z3_sink.src[1].arg + check = solver.check(z3_simplified_expr != z3_expr) + if check == z3.unknown and DEBUG>=1: + skipped += 1 + print("Skipped due to timeout or interrupt:\n" + + f"v1=Variable(\"{u1.arg[0]}\", {u1.arg[1]}, {u1.arg[2]})\n" + + f"v2=Variable(\"{u2.arg[0]}\", {u2.arg[1]}, {u2.arg[2]})\n" + + f"v3=Variable(\"{u3.arg[0]}\", {u3.arg[1]}, {u3.arg[2]})\n" + + f"expr = {expr.render(simplify=False)}\n") + elif check == z3.sat: + m = solver.model() + v1, v2, v3 = z3_sink.src[2].arg, z3_sink.src[3].arg, z3_sink.src[4].arg + n1, n2, n3 = m[v1], m[v2], m[v3] + u1_val, u2_val, u3_val = u1.const_like(n1.as_long()), u2.const_like(n2.as_long()), u3.const_like(n3.as_long()) + with Context(CORRECT_DIVMOD_FOLDING=1): + num = expr.simplify().substitute({u1:u1_val, u2:u2_val, u3:u3_val}).ssimplify() + rn = expr.substitute({u1:u1_val, u2:u2_val, u3:u3_val}).ssimplify() + if num==rn: print("z3 found a mismatch but the expressions are equal!!") + assert False, f"mismatched {expr.render()} at v1={m[v1]}; v2={m[v2]}; v3={m[v3]} = {num} != {rn}\n" +\ + "Reproduce with:\n" +\ + f"v1=Variable(\"{u1.arg[0]}\", {u1.arg[1]}, {u1.arg[2]})\n" +\ + f"v2=Variable(\"{u2.arg[0]}\", {u2.arg[1]}, {u2.arg[2]})\n" +\ + f"v3=Variable(\"{u3.arg[0]}\", {u3.arg[1]}, {u3.arg[2]})\n" +\ + f"expr = {expr}\n" +\ + f"v1_val, v2_val, v3_val = UOp.const(dtypes.int, {n1.as_long()}), UOp.const(dtypes.int, {n2.as_long()})," +\ + f"UOp.const(dtypes.int, {n3.as_long()})\n" +\ + "num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify()\n" +\ + "rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify()\n" +\ + "assert num==rn, f\"{num} != {rn}\"\n" + + if DEBUG >= 2: print(f"validated {expr.render()}") + print(f"Skipped {skipped} expressions due to timeout") diff --git a/test/unit/test_symbolic_failures.py b/test/unit/test_symbolic_failures.py new file mode 100644 index 0000000000..0fc5ee2467 --- /dev/null +++ b/test/unit/test_symbolic_failures.py @@ -0,0 +1,161 @@ +import unittest +from tinygrad import Variable, dtypes +from tinygrad.helpers import Context +from tinygrad.uop.ops import Ops, UOp + + +class TestFuzzFailure(unittest.TestCase): + def setUp(self): + self.context = Context(CORRECT_DIVMOD_FOLDING=1) + self.context.__enter__() + + def tearDown(self): + self.context.__exit__(None, None, None) + + def test_fuzz_failure1(self): + v1=Variable('v1', 0, 8) + v2=Variable('v2', 0, 2) + v3=Variable('v3', 0, 1) + expr = (((((((((((((((((((((((0//4)%2)//8)+-2)+-4)+-3)+v1)+-4)+v2)+-2)+v3)+v2)//3)%7)*1)//2)+v2)*-1)+2)+1)+0)+-3)+v3) + v1_val, v2_val, v3_val = v1.const_like(8), v2.const_like(0), v3.const_like(0) + num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + self.assertEqual(num, rn) + + def test_fuzz_failure2(self): + v1=Variable('v1', 0, 16) + v2=Variable('v2', 0, 5) + v3=Variable('v3', 0, 3) + expr = (((((((((((((((((((((((((0*4)//5)*2)*-1)*-2)+-4)*4)*2)*3)*4)+-4)*4)+v2)+v2)+v3)//3)+v2)+v1)//9)+3)+1)//1)+-4)//4)*2) + expr = (((((v1+(v2+(((v3+(v2*2))+1)//3)))+4)//9)+-57)//(9*4)) + v1_val, v2_val, v3_val = v1.const_like(6), v2.const_like(0), v3.const_like(0) + num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + self.assertEqual(num, rn) + + def test_fuzz_failure3(self): + v1=Variable('v1', 0, 2) + v2=Variable('v2', 0, 1) + v3=Variable('v3', 0, 2) + expr = (((((((((((((((((((0//2)//3)+v3)+0)+-4)*-2)*-2)+-1)+2)+3)+v3)+0)//8)*-3)+0)*-2)*-4)*-2)//5) + v1_val, v2_val, v3_val = v1.const_like(0), v2.const_like(0), v3.const_like(0) + num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + self.assertEqual(num, rn) + + def test_fuzz_failure4(self): + v1=Variable('v1', 0, 2) + v2=Variable('v2', 0, 3) + v3=Variable('v3', 0, 4) + expr = (((((((((((((((((((((((((((((0*-2)+0)*-1)//9)//6)//8)+v1)*-4)+v2)//4)//8)+4)*3)+v1)+v3)//8)//7)+4)+v3)*-4)+1)+v1)*3)+4)*2)//5)//2)//3)*-4) + v1_val, v2_val, v3_val = v1.const_like(2), v2.const_like(0), v3.const_like(2) + num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + self.assertEqual(num, rn) + + def test_fuzz_failure5(self): + v1=Variable('v1', 0, 1) + v2=Variable('v2', 0, 1) + v3=Variable('v3', 0, 3) + expr = ((((((((((((((0+v2)+v1)*0)+v2)//1)//7)+-2)+v2)+v1)*4)+-3)//5)+v2)+1) + v1_val, v2_val, v3_val = v1.const_like(0), v2.const_like(0), v3.const_like(0) + num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + self.assertEqual(num, rn) + + def test_fuzz_failure6(self): + v1=Variable('v1', 0, 8) + v2=Variable('v2', 0, 64) + v3=Variable('v3', 0, 128) + expr = (((((((((((((((((((((((((((((0//3)+4)+v1)//2)+-1)//1)*1)*-1)*4)//5)+v1)//6)+v1)*-1)+-4)+v2)+-2)*-3)+v3)+-4)+-2)*-1)//8)//4)*-4)+3)+v3)* + -2)+v2) + v1_val, v2_val, v3_val = v1.const_like(8), v2.const_like(3), v3.const_like(2) + num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + self.assertEqual(num, rn) + + def test_fuzz_failure7(self): + v1=Variable('v1', 0, 64) + v2=Variable('v2', 0, 5) + v3=Variable('v3', 0, 128) + expr = (((((((((((((((((((((((((((((0+v2)*-4)+0)//9)+-4)*-2)*3)*4)//9)+v3)+v1)//4)+v1)+v3)+-1)*4)//4)+v2)//7)//3)+v1)+v2)+v3)+1)*2)//4)*3)+-1)*1) + v1_val, v2_val, v3_val = v1.const_like(0), v2.const_like(2), v3.const_like(65) + num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + self.assertEqual(num, rn) + + def test_fuzz_failure8(self): + v1=Variable('v1', 0, 2) + v2=Variable('v2', 0, 8) + v3=Variable('v3', 0, 9) + expr = (((((((0+-1)+2)+v1)*-2)//3)+v1)*-4) + v1_val, v2_val, v3_val = v1.const_like(0), v2.const_like(0), v3.const_like(0) + num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + self.assertEqual(num, rn) + + def test_fuzz_failure9(self): + v1=Variable('v1', 0, 256) + v2=Variable('v2', 0, 1) + v3=Variable('v3', 0, 8) + expr = (((((((((((((((((((((((((((((0*-2)//1)+3)*-2)+-3)*-4)*1)+v1)+0)%2)%8)%9)+v2)%9)+-4)//4)+-1)*-2)+0)+v1)+v1)+3)+v1)+4)+-4)+0)*2)+-3)%6) + v1_val, v2_val, v3_val = v1.const_like(0), v2.const_like(1), v3.const_like(0) + num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + self.assertEqual(num, rn) + + def test_fuzz_failure10(self): + v1=Variable("v1", 0, 256) + v2=Variable("v2", 0, 32) + v3=Variable("v3", 0, 32) + expr = UOp(Ops.MUL, dtypes.int, arg=None, src=( + UOp(Ops.MAX, dtypes.int, arg=None, src=( + UOp(Ops.MUL, dtypes.int, arg=None, src=( + UOp(Ops.WHERE, dtypes.int, arg=None, src=( + UOp(Ops.CMPNE, dtypes.bool, arg=None, src=( + UOp(Ops.CMPLT, dtypes.bool, arg=None, src=( + x5:=UOp(Ops.IDIV, dtypes.int, arg=None, src=( + UOp(Ops.WHERE, dtypes.int, arg=None, src=( + UOp(Ops.CMPNE, dtypes.bool, arg=None, src=( + UOp(Ops.CMPLT, dtypes.bool, arg=None, src=( + x9:=UOp(Ops.CONST, dtypes.int, arg=9, src=()), + x10:=UOp(Ops.DEFINE_VAR, dtypes.int, arg=('v1', 0, 256), src=()),)), + x11:=UOp(Ops.CONST, dtypes.bool, arg=True, src=()),)), + UOp(Ops.ADD, dtypes.int, arg=None, src=( + UOp(Ops.MUL, dtypes.int, arg=None, src=( + x10, + x14:=UOp(Ops.CONST, dtypes.int, arg=-4, src=()),)), + x14,)), + UOp(Ops.IDIV, dtypes.int, arg=None, src=( + x10, + x9,)),)), + x9,)), + x14,)), + x11,)), + x5, + UOp(Ops.IDIV, dtypes.int, arg=None, src=( + UOp(Ops.ADD, dtypes.int, arg=None, src=( + UOp(Ops.MOD, dtypes.int, arg=None, src=( + x19:=UOp(Ops.DEFINE_VAR, dtypes.int, arg=('v2', 0, 32), src=()), + UOp(Ops.CONST, dtypes.int, arg=3, src=()),)), + x19,)), + UOp(Ops.CONST, dtypes.int, arg=5, src=()),)),)), + x22:=UOp(Ops.CONST, dtypes.int, arg=-1, src=()),)), + UOp(Ops.MUL, dtypes.int, arg=None, src=( + UOp(Ops.ADD, dtypes.int, arg=None, src=( + UOp(Ops.ADD, dtypes.int, arg=None, src=( + UOp(Ops.MOD, dtypes.int, arg=None, src=( + UOp(Ops.MUL, dtypes.int, arg=None, src=( + x10, + UOp(Ops.CONST, dtypes.int, arg=-2, src=()),)), + UOp(Ops.CONST, dtypes.int, arg=6, src=()),)), + UOp(Ops.MOD, dtypes.int, arg=None, src=( + UOp(Ops.DEFINE_VAR, dtypes.int, arg=('v3', 0, 32), src=()), + UOp(Ops.CONST, dtypes.int, arg=1, src=()),)),)), + UOp(Ops.CONST, dtypes.int, arg=0, src=()),)), + x22,)),)), + x22,)) + v1_val, v2_val, v3_val = UOp.const(dtypes.int, 9), UOp.const(dtypes.int, 0),UOp.const(dtypes.int, 0) + num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() + self.assertEqual(num, rn)