diff --git a/test/null/test_uop_symbolic.py b/test/null/test_uop_symbolic.py index d921a5cb17..9334440753 100644 --- a/test/null/test_uop_symbolic.py +++ b/test/null/test_uop_symbolic.py @@ -347,6 +347,8 @@ class TestSymbolic(unittest.TestCase): def test_mul_lt(self): self.helper_test_variable(Variable("a", 0, 5)*4 < 13, 0, 1, "(a<4)") self.helper_test_variable(Variable("a", 0, 5)*4 < 16, 0, 1, "(a<4)") + self.helper_test_variable(Variable("a", -5, 5)*4 < -13, 0, 1, "(a<-3)") + self.helper_test_variable(Variable("a", -5, 5)*-4 < 13, 0, 1, "((a*-1)<4)") c0, c1 = 2, 2**54+1 self.helper_test_variable(Variable("a", 0, c1)*c0 < c1, 0, 1, f"(a<{2**53+1})") c0, c1 = -2, -(2**54-1) diff --git a/tinygrad/uop/symbolic.py b/tinygrad/uop/symbolic.py index 67b196bd55..65da3f2f56 100644 --- a/tinygrad/uop/symbolic.py +++ b/tinygrad/uop/symbolic.py @@ -262,12 +262,9 @@ symbolic = symbolic_simple+commutative+PatternMatcher([ # (x//c1)//c2 -> x//(c1*c2) for c2>0 ((UPat.var("x") // UPat.cvar("c1")) // UPat.cvar("c2"), lambda x,c1,c2: x//(c1*c2) if c2.vmin>0 else None), # ** lt ** - # c0*x sign(c0)*x < ceil(c1/abs(c0)) ((UPat.cvar("c0")*UPat.var("x", dtype=dtypes.weakint)) 0 and c1.arg > 0 else None), - # c0*x 0 else -x)<-(-c1.arg//abs(c0.arg)) if abs(c0.arg) > 1 else None), # x//d x0, and -> c*d 0 else (x>c.arg*d.arg) if d.arg < 0 else None),