diff --git a/extra/gemm/amd_uop_matmul.py b/extra/gemm/amd_uop_matmul.py index d217edb8c5..9ad5351b57 100644 --- a/extra/gemm/amd_uop_matmul.py +++ b/extra/gemm/amd_uop_matmul.py @@ -5,7 +5,7 @@ from tinygrad.dtype import AddrSpace from tinygrad.helpers import getenv, colored, prod, unwrap from tinygrad.shape.shapetracker import ShapeTracker, View from tinygrad.shape.view import strides_for_shape -from tinygrad.codegen.opt.kernel import axis_colors +from tinygrad.codegen.opt.kernel import axis_colors, Opt, OptOps from tinygrad.codegen.opt.swizzler import merge_views, view_left def to_colored(full_shape, axis_types): return '_'.join([colored(str(s), axis_colors[at]) for s,at in zip(full_shape, axis_types)]) @@ -44,6 +44,21 @@ pm = PatternMatcher([ (UPat(Ops.VIEW, src=(UPat(Ops.REDUCE_AXIS, src=(UPat.var("src"),), name="r"),), name="view"), swizzle_reduceop), ]) +def rangeify_kernel3(): + a = Tensor.empty(N,N) + b = Tensor.empty(N,N) + c = a@b + #c = c.reshape((32,2,16,4,32,2,16,4)).contiguous() + with Context(RANGEIFY=1): + sink = c.schedule()[-1].ast + #print(sink) + + opts = [Opt(OptOps.UPCAST, 0, 4), Opt(OptOps.LOCAL, 0, 16), Opt(OptOps.UPCAST, 0, 2)] + opts += [Opt(OptOps.UPCAST, 1, 4), Opt(OptOps.LOCAL, 1, 16), Opt(OptOps.UPCAST, 1, 2)] + opts += [Opt(OptOps.UNROLL, 0, 8)] + + return sink.replace(arg=KernelInfo(opts_to_apply=tuple(opts))) + def top_spec_kernel3(): a = Tensor.empty(N,N) b = Tensor.empty(N,N) @@ -309,10 +324,15 @@ def hand_spec_kernel3(kernel4=getenv("K4", 0), kernel5=getenv("K5", 0)): if __name__ == "__main__": HL = getenv("HL") - if HL == 2: hprg = top_spec_kernel3() + if HL == 3: hprg = rangeify_kernel3() + elif HL == 2: hprg = top_spec_kernel3() elif HL == 1: hprg = hl_spec_kernel3() else: hprg = hand_spec_kernel3() - prg = get_program(hprg, Device.default.renderer) + if HL == 3: + with Context(RANGEIFY=1, BLOCK_REORDER=0): + prg = get_program(hprg, Device.default.renderer) + else: + prg = get_program(hprg, Device.default.renderer) print(prg.src) if getenv("SRC"): exit(0) hrunner = CompiledRunner(prg) diff --git a/tinygrad/schedule/rangeify.py b/tinygrad/schedule/rangeify.py index 53efb1fc67..aa3f10a4e0 100644 --- a/tinygrad/schedule/rangeify.py +++ b/tinygrad/schedule/rangeify.py @@ -1,12 +1,12 @@ from typing import Any from dataclasses import dataclass, field -from tinygrad.dtype import dtypes, PtrDType -from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, resolve, GroupOp, RewriteNotReady, _substitute, AxisType +from tinygrad.dtype import dtypes, PtrDType, AddrSpace +from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, resolve, GroupOp, RewriteNotReady, _substitute from tinygrad.helpers import argsort, prod, all_same, pluralize, getenv, colored, RANGEIFY from tinygrad.schedule.multi import multi_pm from tinygrad.schedule.kernelize import Kernel -from tinygrad.uop.ops import track_rewrites, graph_rewrite_map, graph_rewrite, KernelInfo, identity_element, sint +from tinygrad.uop.ops import track_rewrites, graph_rewrite_map, graph_rewrite, KernelInfo, identity_element, sint, AxisType # 0. do some cleanup rewrites, mostly copied from the old stuff @@ -332,13 +332,17 @@ pm_cleanups = double_reshape+pm_mops+PatternMatcher([ def bufferize_to_store(x:UOp): rngs = x.src[1:] shape = tuple([int(r.vmax+1) for r in rngs]) - sdtype = x.dtype.ptr(size=prod(shape)) + sdtype = x.dtype.ptr(size=prod(shape), addrspace=AddrSpace.GLOBAL if not isinstance(x.arg, AddrSpace) else x.arg) assert prod(shape) > 0, f"no zero sized buffers {shape}" if x.src[0].op is Ops.ASSIGN: assign_target, assign_src = x.src[0].src assert assign_target.op is Ops.INDEX return assign_target.replace(dtype=sdtype).store(assign_src, *rngs, dtype=sdtype) - buf = UOp.new_buffer(x.arg, prod(shape), x.dtype) + if sdtype.addrspace == AddrSpace.GLOBAL: + buf = UOp.new_buffer(x.arg, prod(shape), x.dtype) + else: + # TODO: how to dedup this + buf = UOp(Ops.DEFINE_LOCAL, sdtype, arg=UOp.unique().arg) return buf.reshape(shape).index(*rngs, dtype=sdtype).store(x.src[0], *rngs, dtype=sdtype).forced_reshape(shape, dtype=x.dtype) pm_add_buffers = pm_mops+PatternMatcher([ @@ -380,6 +384,12 @@ to_define_global = PatternMatcher([ (UPat(Ops.BIND, name="b"), unbind_kernel), (UPat((Ops.ASSIGN, Ops.MSTACK, Ops.MSELECT), name="assign"), handle_assign), + # HACK in case any CONSTs were replaced + # this is only needed if you are using symbolic + #(UPat(Ops.CONST, name="c"), lambda c: c.replace(src=()) if len(c.src) else None), +]) + +rangeify_codegen = PatternMatcher([ # add loads to non ptr indexes # TODO: this can be moved into codegen? (UPat((Ops.DEFINE_GLOBAL, Ops.STORE), name="dg").f(Ops.INDEX, name="idx", allow_any_len=True), @@ -387,21 +397,17 @@ to_define_global = PatternMatcher([ # TODO: this can be moved into codegen (UPat(Ops.STORE, name="store").f(Ops.INDEX, allow_any_len=True, name="idx").f(Ops.LOAD), - lambda store,idx: idx.replace(src=(store.as_buf(),)+idx.src[1:]).load(store)), - - # HACK in case any CONSTs were replaced - # this is only needed if you are using symbolic - #(UPat(Ops.CONST, name="c"), lambda c: c.replace(src=()) if len(c.src) else None), + lambda store,idx: idx.replace(src=(store.as_buf(),)+idx.src[1:]).load(store if idx.dtype.addrspace != AddrSpace.LOCAL else store.barrier())), ]) def split_store(x:UOp): if len(x.ranges): return None ctx = LocalAddBufferContext() - ret = graph_rewrite(x, to_define_global, ctx=ctx, name="kernel split", bottom_up=True) + ret = graph_rewrite(x, to_define_global+rangeify_codegen, ctx=ctx, name="kernel split", bottom_up=True) - store_rngs = ret.src[2:] + # get name rng = sorted([u for u in ret.toposort() if u.op is Ops.RANGE], key=lambda x: x.arg) - name = "k"+colored('_', 'BLACK').join(['']+[colored(s.src[0].render(), "WHITE" if s in store_rngs else "red") for s in rng]) + name = "k"+colored('_', 'BLACK').join(['']+[colored(s.src[0].render(), "WHITE" if s in ret.src[2:] else "red") for s in rng]) # NOTE: the hack for COPY is here ret = ret.sink(arg=KernelInfo(name=name)) if ret.src[1].op is not Ops.COPY else ret.src[1]