forked from tinygrad/tinygrad
divide by gcd in UOp div folding (#5949)
* divide by gcd in UOp div folding `(6x+6y)//16 -> (3x+3y)//8` etc simpler version * only factor out const * don't apply for unsigned * don't need that if * space
This commit is contained in:
@@ -203,7 +203,6 @@ class TestSymbolic(unittest.TestCase):
|
||||
def test_sum_div_some_factor(self):
|
||||
self.helper_test_variable(Node.sum([Variable("a", 0, 7)*5, Variable("b", 0, 3)*4]) // 2, 0, 23, {"(((a*5)//2)+(b*2))", "((b*2)+((a*5)//2))"})
|
||||
|
||||
@unittest.expectedFailure
|
||||
def test_sum_div_some_partial_factor(self):
|
||||
self.helper_test_variable(Node.sum([Variable("a", 0, 7)*6, Variable("b", 0, 7)*6]) // 16, 0, 5, "(((a*3)+(b*3))//8)")
|
||||
self.helper_test_variable(Node.sum([NumNode(16), Variable("a", 0, 7)*6, Variable("b", 0, 7)*6]) // 16, 1, 6, "((((a*3)+(b*3))//8)+1)")
|
||||
@@ -337,12 +336,11 @@ class TestSymbolic(unittest.TestCase):
|
||||
self.helper_test_variable((-Variable("idx", 0, 100)+200)//-4 + 50, 0, 25, "((((-idx)+200)//(-4))+50)")
|
||||
self.helper_test_variable((-Variable("idx", 0, 100)+201)//-4 + 50, 0, 25, "((((-idx)+201)//(-4))+50)")
|
||||
|
||||
@unittest.expectedFailure
|
||||
def test_sum_div_big_const(self):
|
||||
gidx0 = Variable("gidx0", 0, 24)
|
||||
self.helper_test_variable((gidx0+19)//20, 0, 2, "((19+gidx0)//20)")
|
||||
self.helper_test_variable((gidx0+19)//20, 0, 2, {"((19+gidx0)//20)", "((gidx0+19)//20)"})
|
||||
self.helper_test_variable((gidx0+20)//20, 1, 2, "((gidx0//20)+1)")
|
||||
self.helper_test_variable((gidx0+21)//20, 1, 2, "(((1+gidx0)//20)+1)")
|
||||
self.helper_test_variable((gidx0+21)//20, 1, 2, {"(((1+gidx0)//20)+1)", "(((gidx0+1)//20)+1)"})
|
||||
|
||||
@unittest.expectedFailure
|
||||
def test_sum_div_complex1(self):
|
||||
|
||||
@@ -97,18 +97,26 @@ def mod_folding(x:UOp, c:int) -> Optional[UOp]:
|
||||
|
||||
def div_folding(x:UOp, c:int) -> Optional[UOp]:
|
||||
# simplify x // c, None means no change
|
||||
quotient, remainder, something_changed = [], [], False
|
||||
quotient, remainder, something_changed, gcd = [], [], False, c
|
||||
for u in _get_add_chain(x):
|
||||
if (factor:=u.const_factor())%c == 0:
|
||||
if factor: quotient.append(u.divides(c))
|
||||
something_changed = True
|
||||
else: remainder.append(u)
|
||||
if not something_changed: return None
|
||||
else:
|
||||
# (5*a+5)//2 -> (5*a+1)//2+2, only factor out const
|
||||
if u.op is UOps.CONST and c <= abs(u.arg):
|
||||
quotient.append(u.const(u.arg//c))
|
||||
remainder.append(u.const(u.arg%c))
|
||||
something_changed = True
|
||||
else: remainder.append(u)
|
||||
gcd = math.gcd(gcd, factor)
|
||||
|
||||
if not something_changed: return cast(UOp, x.divides(gcd))//(c//gcd) if gcd != c and gcd != 1 else None
|
||||
rem:Optional[UOp] = functools.reduce(operator.add, remainder) if remainder else None
|
||||
if rem is not None and 0 <= rem.vmin.arg and rem.vmax.arg < c: rem = None
|
||||
quo:Optional[UOp] = functools.reduce(operator.add, quotient) if quotient else None
|
||||
if quo is None: return x.const(0) if rem is None else rem//c
|
||||
return quo if rem is None else quo+rem//c
|
||||
if quo is None: return x.const(0) if rem is None else cast(UOp, rem.divides(gcd))//(c//gcd)
|
||||
return quo if rem is None else quo+cast(UOp, rem.divides(gcd))//(c//gcd)
|
||||
|
||||
# ***** transcendental *****
|
||||
|
||||
@@ -260,7 +268,8 @@ constant_folder = PatternMatcher([
|
||||
(NOp.lt(-NOp.var('x'), NOp.cvar('c', dtypes.int)), lambda c,x: UOp.lt(c.const(-c.arg), x)),
|
||||
# ** div **
|
||||
# # div folding
|
||||
(NOp.var('x') // NOp.cvar('c'), lambda x,c: newx if 0 < c.arg and (newx:=div_folding(x,c.arg)) is not None else None),
|
||||
(NOp.var('x') // NOp.cvar('c'), lambda x,c:
|
||||
newx if 0 < c.arg and not dtypes.is_unsigned(x.dtype) and (newx:=div_folding(x,c.arg)) is not None else None),
|
||||
# mul div
|
||||
((NOp.var("x") * NOp.cvar("c0")) // NOp.cvar("c1"),
|
||||
lambda x,c0,c1: x*(c0.arg//gcd)//(c1.arg//gcd) if c1.arg!=0 and (gcd:=math.gcd(c0.arg,c1.arg))> 1 else None),
|
||||
|
||||
Reference in New Issue
Block a user