From 20c45eb705d6c83f10a8ef93afba69db02b5af9a Mon Sep 17 00:00:00 2001 From: George Hotz Date: Fri, 20 Feb 2026 18:02:26 +0800 Subject: [PATCH] guard c1.arg --- test/null/test_uop_symbolic.py | 4 ++++ tinygrad/uop/symbolic.py | 6 +++--- 2 files changed, 7 insertions(+), 3 deletions(-) diff --git a/test/null/test_uop_symbolic.py b/test/null/test_uop_symbolic.py index 79e98ad40b..e25dae745e 100644 --- a/test/null/test_uop_symbolic.py +++ b/test/null/test_uop_symbolic.py @@ -664,6 +664,10 @@ class TestSymbolic(unittest.TestCase): # result is x//a*c2 not just x x2 = Variable("x2", 0, 5*6*7-1) self.helper_test_variable(x2//7%6*14 + x2//42*84, 0, (5*6*7-1)//7*14, "(x2//7*14)") + # negative variable range + xn = Variable("x", -1000, 1000) + self.helper_test_variable(xn//3%224*3 + xn%3 + xn//672*672, -1000, 1000, "x") + self.helper_test_variable(xn//3%7*3 + xn//21*21, -999, 999, "(x//3*3)") # should NOT simplify: a*c1 != b (3*224 != 600) self.helper_test_variable(gidx//3%224*3 + gidx//600*600, 0, 150669, "(gidx//600*600+gidx//3%224*3)") # should NOT simplify: c1*c2 != c3 (224*3 != 700) diff --git a/tinygrad/uop/symbolic.py b/tinygrad/uop/symbolic.py index 889980dc24..e636b8070b 100644 --- a/tinygrad/uop/symbolic.py +++ b/tinygrad/uop/symbolic.py @@ -51,7 +51,7 @@ symbolic_simple = propagate_invalid + PatternMatcher([ ((UPat.var("x")//UPat.cvar("a"))%UPat.cvar("c")+(UPat.var("x")//UPat.cvar("b"))*UPat.cvar("c"), lambda x,a,b,c: x//a if a.arg*c.arg==b.arg else None), # ((x//a)%c)+(x//a*c)*c = x//a. Note if a = 1 it degenerates to the one above ((UPat.var("x")//UPat.cvar("a"))%UPat.cvar("c1")*UPat.cvar("c2")+(UPat.var("x")//UPat.cvar("b"))*UPat.cvar("c3"), - lambda x,a,b,c1,c2,c3: x//a*c2 if a.arg*c1.arg==b.arg and c1.arg*c2.arg==c3.arg else None), + lambda x,a,b,c1,c2,c3: x//a*c2 if c1.arg>0 and a.arg*c1.arg==b.arg and c1.arg*c2.arg==c3.arg else None), ((UPat.var("x")//UPat.cvar("c1"))*UPat.cvar("c3")+UPat.var("x")%UPat.cvar("c1")*UPat.cvar("c2"), lambda x,c1,c2,c3: x*c2 if c1.arg*c2.arg==c3.arg else None), # (x%c1)*c2+(x//c1)*c3 = x*c2 if c1*c2==c3 ((UPat.var("y")+(UPat.var("x")//UPat.cvar("c"))*UPat.cvar("c"))+UPat.var("x")%UPat.cvar("c"), lambda y,x,c: y+x), @@ -61,9 +61,9 @@ symbolic_simple = propagate_invalid + PatternMatcher([ ((UPat.var("y")+UPat.var("x")%UPat.cvar("c1")*UPat.cvar("c2"))+(UPat.var("x")//UPat.cvar("c1"))*UPat.cvar("c3"), lambda y,x,c1,c2,c3: y+x*c2 if c1.arg*c2.arg==c3.arg else None), ((UPat.var("y")+(UPat.var("x")//UPat.cvar("a"))%UPat.cvar("c1")*UPat.cvar("c2"))+(UPat.var("x")//UPat.cvar("b"))*UPat.cvar("c3"), - lambda y,x,a,b,c1,c2,c3: y+x//a*c2 if a.arg*c1.arg==b.arg and c1.arg*c2.arg==c3.arg else None), + lambda y,x,a,b,c1,c2,c3: y+x//a*c2 if c1.arg>0 and a.arg*c1.arg==b.arg and c1.arg*c2.arg==c3.arg else None), ((UPat.var("y")+(UPat.var("x")//UPat.cvar("b"))*UPat.cvar("c3"))+(UPat.var("x")//UPat.cvar("a"))%UPat.cvar("c1")*UPat.cvar("c2"), - lambda y,x,a,b,c1,c2,c3: y+x//a*c2 if a.arg*c1.arg==b.arg and c1.arg*c2.arg==c3.arg else None), + lambda y,x,a,b,c1,c2,c3: y+x//a*c2 if c1.arg>0 and a.arg*c1.arg==b.arg and c1.arg*c2.arg==c3.arg else None), (UPat.var("x", dtype=dtypes.bool) & UPat.cvar("c", vec=False), lambda x,c: x if c.arg else c), (UPat.var("x", dtype=dtypes.bool) | UPat.cvar("c", vec=False), lambda x,c: c if c.arg else x), (UPat(GroupOp.Idempotent, src=(UPat.var("x"), UPat.var("x"))), lambda x: x),