diff --git a/test/unit/test_symbolic_shapetracker.py b/test/unit/test_symbolic_shapetracker.py index 7d77e81b97..565408cc62 100644 --- a/test/unit/test_symbolic_shapetracker.py +++ b/test/unit/test_symbolic_shapetracker.py @@ -13,7 +13,6 @@ class TestSymbolic(unittest.TestCase): assert st.shape == (x, 3) assert st.real_strides() == (3, 1) - @unittest.expectedFailure def test_real_strides_0(self): st = ShapeTracker(views=(View(shape=(2, (Variable('start_pos', 1, 8)+1), 1, 1), strides=(8, 1, 0, 0), offset=0, mask=((0, 2), (0, Variable('start_pos', 1, 8)), (0, 1), (0, 1)), contiguous=False), View(shape=(2, (Variable('start_pos', 1, 8)+1)), strides=((Variable('start_pos', 1, 8)+1), 1), offset=0, mask=None, contiguous=True))) # noqa: E501 self.assertEqual(st.real_strides(), (8, None)) diff --git a/test/unit/test_uop_symbolic.py b/test/unit/test_uop_symbolic.py index 949e37d68f..b3e393a61a 100644 --- a/test/unit/test_uop_symbolic.py +++ b/test/unit/test_uop_symbolic.py @@ -93,6 +93,37 @@ class TestSymbolic(unittest.TestCase): assert idx1+idx2 is not idx2 assert idx1*idx2 is not idx2*idx1 + def test_uop_gcd_method(self): + a = Variable("a", 0, 8) + b = Variable("b", 0, 8) + self.assertEqual(UOp.gcd(a, a*b, a*3).simplify(), a) + self.assertEqual(UOp.gcd(a*a*a, a*b*a, a*3*a).simplify(), a*a) + self.assertEqual(UOp.gcd(a*a*10, b*a*5, a*a*5).simplify(), a*5) + self.assertEqual(UOp.gcd(a*10, b*5, a*5).simplify(), a.const_like(5)) + self.assertEqual(UOp.gcd(a, b*5, a*5).simplify(), a.const_like(1)) + + def test_divides_exact(self): + a = Variable("a", 1, 8) + b = Variable("b", 1, 8) + self.assertEqual((a*a*3).divide_exact(a).simplify(), a*3) + self.assertEqual((a*a*3).divide_exact(a*a*3).simplify(), a.const_like(1)) + self.assertEqual((a*b*3).divide_exact(a.const_like(3)).simplify(), a*b) + self.assertEqual((a*a*3).divide_exact(a*a.const_like(-3)).simplify(), a*-1) + self.assertEqual((a*a*b*3).divide_exact(a*b).simplify(), a*3) + self.assertEqual((a*3+a*b).divide_exact(a).simplify(), b+3) + self.assertEqual((a*b*3+a*b*b).divide_exact(a*b).simplify(), b+3) + self.assertEqual((((a*-2)+14)*b).divide_exact(((a*-2)+14)).simplify(), b) + + def test_divide_exact_not(self): + a = Variable("a", 1, 8) + b = Variable("b", 1, 8) + x = Variable("x", -20, 0) + self.assertEqual((a).divide_exact(b), None) + self.assertEqual((a+2).divide_exact(a), None) + self.assertEqual((x*-1).divide_exact(a), None) + self.assertEqual((a*5).divide_exact(a*10), None) + self.assertEqual((a*10-1).divide_exact(a*10), None) + def test_factorize(self): a = Variable("a", 0, 8) b = Variable("b", 0, 8) @@ -450,6 +481,33 @@ class TestSymbolic(unittest.TestCase): def test_mul_div_factor_div_neg(self): self.helper_test_variable((Variable("a", 0, 10)*-4+4)//8, -4, 0, "(((a*-1)+1)//2)") + def test_div_symbolic_const_gcd(self): + a = Variable("a", -10, 10) + b = Variable("b", -10, 10) + d = Variable("d", 1, 10) + self.helper_test_variable((3*a+9*b)//(3*d), -40, 40, "((a+(b*3))//d)") + + def test_symbolic_gcd_div(self): + a = Variable("a", -10, 10) + b = Variable("b", -10, 10) + c = Variable("c", -10, 10) + d1 = Variable("d1", 1, 10) + d2 = Variable("d2", -10, -1) + self.helper_test_variable((d1*a*b*d1)//(d1), -1000, 1000, "(a*(b*d1))") + self.helper_test_variable((d1*a*d2*b*d1)//(d1*d2), -1000, 1000, "(a*(b*d1))") + self.helper_test_variable((d1*a + b*d1)//(d1), -20, 20, "(a+b)") + self.helper_test_variable((d1*a + b*d1 + c*d1)//(d1), -30, 30, "(c+(a+b))") + self.helper_test_variable((3*a*d1 + 9*b*d1)//(3*d1*d2), -40, 40, "(((a+(b*3))//(d2*-1))*-1)") + self.helper_test_variable((3*a*d1 + 9*b*d1+3)//(3*d1*d2), -401, 399, "(((((a*d1)+((b*d1)*3))+1)//((d1*d2)*-1))*-1)") + + def test_symbolic_factor_remainder_div(self): + a = Variable("a", 0, 10) + b = Variable("b", 0, 10) + d = Variable("d", 1, 10) + self.helper_test_variable((d*a+b)//d, 0, 20, "(a+(b//d))") + self.helper_test_variable((d*a*20+b)//(5*d), 0, 42, "((a*4)+(b//(d*5)))") + self.helper_test_variable((d*a*20+b*d*5+10)//(5*d), 0, 52, "((b+(a*4))+(2//d))") + def test_mod_gcd_factor_neg(self): self.helper_test_variable((Variable("a", 0, 10)*-4+4)%8, -4, 4, "((((a*-1)+1)%2)*4)") diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index 46c441d798..0b251805b8 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -1,6 +1,6 @@ from __future__ import annotations from typing import Any, Callable, cast, TYPE_CHECKING, Type, Sequence -import sys, time, functools, itertools, math, operator, hashlib, os, types, pickle, pathlib, inspect, weakref +import sys, time, functools, itertools, math, operator, hashlib, os, types, pickle, pathlib, inspect, weakref, collections from dataclasses import dataclass, field from enum import Enum, auto from tinygrad.uop import Ops, GroupOp @@ -549,7 +549,23 @@ class UOp(MathTrait, metaclass=UOpMetaClass): if (d0:=self.src[0].divides(v)) is not None: return d0 * self.src[1] if (d1:=self.src[1].divides(v)) is not None: return self.src[0] * d1 return None # generic None if we aren't sure - def pop_const(self) -> tuple[UOp, int]: return (self.src[0], self.src[1].arg) if self.op is Ops.ADD and self.src[1].op is Ops.CONST else (self, 0) + def pop_const(self, op=Ops.ADD) -> tuple[UOp, ConstType]: + return (self.src[0], self.src[1].arg) if self.op is op and self.src[1].op is Ops.CONST else (self, identity_element(op, self.dtype)) + @staticmethod + def gcd(*uops: UOp) -> UOp: + terms, factors = zip(*[(u.divides(f:=u.const_factor()),f) for u in uops]) + count = functools.reduce(operator.and_, [collections.Counter(term.split_uop(Ops.MUL)) for term in terms]) + return math.prod([*count.elements(), terms[0].const_like(math.gcd(*factors))]) # put the const at the top + def divide_exact(self, v:UOp) -> UOp|None: + if self is v: return self.const_like(1) + if self.op is Ops.ADD: return None if (s0:=self.src[0].divide_exact(v)) is None or (s1:=self.src[1].divide_exact(v)) is None else s0+s1 + if v.op is Ops.CONST: return self.divides(v.arg) + if self.op is Ops.MUL: + (fac, const), (div_fac, div_const) = self.pop_const(Ops.MUL), v.pop_const(Ops.MUL) + new_count = collections.Counter(fac.split_uop(Ops.MUL)) + new_count.subtract(div_fac.split_uop(Ops.MUL)) + if const%div_const==0 and all(v>=0 for v in new_count.values()): return math.prod([*new_count.elements(), self.const_like(const//div_const)]) + return None # generic None if we aren't sure @property def vmin(self) -> ConstType: return self._min_max[0] @property diff --git a/tinygrad/uop/symbolic.py b/tinygrad/uop/symbolic.py index 641fb28159..4fe4081a85 100644 --- a/tinygrad/uop/symbolic.py +++ b/tinygrad/uop/symbolic.py @@ -4,7 +4,7 @@ import math, operator, struct, functools from collections import defaultdict from tinygrad.uop.ops import Ops, PatternMatcher, UPat, UOp, GroupOp, exec_alu from tinygrad.dtype import ConstType, dtypes, PtrDType, AddrSpace, can_safe_cast, Invalid -from tinygrad.helpers import partition, all_same, prod, flatten, get_single_element, cdiv, cmod, CORRECT_DIVMOD_FOLDING +from tinygrad.helpers import partition, all_same, prod, flatten, get_single_element, cdiv, cmod, CORRECT_DIVMOD_FOLDING, unwrap from tinygrad.uop.decompositions import xpow # ******** phase 1 of symbolic used to live in ops, it's the most generic folding rules ******** @@ -164,7 +164,7 @@ def remove_nested_mod(m: UOp, x: UOp, y: UOp) -> UOp|None: def fold_binary_numerator(d: UOp, x: UOp, y: UOp) -> UOp|None: # we can fold if the expression has only one non-constant term and this term can only take on two values - if ((c := y.arg) < 0) or (x.dtype.count > 1): return None + if ((c := y.arg) < 0): return None x,const = x.pop_const() terms, factors = zip(*[(u.divides(f:=u.const_factor()),f) for u in x.split_uop(Ops.ADD)]) if len(terms)==1 and (v:=terms[0]).vmax-v.vmin == 1: @@ -175,7 +175,7 @@ def fold_binary_numerator(d: UOp, x: UOp, y: UOp) -> UOp|None: def fold_divmod_congruence(d: UOp, x: UOp, y: UOp) -> UOp|None: # within a mod we can freely subtract multiples of c, we use this to see if a is congruent to an expression whose vmin/vmax are between 0 and c - if (x.vmin<0 and CORRECT_DIVMOD_FOLDING) or ((c := y.arg) < 0) or (x.dtype.count > 1): return None + if (x.vmin<0 and CORRECT_DIVMOD_FOLDING) or ((c := y.arg) < 0): return None x,const = x.pop_const() terms, factors = zip(*[(u.divides(f:=u.const_factor()),f) for u in x.split_uop(Ops.ADD)]) # a//c = (a-a%c)/c, if we can fold a%c, we can fold a//c @@ -186,14 +186,28 @@ def fold_divmod_congruence(d: UOp, x: UOp, y: UOp) -> UOp|None: def divide_by_gcd(d: UOp, x: UOp, y: UOp) -> UOp|None: # x//y -> (x//gcd)//(y//gcd) or x%y -> gcd*(x//gcd)%(y//gcd) - terms, factors = zip(*[(u.divides(f:=u.const_factor()),f) for u in x.split_uop(Ops.ADD)]) - if (gcd := math.gcd(y.arg, *factors)) == 1: return None - ret = sum(f//gcd * v for f,v in zip(factors, terms)).alu(d.op, y.const_like(y.arg//gcd)) + gcd = UOp.gcd(*x.split_uop(Ops.ADD), y).simplify() + if gcd.op is Ops.CONST and gcd.arg==1: return None + ret = unwrap(x.divide_exact(gcd)).alu(d.op, unwrap(y.divide_exact(gcd))) return ret*gcd if d.op is Ops.MOD else ret +def gcd_with_remainder(d: UOp, x: UOp, y: UOp): + # (gcd*x+r)//(gcd*d) -> (x+(r%d)//gcd)//d + r//(gcd*d) + # (gcd*x+r)%(gcd*d) -> gcd*(x+(r%d)//gcd)%d + r%gcd + # These only work for floordiv (and the corresponding remainder)! Thats why we check the sign of x,y and new_x + if ((c := y.arg) < 0) or x.vmin<0: return None + x_no_const, const = x.pop_const() + gcd = UOp.gcd(*x_no_const.split_uop(Ops.ADD), y).simplify() + assert gcd.op is Ops.CONST + if gcd.arg==1: return None + new_x = unwrap(x_no_const.divide_exact(gcd)).simplify() + (const%c)//gcd + if new_x.vmin<0: return None + ret = new_x.alu(d.op, x.ufix(c//gcd.arg)) + return ret*gcd + const%gcd.arg if d.op is Ops.MOD else ret+const//c + def nest_div_by_smallest_factor(d: UOp, x: UOp, y: UOp) -> UOp|None: # we try and nest the div and see if it allows the numerator to be simplified - if ((c := y.arg) < 0) or (x.dtype.count > 1): return None + if ((c := y.arg) < 0): return None factors = [u.const_factor() for u in x.pop_const()[0].split_uop(Ops.ADD)] # div is the smallest factor of the denominator (greater than 1) out of all "factors" # TODO: there are better ways to pick `div`, this sometimes adds extra divisions @@ -202,27 +216,22 @@ def nest_div_by_smallest_factor(d: UOp, x: UOp, y: UOp) -> UOp|None: if (1 < div < c) and (newxs:=(newx:=(x//div)).simplify()) is not newx and x.vmin>=0 and newx.vmin>=0: return newxs//(c//div) return None -def simplify_remainder(d: UOp, x: UOp, y: UOp) -> UOp|None: - # we try and take out the quotient and see if it allows the numerator to be simplified - if ((c := y.arg) < 0) or (x.dtype.count > 1): return None - x_no_const,const = x.pop_const() - terms, factors = zip(*[(u.divides(f:=u.const_factor()),f) for u in x_no_const.split_uop(Ops.ADD)]) - quotients, remainders = zip(*[divmod(f, c) for f in factors]) - gcd = math.gcd(c, *remainders) # gcd without const! - if const%c==const and gcd==1 and not any(r==0 or (r!=f and d.op is Ops.MOD) for r,f in zip(remainders, factors)): return None - - quo, rem = x.const_like(const//c), x.const_like((const%c)//gcd) - for q,r,f,v in zip(quotients, remainders, factors, terms): - if d.op is Ops.IDIV and r!=0: - rem += f//gcd * v - else: - rem += r//gcd * v - quo += q * v - - # if numerator before/after is negative, and it has remainder, don't simplify because C divmod is different from python divmod. - if (x.vmin < 0 or rem.vmin < 0) and remainders: return None - if d.op is Ops.MOD: return gcd*(rem % (c//gcd)) + const%gcd - return rem//(c//gcd)+quo +def factor_remainder(d: UOp, x: UOp, y: UOp) -> UOp|None: + # (d*x+y)//d -> x+y//d or (d*x+y)%d + # for mod we go further and take the remainder of all factors to reduce their size + # These only work for floordiv (and the corresponding remainder)! Thats why we check the sign of x,y and new_x + if y.vmin<0 or x.vmin<0: return None + quo, rem = [], [] + for u in x.split_uop(Ops.ADD): + if (q:=u.divide_exact(y)) is not None: quo.append(q) + # if this is mod and y is a const, we can make the remainder factor sm + elif d.op is Ops.MOD and y.op is Ops.CONST and (c:=u.const_factor())%y.arg!=c: + rem.append(u.divides(c)*(c%y.arg)) + quo.append(u.const_like(0)) # we append this so we can check if something changed + else: rem.append(u) + new_x = sum(rem)+x.const_like(0) + if len(quo)==0 or new_x.vmin<0: return None + return new_x%y if d.op is Ops.MOD else new_x//y+sum(quo) def gep_through_wmma(gep:UOp, wmma:UOp): out_sz = prod(x[1] for x in wmma.arg[6][-1]) @@ -334,14 +343,17 @@ symbolic = symbolic_simple+commutative+PatternMatcher([ (UPat(Ops.RANGE, src=UPat.var("end"), name="r")%UPat.var("end"), lambda r,end: r), (UPat(Ops.RANGE, src=UPat.var("end"), name="r")//UPat.var("end"), lambda r,end: r.const_like(0)), (UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.var("y"))), cancel_divmod), + (UPat.var("x") // UPat.var("d"), lambda x,d: -(x//(-d)) if d.vmax < 0 else None), (UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.cvar("y", vec=False))), fold_binary_numerator), (UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.cvar("y", vec=False))), fold_divmod_congruence), - (UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.cvar("y", vec=False))), divide_by_gcd), + (UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.var("y"))), divide_by_gcd), + (UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.cvar("y", vec=False))), gcd_with_remainder), (UPat(Ops.MOD, dtypes.index, name="m", src=(UPat.var("x"), UPat.cvar("y", vec=False))), remove_nested_mod), (UPat((Ops.IDIV), dtypes.index, name="d", src=(UPat.var("x"), UPat.cvar("y", vec=False))), nest_div_by_smallest_factor), - (UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.cvar("y", vec=False))), simplify_remainder), - (UPat.var("x") // UPat.var("d"), lambda x,d: -(x//(-d)) if d.vmax < 0 else None), + (UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.var("y"))), factor_remainder), (UPat.var("x") // UPat.var("d"), lambda x,d: -((-x)//d) if x.vmax <=0 else None), + ((UPat.var("x", dtypes.index)+UPat.cvar("c", vec=False)).named("n")//UPat.cvar("d", vec=False), + lambda x,c,n,d: ((x+c.arg%d.arg)//d + c.arg//d.arg) if c.arg%d.arg!=c.arg and x.vmin>=0 and n.vmin>=0 and d.arg>0 else None), ((UPat.var("x", dtypes.index)+UPat.cvar("c", vec=False)).named("n")//UPat.cvar("d", vec=False), lambda x,c,n,d: (-(-(c.arg%d.arg + x - (d.arg-1))//d) + c.arg//d.arg) if x.vmax<=0 and n.vmin>=0 and d.arg>0 else None), # ** mod **