move cancel mod pattern into mod_folding (#6100)

changed some kernel in a good way because x does not go through add chain
This commit is contained in:
chenyu
2024-08-15 19:04:18 -04:00
committed by GitHub
parent 11d62668a3
commit e4a7869893
+6 -5
View File
@@ -84,8 +84,11 @@ def _get_add_chain(x:UOp):
else: yield x
def mod_folding(x:UOp, c:int) -> Optional[UOp]:
# simplify x % c
# None means no change
# simplify x % c, None means no change
# simple cancel mod case
if 0 < c and 0 <= x.vmin.arg and (quotient:=x.vmin.arg//c) == x.vmax.arg//c: return x-quotient*c
remainder, something_changed = [], False
for u in _get_add_chain(x):
if (factor:=u.const_factor())%c != factor:
@@ -100,6 +103,7 @@ 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
# simple cancel div case
if 0 <= x.vmin.arg and x.vmax.arg < c: return x.const(0)
@@ -295,9 +299,6 @@ constant_folder = PatternMatcher([
# ** mod **
# mod folding
(NOp.var('x') % NOp.cvar('c'), lambda x,c: newx if 0 < c.arg and (newx:=mod_folding(x,c.arg)) is not None else None),
# remove mod
(NOp.var('x') % NOp.cvar('c'), lambda x,c:\
x-(x.vmin.arg//c.arg)*c.arg if 0 < c.arg and 0 <= x.vmin.arg and x.vmin.arg//c.arg == x.vmax.arg//c.arg else None),
# mul mod
((NOp.cvar('c0')*NOp.var('x')) % NOp.cvar('c1'), lambda x,c0,c1: (x%(c1.arg//c0.arg))*c0 if c1.arg%c0.arg == 0 else None),
# (x%c)+(x//c)*c = x