diff --git a/extra/gemm/tinygrad_nv_matmul.py b/extra/gemm/tinygrad_nv_matmul.py index 77538f7e74..afb9d99408 100644 --- a/extra/gemm/tinygrad_nv_matmul.py +++ b/extra/gemm/tinygrad_nv_matmul.py @@ -14,7 +14,7 @@ if __name__ == "__main__": if getenv("GEMV"): opts = [ Opt(op=OptOps.UPCAST, axis=1, amt=8), - Opt(op=OptOps.GROUP, axis=0, amt=32), + Opt(op=OptOps.LOCAL, axis=1, amt=32), ] else: opts = [ diff --git a/test/backend/test_linearizer.py b/test/backend/test_linearizer.py index 64180d6cc5..99ac143df7 100644 --- a/test/backend/test_linearizer.py +++ b/test/backend/test_linearizer.py @@ -173,7 +173,7 @@ class TestLinearizer(unittest.TestCase): def test_upcast_with_locals(self): x, y = Tensor.rand(1,128), Tensor.rand(128, 128) r = (x@y).relu() - opts_to_apply = [Opt(op=OptOps.GROUP, axis=0, arg=8), Opt(op=OptOps.LOCAL, axis=0, arg=4), Opt(op=OptOps.UPCAST, axis=0, arg=4)] + opts_to_apply = [Opt(op=OptOps.LOCAL, axis=1, arg=8), Opt(op=OptOps.LOCAL, axis=0, arg=4), Opt(op=OptOps.UPCAST, axis=0, arg=4)] program = to_program(replace_opts(r.schedule_linear().src[-1].src[0], opts_to_apply), renderer=Device[Device.DEFAULT].renderer) stores = [u for u in tuple(program.src[1].src) if u.op is Ops.STORE and u.src[0].addrspace != AddrSpace.REG] @@ -345,7 +345,7 @@ class TestLinearizer(unittest.TestCase): def test_grouped_store_locals_and_globals(self): x, y = Tensor.empty(64, 64), Tensor.empty(64, 64) out = x@y - opt = [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.GROUPTOP, 0, 8), + opt = [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.GROUPTOP, 3, 8), Opt(OptOps.UPCAST, 3, 4), Opt(OptOps.UPCAST, 0, 4), Opt(OptOps.UPCAST, 1, 2)] # upcast accs in both reduces ast = helper_linearizer_opt(out, opts=[opt]) def get_recursive(uop): return set.union(set(uop.src), [uop], *[get_recursive(v) for v in uop.src]) @@ -386,7 +386,7 @@ class TestLinearizer(unittest.TestCase): def test_two_grouped_stores_local(self): # GROUP on both reduces puts two LOCAL buffers in one kernel, and the store to each needs its own barrier a = Tensor.rand(32, 32).realize() - opts = [Opt(OptOps.GROUP, 1, 4), Opt(OptOps.GROUP, 2, 4)] + opts = [Opt(OptOps.LOCAL, 3, 4), Opt(OptOps.LOCAL, 5, 4)] ast = helper_linearizer_opt(single_kernel_softmax(a), [opts]) uops = to_program(replace_opts(ast, opts), renderer=Device[Device.DEFAULT].renderer).src[1].src self.assertEqual(len([u for u in uops if u.op is Ops.BARRIER]), 2) diff --git a/test/backend/test_linearizer_dumb.py b/test/backend/test_linearizer_dumb.py index 583c3c64b6..63908f24af 100644 --- a/test/backend/test_linearizer_dumb.py +++ b/test/backend/test_linearizer_dumb.py @@ -24,7 +24,7 @@ class TestLinearizerFailure(unittest.TestCase): c10 = c9.index((((c3*UOp.const(4704000))+c2)+(c6*UOp.const(784))).valid(UOp.const(True))) c11 = c5.alu(Ops.CMPNE, ((((c3*UOp.const(6000))+c6)+((c7*UOp.const(16))+c8)).alu(Ops.CMPLT, UOp.const(59999)).where(UOp.const(0).cast(dtypes.int), UOp.const(1).cast(dtypes.int)).reduce(c7, c8, arg=Ops.ADD)+UOp.const(-1).cast(dtypes.int))).where(UOp.const(0).cast(dtypes.uchar), c10).reduce(c6, arg=Ops.ADD) c12 = c0.index((((c1*UOp.const(7840))+(c2*UOp.const(10)))+c3).valid(UOp.const(True))).store(c11).end(c1, c2, c3) - ast = c12.sink(arg=KernelInfo(name='test', applied_opts=(Opt(op=OptOps.GROUP, axis=1, arg=16),), opts_to_apply=None)) + ast = c12.sink(arg=KernelInfo(name='test', applied_opts=(Opt(op=OptOps.LOCAL, axis=4, arg=16),), opts_to_apply=None)) _ = to_program(ast, Device["METAL"].renderer) if __name__ == '__main__': diff --git a/test/null/test_uops_stats.py b/test/null/test_uops_stats.py index ff90569350..315f2a657a 100644 --- a/test/null/test_uops_stats.py +++ b/test/null/test_uops_stats.py @@ -230,7 +230,7 @@ class TestStatsOptimized(unittest.TestCase): def test_gemm_group(self): try: - p = to_program(replace_opts(self.ast_gemm, [Opt(OptOps.GROUP, 0, 4)]), renderer=Device[Device.DEFAULT].renderer) + p = to_program(replace_opts(self.ast_gemm, [Opt(OptOps.LOCAL, 2, 4)]), renderer=Device[Device.DEFAULT].renderer) except KernelOptError: raise unittest.SkipTest("no locals") SZ = N*N*4 diff --git a/test/opt/test_kernel_opts.py b/test/opt/test_kernel_opts.py index 864f8c288d..2aed363f26 100644 --- a/test/opt/test_kernel_opts.py +++ b/test/opt/test_kernel_opts.py @@ -18,20 +18,21 @@ class TestKernelOpts(unittest.TestCase): [Opt(OptOps.LOCAL, 0, 2)], [Opt(OptOps.LOCAL, 0, 8)], [Opt(OptOps.LOCAL, 0, 16)], # Checking how it works with locals - [Opt(OptOps.GROUPTOP, 0, 2)], - [Opt(OptOps.GROUPTOP, 0, 32)], - [Opt(OptOps.GROUPTOP, 0, 64)], # Checking how it works with grouped reduce - [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.GROUPTOP, 0, 2)], - [Opt(OptOps.LOCAL, 0, 16), Opt(OptOps.GROUPTOP, 0, 16)], - [Opt(OptOps.LOCAL, 0, 32), Opt(OptOps.GROUPTOP, 0, 2)], + [Opt(OptOps.GROUPTOP, 1, 2)], + [Opt(OptOps.GROUPTOP, 1, 32)], + [Opt(OptOps.GROUPTOP, 1, 64)], # Checking how it works with grouped reduce + [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.GROUPTOP, 2, 2)], + [Opt(OptOps.LOCAL, 0, 16), Opt(OptOps.GROUPTOP, 2, 16)], + [Opt(OptOps.LOCAL, 0, 32), Opt(OptOps.GROUPTOP, 2, 2)], # Checking how it works with locals + grouped reduce - [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.GROUPTOP, 0, 64)], + [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.GROUPTOP, 2, 64)], # Checking how it works with locals + grouped reduce + upcasts - [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.GROUPTOP, 0, 2), Opt(OptOps.UPCAST, 0, 8), Opt(OptOps.UPCAST, 4, 4)], + [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.GROUPTOP, 2, 2), Opt(OptOps.UPCAST, 0, 8), Opt(OptOps.UPCAST, 4, 4)], # many local + many group - [Opt(OptOps.GROUP, 0, 2)] * 4, + [Opt(OptOps.LOCAL, 1, 2), Opt(OptOps.LOCAL, 2, 2), Opt(OptOps.LOCAL, 3, 2), Opt(OptOps.LOCAL, 4, 2)], [Opt(OptOps.LOCAL, 0, 2)] * 4, - [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.GROUP, 0, 2)] * 4, + [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.LOCAL, 2, 2), Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.LOCAL, 4, 2), + Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.LOCAL, 6, 2), Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.LOCAL, 8, 2)], ]) def test_upcasts(self): @@ -71,17 +72,17 @@ class TestKernelOpts(unittest.TestCase): [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 1, 4)], [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 1, 32)], [Opt(OptOps.LOCAL, 0, 16), Opt(OptOps.LOCAL, 1, 8)], # Checking how it works with locals - [Opt(OptOps.GROUPTOP, 0, 2)], - [Opt(OptOps.GROUPTOP, 0, 32)], - [Opt(OptOps.GROUPTOP, 0, 32), Opt(OptOps.UPCAST, 2, 4)], # Checking how it works with grouped_reduce - [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.LOCAL, 1, 2), Opt(OptOps.GROUPTOP, 0, 32)], - [Opt(OptOps.LOCAL, 0, 8), Opt(OptOps.GROUPTOP, 0, 32)], - [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 0, 8), Opt(OptOps.GROUPTOP, 0, 4)], # Checking how it works with local+grouped_reduce + [Opt(OptOps.GROUPTOP, 2, 2)], + [Opt(OptOps.GROUPTOP, 2, 32)], + [Opt(OptOps.GROUPTOP, 2, 32), Opt(OptOps.UPCAST, 2, 4)], # Checking how it works with grouped_reduce + [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.LOCAL, 1, 2), Opt(OptOps.GROUPTOP, 4, 32)], + [Opt(OptOps.LOCAL, 0, 8), Opt(OptOps.GROUPTOP, 3, 32)], + [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 0, 8), Opt(OptOps.GROUPTOP, 4, 4)], # Checking how it works with local+grouped_reduce # Checking all together - [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.GROUPTOP, 0, 8), Opt(OptOps.UPCAST, 4, 4), Opt(OptOps.UPCAST, 0, 4), + [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.GROUPTOP, 4, 8), Opt(OptOps.UPCAST, 4, 4), Opt(OptOps.UPCAST, 0, 4), Opt(OptOps.UPCAST, 1, 2)], # Full global upcast + local - [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.GROUPTOP, 0, 8), Opt(OptOps.UPCAST, 4, 4), Opt(OptOps.UPCAST, 0, 8)], + [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.GROUPTOP, 4, 8), Opt(OptOps.UPCAST, 4, 4), Opt(OptOps.UPCAST, 0, 8)], ]) @unittest.skipUnless(Device[Device.DEFAULT].renderer.has_local, "test requires locals") @@ -93,20 +94,20 @@ class TestKernelOpts(unittest.TestCase): r = a.sum(axis=(1,3)) helper_linearizer_opt(r, [ # openCL / DEV=CL is 256 max threads - [Opt(OptOps.GROUPTOP, 0, 2)], [Opt(OptOps.GROUPTOP, 0, 32)], - [Opt(OptOps.GROUPTOP, 1, 2)], [Opt(OptOps.GROUPTOP, 1, 32)], # Checking how it works with 1 grouped_reduce. - [Opt(OptOps.GROUPTOP, 0, 2), Opt(OptOps.GROUPTOP, 1, 2)], - [Opt(OptOps.GROUPTOP, 0, 16), Opt(OptOps.GROUPTOP, 1, 2)], - [Opt(OptOps.GROUPTOP, 0, 4), Opt(OptOps.GROUPTOP, 1, 64)], # Checking how it works with 2 grouped_reduces. - [Opt(OptOps.GROUPTOP, 0, 16), Opt(OptOps.GROUPTOP, 1, 2), Opt(OptOps.UPCAST, 2, 4)], - [Opt(OptOps.GROUPTOP, 0, 2), Opt(OptOps.GROUPTOP, 1, 32), Opt(OptOps.UPCAST, 4, 4)], # Checking how it works with 2 grouped_reduces + upcasts. - [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 1, 4), Opt(OptOps.GROUPTOP, 0, 4), Opt(OptOps.GROUPTOP, 1, 4)], + [Opt(OptOps.GROUPTOP, 2, 2)], [Opt(OptOps.GROUPTOP, 2, 32)], + [Opt(OptOps.GROUPTOP, 3, 2)], [Opt(OptOps.GROUPTOP, 3, 32)], # Checking how it works with 1 grouped_reduce. + [Opt(OptOps.GROUPTOP, 2, 2), Opt(OptOps.GROUPTOP, 4, 2)], + [Opt(OptOps.GROUPTOP, 2, 16), Opt(OptOps.GROUPTOP, 4, 2)], + [Opt(OptOps.GROUPTOP, 2, 4), Opt(OptOps.GROUPTOP, 4, 64)], # Checking how it works with 2 grouped_reduces. + [Opt(OptOps.GROUPTOP, 2, 16), Opt(OptOps.GROUPTOP, 4, 2), Opt(OptOps.UPCAST, 2, 4)], + [Opt(OptOps.GROUPTOP, 2, 2), Opt(OptOps.GROUPTOP, 4, 32), Opt(OptOps.UPCAST, 4, 4)], # Checking how it works with 2 grouped_reduces + upcasts. + [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 1, 4), Opt(OptOps.GROUPTOP, 4, 4), Opt(OptOps.GROUPTOP, 6, 4)], # Checking how it works with 2 grouped_reduces + upcasts + locals. - [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 1, 4), Opt(OptOps.GROUPTOP, 0, 2), Opt(OptOps.GROUPTOP, 1, 32), Opt(OptOps.UPCAST, 5, 4)], - [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.LOCAL, 1, 2), Opt(OptOps.GROUPTOP, 0, 8), Opt(OptOps.GROUPTOP, 1, 4), Opt(OptOps.UPCAST, 0, 2)], - [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.LOCAL, 1, 2), Opt(OptOps.GROUPTOP, 0, 8), Opt(OptOps.GROUPTOP, 1, 4), Opt(OptOps.UPCAST, 0, 2), + [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 1, 4), Opt(OptOps.GROUPTOP, 4, 2), Opt(OptOps.GROUPTOP, 6, 32), Opt(OptOps.UPCAST, 5, 4)], + [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.LOCAL, 1, 2), Opt(OptOps.GROUPTOP, 4, 8), Opt(OptOps.GROUPTOP, 6, 4), Opt(OptOps.UPCAST, 0, 2)], + [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.LOCAL, 1, 2), Opt(OptOps.GROUPTOP, 4, 8), Opt(OptOps.GROUPTOP, 6, 4), Opt(OptOps.UPCAST, 0, 2), Opt(OptOps.UPCAST, 4, 4), Opt(OptOps.UPCAST, 5, 4)], # Checking how it works with 2 grouped_reduces + upcasts + locals. - [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 1, 4), Opt(OptOps.GROUPTOP, 0, 4), Opt(OptOps.GROUPTOP, 1, 4), Opt(OptOps.UPCAST, 0, 2), + [Opt(OptOps.LOCAL, 0, 4), Opt(OptOps.LOCAL, 1, 4), Opt(OptOps.GROUPTOP, 4, 4), Opt(OptOps.GROUPTOP, 6, 4), Opt(OptOps.UPCAST, 0, 2), Opt(OptOps.UPCAST, 0, 2)], # No globals ]) @@ -221,7 +222,7 @@ class TestKernelOpts(unittest.TestCase): def test_padto_group_full_unroll_sum(self): a = Tensor.ones(2, 28, 4096, dtype=dtypes.bfloat16).realize() out = ((a * 0.5).float().square()).sum(axis=(0, 2)) - opts_to_apply = [Opt(OptOps.GROUPTOP, 1, 256), Opt(OptOps.PADTO, 3, 32), Opt(OptOps.UPCAST, 3, 0), Opt(OptOps.UPCAST, 0, 7)] + opts_to_apply = [Opt(OptOps.GROUPTOP, 2, 256), Opt(OptOps.PADTO, 3, 32), Opt(OptOps.UPCAST, 3, 0), Opt(OptOps.UPCAST, 0, 7)] helper_linearizer_opt(out, [opts_to_apply], check_default_opt=False) def test_padto_sum(self): @@ -282,15 +283,15 @@ class TestKernelOpts(unittest.TestCase): r = a@b opts_shapes = [ ([Opt(OptOps.LOCAL, 0, 2)], [("blue",16),("blue",32),("cyan",2),("red",32)]), - ([Opt(OptOps.LOCAL, 0, 2),Opt(OptOps.GROUP, 0, 2)], [("blue",16),("blue",32),("cyan",2),("green",2),("red",16)]), + ([Opt(OptOps.LOCAL, 0, 2),Opt(OptOps.LOCAL, 3, 2)], [("blue",16),("blue",32),("cyan",2),("green",2),("red",16)]), # check to ensure local_dims are stable for full UNROLL of the first reduce ([Opt(OptOps.LOCAL, 0, 2),Opt(OptOps.UPCAST, 3, 0)], [("blue",16),("blue",32),("cyan",2),("magenta",32)]), ([Opt(OptOps.UPCAST, 2, 0),Opt(OptOps.LOCAL, 0, 2)], [("blue",16),("blue",32),("cyan",2),("magenta",32)]), # check behavior for full UNROLL on an existing GROUP - ([Opt(OptOps.LOCAL, 0, 2),Opt(OptOps.GROUP, 0, 0),Opt(OptOps.UPCAST, 3, 2)], [("blue",16),("blue",32),("cyan",2),("green",16),("magenta",2)]), - ([Opt(OptOps.LOCAL, 0, 2),Opt(OptOps.GROUP, 0, 0),Opt(OptOps.UPCAST, 3, 0)], [("blue",16),("blue",32),("cyan",2),("magenta",32)]), - ([Opt(OptOps.GROUP, 0, 0),Opt(OptOps.LOCAL, 0, 2),Opt(OptOps.UPCAST, 2, 0)], [("blue",16),("blue",32),("cyan",2),("magenta",32)]), - ([Opt(OptOps.GROUP, 0, 2),Opt(OptOps.UPCAST, 2, 0)], [("blue",32),("blue",32),("red",16),("magenta",2)]), + ([Opt(OptOps.LOCAL, 0, 2),Opt(OptOps.LOCAL, 3, 0),Opt(OptOps.UPCAST, 3, 2)], [("blue",16),("blue",32),("cyan",2),("green",16),("magenta",2)]), + ([Opt(OptOps.LOCAL, 0, 2),Opt(OptOps.LOCAL, 3, 0),Opt(OptOps.UPCAST, 3, 0)], [("blue",16),("blue",32),("cyan",2),("magenta",32)]), + ([Opt(OptOps.LOCAL, 2, 0),Opt(OptOps.LOCAL, 0, 2),Opt(OptOps.UPCAST, 2, 0)], [("blue",16),("blue",32),("cyan",2),("magenta",32)]), + ([Opt(OptOps.LOCAL, 2, 2),Opt(OptOps.UPCAST, 2, 0)], [("blue",32),("blue",32),("red",16),("magenta",2)]), ] helper_linearizer_opt(r, [x[0] for x in opts_shapes], color_sizes=[x[1] for x in opts_shapes]) @@ -315,7 +316,7 @@ class TestKernelOpts(unittest.TestCase): helper_linearizer_opt(r, [[Opt(OptOps.UPCAST, 1, 4), Opt(OptOps.GROUPTOP, 0, 16)],]) r = a.sum((1, 2)).sum() with self.assertRaises(KernelOptError): - helper_linearizer_opt(r, [[Opt(OptOps.GROUPTOP, 1, 4), Opt(OptOps.GROUPTOP, 0, 16)],]) + helper_linearizer_opt(r, [[Opt(OptOps.GROUPTOP, 1, 4), Opt(OptOps.GROUPTOP, 1, 16)],]) if __name__ == '__main__': unittest.main() diff --git a/tinygrad/codegen/opt/__init__.py b/tinygrad/codegen/opt/__init__.py index da1f3083f6..61737f2184 100644 --- a/tinygrad/codegen/opt/__init__.py +++ b/tinygrad/codegen/opt/__init__.py @@ -4,7 +4,7 @@ from enum import Enum, auto from dataclasses import dataclass class OptOps(Enum): - TC = auto(); UPCAST = auto(); LOCAL = auto(); GROUP = auto(); GROUPTOP = auto(); PADTO = auto(); SWAP = auto() # noqa: E702 + TC = auto(); UPCAST = auto(); LOCAL = auto(); GROUPTOP = auto(); PADTO = auto(); SWAP = auto() # noqa: E702 def __lt__(self, x:OptOps): return self.value < x.value @dataclass(frozen=True, order=True) diff --git a/tinygrad/codegen/opt/heuristic.py b/tinygrad/codegen/opt/heuristic.py index 63a3b2c063..7c7fb65494 100644 --- a/tinygrad/codegen/opt/heuristic.py +++ b/tinygrad/codegen/opt/heuristic.py @@ -68,7 +68,7 @@ def hand_coded_optimizations(k:Scheduler) -> Scheduler: if DEBUG >= 3: print(f"MATVEC: {k.full_shape=} {first_reduce_rng.render()} {MV_BLOCKSIZE=} {MV_THREADS_PER_ROW=} {MV_ROWS_PER_THREAD=}") try: - if MV_THREADS_PER_ROW > 1: k.apply_opt(Opt(OptOps.GROUP, 0, MV_THREADS_PER_ROW)) + if MV_THREADS_PER_ROW > 1: k.apply_opt(Opt(OptOps.LOCAL, k.axes_of(AxisType.REDUCE)[0], MV_THREADS_PER_ROW)) except KernelOptError: pass if MV_BLOCKSIZE > 1: k.apply_opt(Opt(OptOps.LOCAL, global_idx, MV_BLOCKSIZE)) if MV_ROWS_PER_THREAD > 1: k.apply_opt(Opt(OptOps.UPCAST, global_idx, MV_ROWS_PER_THREAD)) @@ -76,7 +76,7 @@ def hand_coded_optimizations(k:Scheduler) -> Scheduler: # are we grouping? (requires local shape support) if resolve(prod(k.output_shape[i] for i in k.upcastable_dims) <= (240 if k.ren.target.device == "QCOM" else 2048), False): - for axis, sz in itertools.product((0, 1, 2), (16,)): + for axis, sz in itertools.product(k.axes_of(AxisType.REDUCE)[:3], (16,)): try: k.apply_opt(Opt(OptOps.GROUPTOP, axis, sz)) break diff --git a/tinygrad/codegen/opt/postrange.py b/tinygrad/codegen/opt/postrange.py index 085324bedd..5ee2ed5501 100644 --- a/tinygrad/codegen/opt/postrange.py +++ b/tinygrad/codegen/opt/postrange.py @@ -13,6 +13,7 @@ from tinygrad.renderer import Renderer upcast_to = {AxisType.GLOBAL: AxisType.UPCAST, AxisType.LOCAL: AxisType.UPCAST, AxisType.WEAK: AxisType.UPCAST, AxisType.GROUP_REDUCE: AxisType.UNROLL, AxisType.REDUCE: AxisType.UNROLL} +local_to = {AxisType.GLOBAL: AxisType.LOCAL, AxisType.WEAK: AxisType.LOCAL, AxisType.REDUCE: AxisType.GROUP_REDUCE} class Scheduler: def __init__(self, ast:UOp, ren:Renderer): @@ -110,48 +111,44 @@ class Scheduler: if isinstance(s:=self.full_shape[i], int) and s > 1] def real_axis(self, op:OptOps, axis:int|None) -> int: - try: - if axis is None or op is OptOps.TC: return -1 - if op in {OptOps.GROUP, OptOps.GROUPTOP}: return self.axes_of(AxisType.REDUCE)[axis] - check(axis < self.shape_len, f"invalid axis on {axis=} {op=} {self.shape_len=}") - return axis - except IndexError as e: raise KernelOptError from e + if axis is None or op is OptOps.TC: return -1 + check(0 <= axis < self.shape_len, f"invalid axis on {axis=} {op=} {self.shape_len=}") + return axis def apply_opt(self, opt:Opt, append_opt:bool=True): - if opt.op in {OptOps.LOCAL, OptOps.GROUP, OptOps.GROUPTOP}: + if opt.op in {OptOps.LOCAL, OptOps.GROUPTOP}: check(self.ren.has_local, "locals needed for opt") rng = self.rngs[real_axis] if (real_axis:=self.real_axis(opt.op, opt.axis)) >= 0 else UOp(Ops.NOOP) - opt_to_at = { - OptOps.LOCAL: AxisType.LOCAL, OptOps.UPCAST: AxisType.UPCAST, OptOps.GROUP: AxisType.GROUP_REDUCE, - OptOps.GROUPTOP: AxisType.GROUP_REDUCE} + opt_to_at = {OptOps.LOCAL: AxisType.LOCAL, OptOps.UPCAST: AxisType.UPCAST, OptOps.GROUPTOP: AxisType.GROUP_REDUCE} ret = None if opt.op in opt_to_at: amt:int = int(rng.vmax+1) if opt.arg == 0 else cast(int, opt.arg) new_type = opt_to_at[opt.op] - # copied from kernel.py. prevents METAL compiler hangs - if self.reduceop is not None and (opt.op in {OptOps.GROUP, OptOps.GROUPTOP} or (self.group_for_reduces and opt.op != OptOps.PADTO)): - upcast_local_sz = prod([self.full_shape[a] for a in self.axes_of(AxisType.UPCAST, AxisType.WARP, AxisType.LOCAL, AxisType.GROUP_REDUCE)]) - smem_sz = amt*upcast_local_sz*self.reduceop.dtype.itemsize - check(smem_sz <= self.ren.shared_max, f"exceeds maximum shared memory size: needs {smem_sz}, max {self.ren.shared_max}") - if self.reduceop is not None and (opt.op in {OptOps.GROUP, OptOps.GROUPTOP}): - # We currently dont support a group within another rudece, TODO: fix if-contexts - reduce = [u for u in self.ast.backward_slice if u.op is Ops.REDUCE and rng in merge_dicts([r.ranges for r in u.src[1:]])][0] - check(not any(u.arg[-1] in (AxisType.REDUCE, AxisType.UNROLL, AxisType.GROUP_REDUCE) for u in reduce.ranges), - "cannot have a GROUP_REDUCE inside another reduce") - if opt.op is OptOps.UPCAST: check(rng.arg[-1] in upcast_to, f"upcast is for GLOBAL/LOCAL/LOOP/REDUCE, not {rng.arg[-1]}") if (new_type:=upcast_to[rng.arg[-1]]) is AxisType.UNROLL: check(amt <= 32, "don't unroll more than 32") else: check((self.ren is not None and self.ren.target.device == "DSP") or amt <= 16, "don't upcast more than 16") if opt.op is OptOps.LOCAL: - check(rng.arg[-1] in {AxisType.GLOBAL, AxisType.WEAK}, "local is for globals") - if opt.op in {OptOps.GROUP, OptOps.GROUPTOP}: + check(rng.arg[-1] in local_to, f"local is for GLOBAL/LOOP/REDUCE, not {rng.arg[-1]}") + new_type = local_to[rng.arg[-1]] + if opt.op is OptOps.GROUPTOP: check(rng.arg[-1] is AxisType.REDUCE, "grouptop is for reduce") + if new_type is AxisType.GROUP_REDUCE: check(all(x.op is not OptOps.TC for x in self.applied_opts), "no grouping with tensor cores") # TODO: why is this wrong? - check(rng.arg[-1] == AxisType.REDUCE, "group is for reduce") + + # copied from kernel.py. prevents METAL compiler hangs + if self.reduceop is not None and (new_type is AxisType.GROUP_REDUCE or (self.group_for_reduces and opt.op != OptOps.PADTO)): + upcast_local_sz = prod([self.full_shape[a] for a in self.axes_of(AxisType.UPCAST, AxisType.WARP, AxisType.LOCAL, AxisType.GROUP_REDUCE)]) + smem_sz = amt*upcast_local_sz*self.reduceop.dtype.itemsize + check(smem_sz <= self.ren.shared_max, f"exceeds maximum shared memory size: needs {smem_sz}, max {self.ren.shared_max}") + if self.reduceop is not None and new_type is AxisType.GROUP_REDUCE: + # We currently dont support a group within another rudece, TODO: fix if-contexts + reduce = [u for u in self.ast.backward_slice if u.op is Ops.REDUCE and rng in merge_dicts([r.ranges for r in u.src[1:]])][0] + check(not any(u.arg[-1] in (AxisType.REDUCE, AxisType.UNROLL, AxisType.GROUP_REDUCE) for u in reduce.ranges), + "cannot have a GROUP_REDUCE inside another reduce") ret = self.shift_to(rng, amt, new_type, top=opt.op is OptOps.GROUPTOP) elif opt.op is OptOps.TC: check(len(self.applied_opts) == 0, "tensor core opts must be first") # TODO: remove the need for this by having warps diff --git a/tinygrad/codegen/opt/search.py b/tinygrad/codegen/opt/search.py index 2cc29d54c6..4e505e5dfe 100644 --- a/tinygrad/codegen/opt/search.py +++ b/tinygrad/codegen/opt/search.py @@ -12,11 +12,10 @@ from tinygrad.codegen import to_program from tinygrad.codegen.opt.postrange import Scheduler actions = [Opt(op=OptOps.UPCAST, axis=axis, arg=amt) for amt in [0,2,3,4,5,7] for axis in range(10)] -actions += [Opt(op=OptOps.LOCAL, axis=axis, arg=amt) for amt in [2,3,4,8,13,16,29] for axis in range(6)] -actions += [Opt(op=OptOps.GROUPTOP, axis=axis, arg=amt) for amt in [13,16,28,29,32,49,64,256] for axis in range(3)] -actions += [Opt(op=OptOps.GROUP, axis=axis, arg=amt) for amt in [0,4,8,16] for axis in range(3)] +actions += [Opt(op=OptOps.LOCAL, axis=axis, arg=amt) for amt in [0,2,3,4,8,13,16,29] for axis in range(8)] +actions += [Opt(op=OptOps.GROUPTOP, axis=axis, arg=amt) for amt in [13,16,28,29,32,49,64,256] for axis in range(8)] if getenv("BEAM_PADTO", 0): actions += [Opt(op=OptOps.PADTO, axis=axis, arg=amt) for amt in [32] for axis in range(7)] -actions += [Opt(op=OptOps.LOCAL, axis=0, arg=32), Opt(op=OptOps.LOCAL, axis=6, arg=2)] +actions += [Opt(op=OptOps.LOCAL, axis=0, arg=32)] actions += [Opt(op=OptOps.TC, axis=0, arg=(-1, 0, getenv("TC", 1)))] # covers resnet kernels (3 global * 3 reduce) actions += [Opt(op=OptOps.TC, axis=axis, arg=(-1, getenv("TC_OPT", 2), getenv("TC", 1))) for axis in range(9)]