Compare commits

..
Author SHA1 Message Date
geohot 872ca998db simpler reduce that just removes ones 2025-07-25 19:17:38 -07:00
2 changed files with 9 additions and 16 deletions
+5 -6
View File
@@ -308,14 +308,13 @@ class MetalRenderer(CStyleLanguage):
]) + base_rewrite
def render_kernel(self, function_name, kernel, bufs, uops, prefix=None):
wmma_args = set([(uop.dtype, uop.src[0].dtype, uop.arg[0]) for uop in uops if uop.op is Ops.WMMA])
prefix = ["#include <metal_stdlib>","using namespace metal;"]
for dtype_out, dtype_in, name in wmma_args: prefix.append(
f"""{(dstr_out:=self.render_dtype(dtype_out))} __{name}({(dstr_in:=self.render_dtype(dtype_in))} a, {dstr_in} b, {dstr_out} c){{
simdgroup_{self.render_dtype(dtype_in.scalar())}8x8 mat_a, mat_b; simdgroup_{self.render_dtype(dtype_out.scalar())}8x8 mat_c;
prefix, wmma_args = ["#include <metal_stdlib>","using namespace metal;"], set([uop.arg for uop in uops if uop.op is Ops.WMMA])
for arg in wmma_args: prefix.append(
f"""{(dtype_out:=self.render_dtype(arg[3].vec(2)))} __{arg[0]}({(dtype_in:=self.render_dtype(arg[2].vec(2)))} a, {dtype_in} b, {dtype_out} c){{
simdgroup_{self.render_dtype(arg[2])}8x8 mat_a, mat_b; simdgroup_{self.render_dtype(arg[3])}8x8 mat_c;
mat_a.thread_elements()[0] = a[0]; mat_b.thread_elements()[0] = b[0]; mat_c.thread_elements()[0] = c[0];
mat_a.thread_elements()[1] = a[1]; mat_b.thread_elements()[1] = b[1]; mat_c.thread_elements()[1] = c[1];
simdgroup_multiply_accumulate(mat_c, mat_a, mat_b, mat_c);\n return {dstr_out}(mat_c.thread_elements()[0], mat_c.thread_elements()[1]);\n}}""")
simdgroup_multiply_accumulate(mat_c, mat_a, mat_b, mat_c);\n return {dtype_out}(mat_c.thread_elements()[0], mat_c.thread_elements()[1]);\n}}""")
return super().render_kernel(function_name, kernel, bufs, uops, prefix)
_nms = "xyzwabcdefghijkl"
+4 -10
View File
@@ -257,18 +257,12 @@ class UOp(MathTrait, metaclass=UOpMetaClass):
@staticmethod
def range(dtype:DType, end:sint, idx:int): return UOp(Ops.RANGE, dtype=dtype, src=(sint_to_uop(end),), arg=idx)
def r(self, op:Ops, axis:tuple[int, ...], permute=True):
assert permute
axis = tuple(sorted([x for x in axis if resolve(self.shape[x] != 1)]))
if len(axis) == 0: return self
# move any non reduce axis before the first reduce axis
move_early, rest = partition(range(axis[0], len(self.shape)), lambda i: i not in axis and resolve(self.shape[i] != 1))
if move_early and permute:
permaxis = tuple(range(axis[0])) + tuple(move_early) + tuple(rest)
ret = self.permute(permaxis)
new_axis = tuple([x for x in range(axis[0]+len(move_early), len(self.shape)) if resolve(ret.shape[x] != 1)])
assert len(axis) == len(new_axis)
else:
ret, new_axis = self, axis
ret = UOp(Ops.REDUCE_AXIS, self.dtype, (ret,), (op, new_axis))
ret = self.permute(tuple([i for i in range(len(self.shape)) if i not in axis])+axis)
ret = ret.reshape(tuple([s for s in ret.shape if resolve(s != 1)]))
ret = UOp(Ops.REDUCE_AXIS, self.dtype, (ret,), (op, tuple(range(len(ret.shape)-len(axis), len(ret.shape)))))
return ret.reshape(tuple([x if i not in axis else 1 for i,x in enumerate(self.shape)]))
def reduce(self, *src:UOp, **kwargs): return UOp(Ops.REDUCE, kwargs.pop('dtype', self.dtype), src=(self,)+src, **kwargs)
def contiguous(self): return self.alu(Ops.CONTIGUOUS)