diff --git a/tinygrad/renderer/cstyle.py b/tinygrad/renderer/cstyle.py index ec90358f41..9e0827834e 100644 --- a/tinygrad/renderer/cstyle.py +++ b/tinygrad/renderer/cstyle.py @@ -191,10 +191,7 @@ class OpenCLLanguage(CStyleLanguage): OpenCLRenderer = functools.partial(uops_to_cstyle, OpenCLLanguage()) class MetalLanguage(CStyleLanguage): - kernel_prefix = """#include \nusing namespace metal;\ntemplate U __metal_wmma(T m, T n, U o) { - S a,b,c; a.thread_elements()[0] = m.x; a.thread_elements()[1] = m.y; b.thread_elements()[0] = n.x; b.thread_elements()[1] = n.y; - c.thread_elements()[0] = o.x; c.thread_elements()[1] = o.y; simdgroup_multiply_accumulate(c, a, b, c); - return U(c.thread_elements()[0], c.thread_elements()[1]);\n}\nkernel """ + kernel_prefix = "kernel " buffer_prefix = "device " smem_prefix = "threadgroup " arg_int_prefix = "constant int&" @@ -205,6 +202,14 @@ class MetalLanguage(CStyleLanguage): extra_args = ['uint3 gid [[threadgroup_position_in_grid]]', 'uint3 lid [[thread_position_in_threadgroup]]'] def render_cast(self, x: List[str], var_dtype: DType, bitcast=False) -> str: return f"as_type<{var_dtype.name}>({x[0]})" if bitcast else super().render_cast(x, var_dtype) + + def render_kernel(self, function_name, kernel, bufs, local_size, uops, prefix=None): + prefix = ["#include ","using namespace metal;"] + if any(uop.uop == UOps.WMMA for uop in uops): prefix.append("""template U __metal_wmma(T m, T n, U o) { + S a,b,c; a.thread_elements()[0] = m.x; a.thread_elements()[1] = m.y; b.thread_elements()[0] = n.x; b.thread_elements()[1] = n.y; + c.thread_elements()[0] = o.x; c.thread_elements()[1] = o.y; simdgroup_multiply_accumulate(c, a, b, c); + return U(c.thread_elements()[0], c.thread_elements()[1]);\n}""") + return super().render_kernel(function_name, kernel, bufs, local_size, uops, prefix) MetalRenderer = functools.partial(uops_to_cstyle, MetalLanguage()) code_for_op_half = {