diff --git a/test/test_linearizer.py b/test/test_linearizer.py index 13a421a340..25fbfe14ce 100644 --- a/test/test_linearizer.py +++ b/test/test_linearizer.py @@ -190,7 +190,6 @@ class TestLinearizer(unittest.TestCase): helper_linearizer_ast((store, ), [dataset, idxs], wanna_output=[real_index]) # AssertionError: repeated stores in uops - @unittest.expectedFailure def test_argmax_multireduce_axis0(self): t = Tensor.randn(10, 20).realize() t_max = t.max((0,)).realize() diff --git a/test/test_linearizer_failures.py b/test/test_linearizer_failures.py index 92c21b09b6..34d8c71730 100644 --- a/test/test_linearizer_failures.py +++ b/test/test_linearizer_failures.py @@ -278,7 +278,7 @@ class TestLinearizerFailures(unittest.TestCase): LazyOp(op=BufferOps.CONST, src=(), arg=ConstBuffer(val=1.0, dtype=dtypes.float, st=ShapeTracker(views=(View(shape=(32640,), strides=(0,), offset=0, mask=None, contiguous=False),))))), arg=None), LazyOp(op=BufferOps.CONST, src=(), arg=ConstBuffer(val=0.0, dtype=dtypes.float, st=ShapeTracker(views=(View(shape=(32640,), strides=(0,), offset=0, mask=None, contiguous=False),))))), arg=None)), arg=None),), arg=(0,)),), arg=MemBuffer(idx=0, dtype=dtypes.float, st=ShapeTracker(views=(View(shape=(1,), strides=(0,), offset=0, mask=None, contiguous=True),)))),), arg=None) opts = [Opt(op=OptOps.GROUPTOP, axis=0, amt=16)] - helper_test_lin(Kernel(ast), opts=opts, failed_platforms=["METAL", "GPU", "CUDA", "AMD", "NV"]) + helper_test_lin(Kernel(ast), opts=opts, failed_platforms=[]) # from fuzzing on metal def test_failure_34(self, unroll=False): @@ -337,7 +337,7 @@ class TestLinearizerFailures(unittest.TestCase): ast = LazyOp(op=MetaOps.KERNEL, src=(LazyOp(op=BufferOps.STORE, src=(LazyOp(op=BinaryOps.ADD, src=(LazyOp(op=ReduceOps.SUM, src=(LazyOp(op=BufferOps.CONST, src=(), arg=ConstBuffer(val=1, dtype=dtypes.int, st=ShapeTracker(views=(View(shape=(60001, 119999), strides=(0, 0), offset=0, mask=((0, 60001), (59999, 119999)), contiguous=False), View(shape=(60000, 60000), strides=(1, 120000), offset=0, mask=None, contiguous=False))))),), arg=(1,)), LazyOp(op=BufferOps.CONST, src=(), arg=ConstBuffer(val=-1, dtype=dtypes.int, st=ShapeTracker(views=(View(shape=(60000, 1), strides=(0, 0), offset=0, mask=None, contiguous=False),))))), arg=None),), arg=MemBuffer(idx=0, dtype=dtypes.int, st=ShapeTracker(views=(View(shape=(60000, 1), strides=(1, 0), offset=0, mask=None, contiguous=True),)))),), arg=None) for amt in [16,32]: opts = [Opt(op=OptOps.GROUPTOP, axis=0, amt=amt), Opt(op=OptOps.UNROLL, axis=0, amt=0)] - helper_test_lin(Kernel(ast), opts=opts, failed_platforms=["METAL", "GPU"]) + helper_test_lin(Kernel(ast), opts=opts, failed_platforms=[]) # END METAL=1 ./examples/beautiful_mnist.py failures if __name__ == '__main__': diff --git a/test/test_ops.py b/test/test_ops.py index 782f2b093d..43e2d3795e 100644 --- a/test/test_ops.py +++ b/test/test_ops.py @@ -192,6 +192,9 @@ class TestOps(unittest.TestCase): def test_arange_big(self): helper_test_op([], lambda: torch.arange(256, dtype=torch.int32), lambda: Tensor.arange(256), forward_only=True) + def test_arange_4096(self): + helper_test_op([], lambda: torch.arange(4096, dtype=torch.int32), lambda: Tensor.arange(4096), forward_only=True) + def test_sum_fake(self): helper_test_op([(256, 1)], lambda x: x.sum(axis=1)) diff --git a/test/test_pattern_matcher.py b/test/test_pattern_matcher.py index 4b00c5f85e..d82193c455 100644 --- a/test/test_pattern_matcher.py +++ b/test/test_pattern_matcher.py @@ -1,7 +1,7 @@ import unittest, itertools from test.helpers import TestUOps from tinygrad.dtype import dtypes -from tinygrad.ops import BinaryOps, TernaryOps, UnaryOps # noqa: F401 +from tinygrad.ops import BinaryOps, TernaryOps, ReduceOps, UnaryOps # noqa: F401 from tinygrad.codegen.uops import UOps, UOp, PatternMatcher, UPat, _match from tinygrad.codegen.uopgraph import UOpGraph, constant_folder diff --git a/tinygrad/codegen/uopgraph.py b/tinygrad/codegen/uopgraph.py index 8e92b1baed..b0292cee29 100644 --- a/tinygrad/codegen/uopgraph.py +++ b/tinygrad/codegen/uopgraph.py @@ -117,14 +117,20 @@ def sum_collapse(phi_input, loop, val1, val2): return UOp(UOps.PHI, phi_input.dtype, (phi_input, v2))+ret return None -def loop_collapse(loop_start, loop_end, compval, idx, mval, multconst, rng): - if getenv("DISABLE_LOOP_COLLAPSE") or not rng.arg[1]: return None # must be a REDUCE +def loop_collapse(loop_start, loop_end, compval, idx, mval, multconst, rng, reduce_allow_any_len): + if getenv("DISABLE_LOOP_COLLAPSE") or rng not in reduce_allow_any_len.src: return None # must be the right REDUCE if mval.arg >= 0 or loop_start.arg != 0: # TODO: support and test this with other mvals and loop_starts if DEBUG >= 1: print(f"WARNING, NOT FOLDING: mval:{mval.arg} loop_start:{loop_start.arg}") return None comprange = UOp.min(loop_end, UOp.max(UOp.alu(BinaryOps.IDIV, idx-compval-mval, mval) + (loop_end-loop_start), loop_start)) - return UOp(UOps.UNMUL, multconst.dtype, (comprange.cast(multconst.dtype) * multconst, loop_end-loop_start)) + return UOp(UOps.REDUCE, reduce_allow_any_len.dtype, (comprange.cast(multconst.dtype) * multconst,) + + tuple(x for x in reduce_allow_any_len.src[1:] if x is not rng), reduce_allow_any_len.arg) + +def index_collapse(idx,rng,buf,add,mul,ld,reduce_allow_any_len): + if rng not in reduce_allow_any_len.src: return None + return UOp(reduce_allow_any_len.op, reduce_allow_any_len.dtype, (UOp(ld.op, ld.dtype, (buf, add+mul*idx)),)+ + tuple(x for x in reduce_allow_any_len.src[1:] if x is not rng), reduce_allow_any_len.arg) # this is symbolic 2.0 constant_folder = PatternMatcher([ @@ -154,29 +160,24 @@ constant_folder = PatternMatcher([ lambda add, wmma: UOp(wmma.op, wmma.dtype, (wmma.src[0], wmma.src[1], wmma.src[2]+add), wmma.arg)), # threefry (UOp(UOps.ALU, dtype=dtypes.uint64, src=(UOp.var("x"), UOp.var("seed")), arg=BinaryOps.THREEFRY), threefry2x32), - # arange loop folding (early) - ((UOp.var("idx") + UOp.cvar("mval") * UOp(UOps.RANGE, src=(UOp.var("loop_start"), UOp.var("loop_end"))).name("rng")).lt(UOp.cvar("compval")).where( - UOp.cvar("multconst"), UOp.const(None, 0)), loop_collapse), - ((UOp.var("idx") - UOp(UOps.RANGE, src=(UOp.var("loop_start"), UOp.var("loop_end"))).name("rng")).lt(UOp.cvar("compval")).where( - UOp.cvar("multconst"), UOp.const(None, 0)), lambda **kwargs: loop_collapse(mval=UOp.const(dtypes.int, -1), **kwargs)), # sum collapse to mul (with possible GEP) (UPat(UOps.PHI, src=(UPat(UOps.DEFINE_ACC, name="phi_input", src=[UPat(UOps.CONST), UPat(UOps.RANGE, name="loop")]), UPat(UOps.ALU, BinaryOps.ADD, src=(UPat(name="val1"), UPat(name="val2"))))), sum_collapse), (UPat(UOps.PHI, src=(UPat(UOps.GEP, name="phi_input", src=(UPat(UOps.DEFINE_ACC, src=[UPat(UOps.CONST), UPat(UOps.RANGE, name="loop")]),)), UPat(UOps.ALU, BinaryOps.ADD, src=(UPat(name="val1"), UPat(name="val2"))))), sum_collapse), - # deal with UNMUL - (UOp.cvar('c1') * UOp(UOps.UNMUL, src=(UOp.cvar('c2'), UOp.var('v'))), lambda c1,c2,v: v if c1.arg == c2.arg else None), - (UOp.cvar('c1') * (UOp.var('add') + UOp(UOps.UNMUL, src=(UOp.cvar('c2'), UOp.var('v')))), - lambda c1, add, c2, v: (add*c1+v) if c1.arg == c2.arg else None), - (UOp(UOps.UNMUL, src=(UOp.const(None, 0).name('zero'), UOp.var())), lambda zero: zero), - (UOp(UOps.UNMUL).name('unmul').cast().name('root'), lambda root,unmul: UOp(UOps.UNMUL, root.dtype, (unmul.src[0].cast(root.dtype), unmul.src[1]))), + # arange loop folding (reduce) + (UOp(UOps.REDUCE, src=((UOp.var("idx") + UOp.cvar("mval") * UOp(UOps.RANGE, src=(UOp.var("loop_start"), UOp.var("loop_end"))).name("rng")) + .lt(UOp.cvar("compval")).where(UOp.cvar("multconst"), UOp.const(None, 0)),), arg=ReduceOps.SUM).name("reduce_allow_any_len"), loop_collapse), + (UOp(UOps.REDUCE, src=((UOp.var("idx") - UOp(UOps.RANGE, src=(UOp.var("loop_start"), UOp.var("loop_end"))).name("rng")) + .lt(UOp.cvar("compval")).where(UOp.cvar("multconst"), UOp.const(None, 0)),), arg=ReduceOps.SUM).name("reduce_allow_any_len"), + lambda **kwargs: loop_collapse(mval=UOp.const(dtypes.int, -1), **kwargs)), # indexing (with a multiply offset)! - (UOp.var('idx').eq(UOp(UOps.RANGE).name("rng")).cast()* - UOp(UOps.LOAD, src=(UOp.var("buf"), UOp.var('add')+UOp.var('mul')*UOp(UOps.RANGE).name("rng"))).name("ld"), - lambda idx,rng,buf,add,mul,ld: UOp(UOps.UNMUL, ld.dtype, (UOp(ld.op, ld.dtype, (buf, add+mul*idx)), rng.src[1]-rng.src[0]))), - (UOp.var('idx').eq(UOp(UOps.RANGE).name("rng")).where( - UOp(UOps.LOAD, src=(UOp.var("buf"), UOp.var('add')+UOp.var('mul')*UOp(UOps.RANGE).name("rng"))).name("ld"), UOp.const(None, 0.0)), - lambda idx,rng,buf,add,mul,ld: UOp(UOps.UNMUL, ld.dtype, (UOp(ld.op, ld.dtype, (buf, add+mul*idx)), rng.src[1]-rng.src[0]))), + (UOp(UOps.REDUCE, src=(UOp.var('idx').eq(UOp(UOps.RANGE).name("rng")).cast()* + UOp(UOps.LOAD, src=(UOp.var("buf"), UOp.var('add')+UOp.var('mul')*UOp(UOps.RANGE).name("rng"))).name("ld"),), + arg=ReduceOps.SUM).name("reduce_allow_any_len"), index_collapse), + (UOp(UOps.REDUCE, src=(UOp.var('idx').eq(UOp(UOps.RANGE).name("rng")).where( + UOp(UOps.LOAD, src=(UOp.var("buf"), UOp.var('add')+UOp.var('mul')*UOp(UOps.RANGE).name("rng"))).name("ld"), UOp.const(None, 0.0)),), + arg=ReduceOps.SUM).name("reduce_allow_any_len"), index_collapse), # other arange folders (UOp.cvar("c1") - (UOp.var("x") + UOp.cvar("c2")), lambda c1, c2, x: (c1-c2)-x), # c1 - (x + c2) -> (c1-c2) - x # max on special can go away (TODO: special should be variable, same thing applies) @@ -539,7 +540,7 @@ class UOpGraph: for u, x in scope_end.items(): self._uops.insert(self._uops.index(x)+1, UOp(END_FOR_UOP[u.op][1], None, (u,))) # sanity checks (NOTE: these can cause things to be skipped in BEAM) - bad_ops = dedup([x.op for x in self._uops if x.op in {UOps.EXPAND, UOps.CONTRACT, UOps.REDUCE, UOps.UNMUL}]) + bad_ops = dedup([x.op for x in self._uops if x.op in {UOps.EXPAND, UOps.CONTRACT, UOps.REDUCE}]) try: type_verify(self.uops) assert self._uops[-1].op is UOps.SINK, f"didn't end with SINK, ended with {self._uops[-1]}" diff --git a/tinygrad/codegen/uops.py b/tinygrad/codegen/uops.py index dcbc2f7eb8..a670776f1c 100644 --- a/tinygrad/codegen/uops.py +++ b/tinygrad/codegen/uops.py @@ -15,7 +15,7 @@ class UOps(Enum): SINK = auto(); VAR = auto(); EXPAND = auto(); CONTRACT = auto() # noqa: E702 DEFINE_GLOBAL = auto(); DEFINE_VAR = auto(); DEFINE_LOCAL = auto(); DEFINE_ACC = auto() # noqa: E702 CONST = auto(); SPECIAL = auto() # noqa: E702 - NOOP = auto(); UNMUL = auto(); GEP = auto() # noqa: E702 + NOOP = auto(); GEP = auto() # noqa: E702 # math ops CAST = auto(); BITCAST = auto(); VECTORIZE = auto() # noqa: E702 ALU = auto(); REDUCE = auto(); WMMA = auto() # noqa: E702 @@ -35,7 +35,7 @@ class UOp: src: Tuple[UOp, ...] = tuple() arg: Any = None def commutative(self) -> bool: - return self.op is UOps.UNMUL or (self.op is UOps.ALU and \ + return (self.op is UOps.ALU and \ self.arg in {BinaryOps.ADD, BinaryOps.MUL, BinaryOps.MAX, BinaryOps.CMPNE, BinaryOps.XOR, BinaryOps.AND, BinaryOps.OR}) @functools.cached_property def cmp_tuple(self):