From 1e6b70568bf484c5ed012579077be1cddd9e48b5 Mon Sep 17 00:00:00 2001 From: George Hotz Date: Sun, 24 Aug 2025 16:42:28 -0700 Subject: [PATCH] working WMMA --- tinygrad/codegen/opt/kernel.py | 54 +++++++++++++------------ tinygrad/codegen/opt/postrange.py | 66 ++++++++++++++++++++++++++++++- 2 files changed, 92 insertions(+), 28 deletions(-) diff --git a/tinygrad/codegen/opt/kernel.py b/tinygrad/codegen/opt/kernel.py index d3d704f9c1..d0580b09e3 100644 --- a/tinygrad/codegen/opt/kernel.py +++ b/tinygrad/codegen/opt/kernel.py @@ -290,23 +290,23 @@ class Kernel: # it's disabled for now since it makes BEAM slow for little gain check(self.opts.has_local, "target does not support local") check(self.axis_types[axis] is AxisType.GLOBAL, "local is for globals") - self.shift_to(axis, amt, AxisType.LOCAL, insert_at=max(self.axes_of(AxisType.GLOBAL, AxisType.LOCAL))+1) + return self.shift_to(axis, amt, AxisType.LOCAL, insert_at=max(self.axes_of(AxisType.GLOBAL, AxisType.LOCAL))+1) elif opt.op in {OptOps.GROUP, OptOps.GROUPTOP}: # green check(self.opts.has_local and self.opts.has_shared, "target does not support local or shared mem") check(self.axis_types[axis] is AxisType.REDUCE, "must be reduce axis to group") check(not self.tensor_core, "can't group with tensor cores") check(len(reduce_axes:=[i for r in self.reduceops for i in r.axis_arg]) == len(set(reduce_axes)), "can't group with parallel reduces") - self.shift_to(axis, amt, AxisType.GROUP_REDUCE, top=(opt.op is OptOps.GROUPTOP), insert_at=min(self.axes_of(AxisType.REDUCE))) + return self.shift_to(axis, amt, AxisType.GROUP_REDUCE, top=(opt.op is OptOps.GROUPTOP), insert_at=min(self.axes_of(AxisType.REDUCE))) elif opt.op is OptOps.UNROLL: # purple check(self.axis_types[axis] not in (AxisType.UPCAST, AxisType.UNROLL), "can't upcasted already upcasted") check(amt <= 32, "don't unroll more than 32") - self.shift_to(axis, amt, AxisType.UNROLL, insert_at=None) + return self.shift_to(axis, amt, AxisType.UNROLL, insert_at=None) elif opt.op is OptOps.UPCAST: # yellow check(axis in self.upcastable_dims, f"{axis=} not in {self.upcastable_dims=}") # NOTE: assume the first get_local_axes() LOCAL are for TC check(not (self.tensor_core and axis in self.axes_of(AxisType.LOCAL)[:len(self.tensor_core.get_local_axes())]), "can't upcast TC locals") check((self.opts is not None and self.opts.device == "DSP") or amt <= 16, "don't upcast more than 16") - self.shift_to(axis, amt, AxisType.UPCAST, insert_at=max(self.axes_of(AxisType.GLOBAL, AxisType.LOCAL, AxisType.LOOP, AxisType.UPCAST))+1) + return self.shift_to(axis, amt, AxisType.UPCAST, insert_at=max(self.axes_of(AxisType.GLOBAL, AxisType.LOCAL, AxisType.LOOP, AxisType.UPCAST))+1) elif opt.op is OptOps.NOLOCALS: check(self.opts.has_local and not self.dont_use_locals, "NOLOCALS is meaningless if target does not support local or already not using locals") check(AxisType.LOCAL not in self.axis_types and self.group_for_reduces == 0, "can't have no locals with locals") @@ -444,6 +444,29 @@ class Kernel: return ret def shape_str_to_axis(self, nms:list[str]) -> tuple[int, ...]: return tuple([self.shape_str().index(x) for x in nms]) + def reduce_to_wmma(self, ret:UOp, tc:TensorCore, axes:list[int]): + # permute the srcs + srcs = list((ret.src[0] if ret.src[0].op is not Ops.CAST else ret.src[0].src[0]).src) + for i, (src, permaxis) in enumerate(zip(srcs, tc.permutes_for_shape_str(self.shape_str()))): + src_st = (src if src.op is Ops.LOAD else src.src[0]).st_arg + srcs[i] = src.view(ShapeTracker.from_shape(src_st.shape).permute(permaxis)) + + # get reduce/upcast axes for the tensor cores + tc_reduce_axes = self.shape_str_to_axis([f"r{i}" for i in range(len(tc.get_reduce_axes()))]) + base_upcast_axes = tuple([(s,2) for s in self.shape_str_to_axis(tc.base_upcast_axes())]) + tc_upcast_axes = tuple([base_upcast_axes[:int(math.log2(tc.elements_per_thread[i]))] for i in range(3)]) + + # construct the op + wmma_arg = (str(tc), tc.dims, tc.dtype_in, tc.dtype_out, self.opts.device, tc.threads, tc_upcast_axes, tc_reduce_axes) + wmma = UOp(Ops.WMMA, dtype=tc.dtype_out.vec(tc.elements_per_thread[2]), src=( + UOp(Ops.CONTRACT, dtype=srcs[0].dtype.vec(tc.elements_per_thread[0]), src=(srcs[0],), arg=tc_upcast_axes[0]), + UOp(Ops.CONTRACT, dtype=srcs[1].dtype.vec(tc.elements_per_thread[1]), src=(srcs[1],), arg=tc_upcast_axes[1]), + UOp.const(tc.dtype_out.vec(tc.elements_per_thread[2]), 0.0)), arg=wmma_arg) + tc_uop = UOp(Ops.UNROLL, tc.dtype_out, (wmma,), arg=tc_upcast_axes[2]) + + # preserve any other reduce + return ret.replace(src=(tc_uop,), arg=(Ops.ADD, new_axes)) if (new_axes := tuple(i for i in axes if i not in tc_reduce_axes)) else tc_uop + def get_optimized_ast(self, name_override:str|None=None) -> UOp: @functools.cache def fixup_ast(op:UOp) -> UOp: @@ -463,28 +486,7 @@ class Kernel: axes = tuple(i for i in self.axes_of(AxisType.REDUCE, AxisType.UNROLL) if i in changed) grouped_axes = tuple(i for i in self.axes_of(AxisType.GROUP_REDUCE) if i in changed) if (tc := self.tensor_core) and self.use_tensor_cores == 1: - # get reduce/upcast axes for the tensor cores - tc_reduce_axes = self.shape_str_to_axis([f"r{i}" for i in range(len(tc.get_reduce_axes()))]) - base_upcast_axes = tuple([(s,2) for s in self.shape_str_to_axis(tc.base_upcast_axes())]) - tc_upcast_axes = tuple([base_upcast_axes[:int(math.log2(tc.elements_per_thread[i]))] for i in range(3)]) - - # permute the srcs - srcs = list((ret.src[0] if ret.src[0].op is not Ops.CAST else ret.src[0].src[0]).src) - for i, (src, permaxis) in enumerate(zip(srcs, tc.permutes_for_shape_str(self.shape_str()))): - src_st = (src if src.op is Ops.LOAD else src.src[0]).st_arg - srcs[i] = src.view(ShapeTracker.from_shape(src_st.shape).permute(permaxis)) - - # construct the op - wmma_arg = (str(tc), tc.dims, tc.dtype_in, tc.dtype_out, self.opts.device, tc.threads, tc_upcast_axes, tc_reduce_axes) - wmma = UOp(Ops.WMMA, dtype=tc.dtype_out.vec(tc.elements_per_thread[2]), src=( - UOp(Ops.CONTRACT, dtype=srcs[0].dtype.vec(tc.elements_per_thread[0]), src=(srcs[0],), arg=tc_upcast_axes[0]), - UOp(Ops.CONTRACT, dtype=srcs[1].dtype.vec(tc.elements_per_thread[1]), src=(srcs[1],), arg=tc_upcast_axes[1]), - UOp.const(tc.dtype_out.vec(tc.elements_per_thread[2]), 0.0)), arg=wmma_arg) - tc_uop = UOp(Ops.UNROLL, tc.dtype_out, (wmma,), arg=tc_upcast_axes[2]) - - # preserve any other reduce - return ret.replace(src=(tc_uop,), arg=(Ops.ADD, new_axes)) if (new_axes := tuple(i for i in axes if i not in tc_reduce_axes)) else tc_uop - + return self.reduce_to_wmma(ret, tc, axes) ret = ret.replace(arg = (op.arg[0], axes)) if self.group_for_reduces and grouped_axes: local_axes = tuple([i for i,t in enumerate(self.axis_types) if t in (AxisType.LOCAL, AxisType.UPCAST) or i in grouped_axes]) diff --git a/tinygrad/codegen/opt/postrange.py b/tinygrad/codegen/opt/postrange.py index a77759d749..d1e441dc05 100644 --- a/tinygrad/codegen/opt/postrange.py +++ b/tinygrad/codegen/opt/postrange.py @@ -1,5 +1,6 @@ +import math from tinygrad.uop.ops import UOp, Ops, sint, ssimplify, AxisType, KernelInfo -from tinygrad.codegen.opt.kernel import Kernel +from tinygrad.codegen.opt.kernel import Kernel, Opt, OptOps from tinygrad.renderer import Renderer from tinygrad.dtype import dtypes @@ -16,6 +17,66 @@ class RKernel(Kernel): self.replaces.update(dict(zip(self.rng, rng))) self.rng = rng + def _apply_tc_opt(self, use_tensor_cores:int, axis:int, tc_select:int, opt_level:int) -> bool: + reduceop = [x for x in self.ast.toposort() if x.op is Ops.REDUCE][0] + if use_tensor_cores and reduceop is not None and reduceop.arg is Ops.ADD: + tensor_cores = self.opts.tensor_cores if tc_select == -1 else [self.opts.tensor_cores[tc_select]] + for tc in tensor_cores: + if tc.dtype_in == dtypes.float and tc.dtype_out == dtypes.float: + axes = [0,1] + ne = [] + un, ln = 0, 0 + ss = [] + for opt in tc.opts: + ne.append(self.apply_opt(Opt({"u":OptOps.UPCAST, "l":OptOps.LOCAL}[opt[0]], axes[int(opt[1])], 2), append_opt=False)) + if opt[0] == 'u': + ss.append(f"u{un}") + un += 1 + if opt[0] == 'l': + ss.append(f"l{ln}") + ln += 1 + for i, (_, amt) in enumerate(tc.get_reduce_axes()): + ne.append(self.apply_opt(Opt(OptOps.UNROLL, 0, amt), append_opt=False)) # TODO: this should be the reduce, not 0 + ss.append(f"r{i}") + + # early realize for TC + self.ast = self.ast.substitute(self.replaces) + self.replaces = {} + reduceop = [x for x in self.ast.toposort() if x.op is Ops.REDUCE][0] + tne = [x.replace(tag=1) for x in ne] + treduceop = reduceop.substitute(dict(zip(ne, tne))) + mul = treduceop.src[0] + assert mul.op is Ops.MUL + p1, p2 = tc.permutes_for_shape_str(ss) + m1 = mul.src[0].substitute(dict(zip(tne, [ne[i] for i in p1]))) + m2 = mul.src[1].substitute(dict(zip(tne, [ne[i] for i in p2]))) + srcs = [m1, m2] + + # get reduce/upcast axes for the tensor cores + tc_reduce_axes = self.shape_str_to_axis([f"r{i}" for i in range(len(tc.get_reduce_axes()))]) + base_upcast_axes = tuple([(s,2) for s in self.shape_str_to_axis(tc.base_upcast_axes())]) + tc_upcast_axes = tuple([base_upcast_axes[:int(math.log2(tc.elements_per_thread[i]))] for i in range(3)]) + + tc_reduce_axes = tuple([self.rng[x].arg[0] for x in tc_reduce_axes]) + # TODO: remove tc_upcast_axes from the arg + #tc_upcast_axes = [(self.rng[x[0]].arg[0], x[1]) for x in tc_upcast_axes] + + # construct the op + wmma_arg = (str(tc), tc.dims, tc.dtype_in, tc.dtype_out, self.opts.device, tc.threads, tc_upcast_axes, tc_reduce_axes) + wmma = UOp(Ops.WMMA, dtype=tc.dtype_out.vec(tc.elements_per_thread[2]), src=( + UOp(Ops.CONTRACT, dtype=srcs[0].dtype.vec(tc.elements_per_thread[0]), src=(srcs[0],), arg=tc_upcast_axes[0]), + UOp(Ops.CONTRACT, dtype=srcs[1].dtype.vec(tc.elements_per_thread[1]), src=(srcs[1],), arg=tc_upcast_axes[1]), + UOp.const(tc.dtype_out.vec(tc.elements_per_thread[2]), 0.0)), arg=wmma_arg) + tc_uop = UOp(Ops.UNROLL, tc.dtype_out, (wmma,), arg=tc_upcast_axes[2]) + + reduce_ranges = [x for x in UOp.sink(*reduceop.src[1:]).toposort() if x.op is Ops.RANGE and x.arg[0] not in tc_reduce_axes] + if len(reduce_ranges): + tc_uop = UOp(Ops.REDUCE, tc_uop.dtype, (tc_uop,)+tuple(reduce_ranges), Ops.ADD) + + self.ast = self.ast.substitute({reduceop: tc_uop}) + return True + return False + def shift_to(self, axis:int, amount:int, new_type:AxisType, top:bool=False, insert_at:int|None=None): old_sz = self.rng[axis].src[0].arg // amount assert old_sz > 0, f"bad old_sz on {axis} {amount} {self.rng[axis]}" @@ -29,9 +90,10 @@ class RKernel(Kernel): del self.rng[axis] else: replaced_rng = self.rng[axis].replace(src=(UOp.const(dtypes.int, old_sz),)) - self.replaces[self.rng[axis]] = replaced_rng * amount + new_rng + self.replaces[self.rng[axis]] = (new_rng * old_sz + replaced_rng) if top else (replaced_rng * amount + new_rng) self.rng[axis] = replaced_rng self.rng.insert(insert_at if insert_at is not None else len(self.rng), new_rng) + return new_rng @property def axis_types(self) -> list[AxisType]: return [x.arg[1] for x in self.rng]