diff --git a/test/test_uops.py b/test/test_uops.py index 65fef7a9a3..dfd6a7f967 100644 --- a/test/test_uops.py +++ b/test/test_uops.py @@ -402,6 +402,14 @@ class TestAssembly(unittest.TestCase): self.assertIn(Ops.SHR, ops) self.assertNotIn(Ops.IDIV, ops) + def test_fast_idiv_remove_powers_of_two(self): + ridx = UOp.range(dtypes.int, 2**20, 0) + uops = to_uops_list([ridx//(7*64)], opts=Device[Device.DEFAULT].renderer) + ops = [x.op for x in uops] + # this requires shifting out the powers of two before doing fast_idiv + # (((ridx0>>6)*18725)>>17) instead of (int)((((long)(ridx0)*1198373)>>29)) + self.assertNotIn(Ops.CAST, ops) + def test_mulacc_unrolled(self): # test that acc = acc + a0*b0 + a1*b1 + a2*b2 + a3*b3 # is not acc = acc + (a0*b0 + a1*b1 + a2*b2 + a3*b3) diff --git a/tinygrad/uop/decompositions.py b/tinygrad/uop/decompositions.py index 4431ea180c..dd807f905c 100644 --- a/tinygrad/uop/decompositions.py +++ b/tinygrad/uop/decompositions.py @@ -280,7 +280,7 @@ def magicgu(vmax:int, d:int) -> tuple[int,int]: return m, s assert False -def fast_idiv(device: str, x: UOp, d: int) -> UOp|None: +def fast_idiv(device: str, x: UOp, d: int, dont_cast=False) -> UOp|None: # If d is a power of two this is not valid for signed ints! is_unsigned = True if x.vmin>=0 or x.dtype in dtypes.uints else False assert d>0, "Sign should have been taken out of divisor" @@ -288,6 +288,10 @@ def fast_idiv(device: str, x: UOp, d: int) -> UOp|None: m,s = magicgu(max(vmax, abs(vmin)), d) if m*vmin >= dtypes.min(x.dtype) and m*vmax <= dtypes.max(x.dtype): return ((x*m) >> s) if is_unsigned else ((x*m) >> s) + (x<0).where(x.ufix(1), 0) + # before we try casting to a larger dtype (slow), we see if there are powers of two in d we can shift to make x smaller + if (largest_factor_of_two_in_d := (d & -d)) > 1: + if (ret:=fast_idiv(device, x//largest_factor_of_two_in_d, d//largest_factor_of_two_in_d, dont_cast=True)) is not None: return ret + if dont_cast: return None # promo_lattice needs to return an unsigned type if the type is unsigned if dtypes.is_int(next_dtype := promo_lattice[x.dtype][-1]) and is_dtype_supported(next_dtype, None if device=='' else device): if m*vmin >= dtypes.min(next_dtype) and m*vmax <= dtypes.max(next_dtype):