From 9cd365c12ecc4a20d10bbdf636b5db690241610f Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Fri, 3 Oct 2025 10:31:51 +0800 Subject: [PATCH] little changes from double gemm (#12429) * little changes from double gemm * split pm_group_for_reduce * pm_add_buffers_local * Revert "pm_add_buffers_local" This reverts commit 4d30a91db2c720e146a63a343dc3593e03a4df61. --- tinygrad/codegen/__init__.py | 4 ++-- tinygrad/codegen/late/expander.py | 3 +++ tinygrad/renderer/cstyle.py | 3 ++- tinygrad/uop/ops.py | 2 +- 4 files changed, 8 insertions(+), 4 deletions(-) diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index f93b639007..bcd3a9ea5c 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -12,7 +12,7 @@ from tinygrad.codegen.quantize import pm_quant from tinygrad.codegen.gpudims import pm_add_gpudims from tinygrad.uop.symbolic import sym, symbolic_simple, gep_pushing from tinygrad.uop.decompositions import get_late_rewrite_patterns -from tinygrad.codegen.late.expander import migrate_indexing, expander, pm_pre_expander +from tinygrad.codegen.late.expander import migrate_indexing, expander, pm_pre_expander, pm_group_for_reduce from tinygrad.codegen.late.devectorizer import load_store_folding, load_store_indexing, devectorize, pm_reduce, \ ReduceContext, correct_load_store, pm_render from tinygrad.codegen.late.linearize import block_create, pm_blockend_merge, block_merge, pm_finalize, BlockContext @@ -78,7 +78,7 @@ def _get_rewrites_for_renderer(opts:Renderer, optimize:bool, linearizer:bool, _Q ret.append(RewriteStep(sym+migrate_indexing, name="postopt symbolic")) # expand - ret.append(RewriteStep(sym+pm_pre_expander+expander, name="expander")) + ret.append(RewriteStep(sym+pm_pre_expander+pm_group_for_reduce+expander, name="expander")) # add locals ret.append(RewriteStep(pm_add_buffers+rangeify_codegen, name="add local buffers")) diff --git a/tinygrad/codegen/late/expander.py b/tinygrad/codegen/late/expander.py index d2d0a41162..9a42d414ce 100644 --- a/tinygrad/codegen/late/expander.py +++ b/tinygrad/codegen/late/expander.py @@ -157,6 +157,9 @@ pm_pre_expander = PatternMatcher([ # fix REDUCEs with UNROLLs (UPat(Ops.REDUCE, name="x"), fix_reduce_unroll), (UPat(Ops.STORE, name="x"), fix_store_unroll), +]) + +pm_group_for_reduce = PatternMatcher([ # fix group for reduce (UPat(Ops.REDUCE, name="x"), fix_group_for_reduce), ]) diff --git a/tinygrad/renderer/cstyle.py b/tinygrad/renderer/cstyle.py index 92ee79474e..0916f18ccd 100644 --- a/tinygrad/renderer/cstyle.py +++ b/tinygrad/renderer/cstyle.py @@ -320,7 +320,8 @@ class MetalRenderer(CStyleLanguage): def render_kernel(self, function_name, kernel, bufs, uops, prefix=None): prefix = ["#include ","using namespace metal;"] - for name, _, dtype_in, dtype_out, _, _, _, _ in wmma_args(uops): prefix.append( + deduped_wmma_args = dedup([(name, dtype_in, dtype_out) for name, _, dtype_in, dtype_out, _, _, _, _ in wmma_args(uops)]) + for name, dtype_in, dtype_out in deduped_wmma_args: prefix.append( f"""{(dstr_out:=self.render_dtype(dtype_out.vec(2)))} __{name}({(dstr_in:=self.render_dtype(dtype_in.vec(2)))} a, {dstr_in} b, {dstr_out} c){{ simdgroup_{self.render_dtype(dtype_in)}8x8 mat_a, mat_b; simdgroup_{self.render_dtype(dtype_out)}8x8 mat_c; mat_a.thread_elements()[0] = a[0]; mat_b.thread_elements()[0] = b[0]; mat_c.thread_elements()[0] = c[0]; diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index c7326d9379..c0a21bfa29 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -481,7 +481,7 @@ class UOp(MathTrait, metaclass=UOpMetaClass): if self.op is Ops.MSTACK: return UOp(Ops.MSTACK, self.dtype, src=tuple(x.as_buf() for x in self.src)) # TODO: this should be the only one of these. this is the one RANGEIFY uses s = self - while len(s.src) and s.op not in {Ops.BUFFER, Ops.MSTACK}: s = s.src[0] + while len(s.src) and s.op not in {Ops.BUFFER, Ops.BUFFERIZE, Ops.MSTACK}: s = s.src[0] return s @property