[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 <[email protected]>
This commit is contained in:
Sieds Lykles
2025-05-28 16:28:37 -04:00
committed by GitHub
co-authored by chenyu
parent 74cf5dbd9e
commit ae02a1e232
3 changed files with 240 additions and 65 deletions
+2 -2
View File
@@ -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
+77 -63
View File
@@ -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")
+161
View File
@@ -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)