diff --git a/setup.py b/setup.py index 2b8c70bd63..9aaaaac244 100644 --- a/setup.py +++ b/setup.py @@ -48,7 +48,8 @@ setup(name='tinygrad', 'testing_unit': testing_minimal + [ "tqdm", "safetensors", - "tabulate" # for sz.py + "z3-solver", + "tabulate", # for sz.py ], 'testing': testing_minimal + [ "pillow", diff --git a/test/unit/test_simplify_valid_idx.py b/test/unit/test_simplify_valid_idx.py index 9093f4d037..b13f9d87f4 100644 --- a/test/unit/test_simplify_valid_idx.py +++ b/test/unit/test_simplify_valid_idx.py @@ -138,6 +138,33 @@ class TestValidIdxSimplification(unittest.TestCase): "(ridx0*1568)", "((ridx2<1)&(ridx1<6))") + def test_valid_becomes_const1_z3(self): + from z3 import Ints, Solver, And, If, Not, unsat + ridx0, ridx1, ridx2, alu11, alu15 = Ints('ridx0 ridx1 ridx2 alu11 alu15') + alu11 = (ridx1+ridx2) + alu15 = ((alu11+1)/7) + idx = (alu15*-31)+(((((alu11+218)/224)+ridx0)%30)*1568) + valid = (ridx2<1)&(ridx1<6) + load = If(valid, idx, 0) + + # correct simplification + s = Solver() + s.add(And(0<=ridx0, ridx0<30, 0<=ridx1, ridx1<7, 0<=ridx2, ridx2<2)) + simplifed_idx = (ridx0*1568) + simplifed_load = If(valid, simplifed_idx, 0) + s.add(Not(load == simplifed_load)) # Check if they are NOT equivalent + assert s.check() == unsat, f"The expressions are not equivalent. {s.model()=}" + + # new solver for a wrong simplified expression + s = Solver() + s.add(And(0<=ridx0, ridx0<30, 0<=ridx1, ridx1<7, 0<=ridx2, ridx2<2)) + wrong_simplifed_idx = (ridx0*1567)+ridx1 + wrong_simplifed_load = If(valid, wrong_simplifed_idx, 0) + s.add(Not(load == wrong_simplifed_load)) # Check if they are NOT equivalent + assert s.check() != unsat, "The expressions are equivalent??" + print("The expressions are not equivalent.") + print(s.model()) + class TestImageSimplification(unittest.TestCase): def check(self, load, svalid, sidx0, sidx1): load = full_graph_rewrite(load.sink()).src[0]