forked from tinygrad/tinygrad
[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:
@@ -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
|
||||
|
||||
Vendored
+77
-63
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user