From 257f7d6d0339517a5b2ed1005cd986ffde571f8a Mon Sep 17 00:00:00 2001 From: George Hotz Date: Wed, 27 Aug 2025 21:01:26 -0700 Subject: [PATCH] simplify_merge_adjacent --- tinygrad/codegen/__init__.py | 4 +- tinygrad/codegen/opt/postrange.py | 71 +++++++++++++++++++++++++------ 2 files changed, 60 insertions(+), 15 deletions(-) diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index e19d3a98de..9a3ba916b7 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -18,7 +18,7 @@ from tinygrad.codegen.late.devectorizer import load_store_folding, load_store_in from tinygrad.codegen.late.linearize import block_create, pm_blockend_merge, block_merge, pm_finalize, BlockContext from tinygrad.codegen.opt import pm_get_optimization, pm_do_optimize from tinygrad.codegen.opt.swizzler import view_left, view_right, fix_kernel_ops -from tinygrad.codegen.opt.postrange import pm_postrange_opt_early, pm_postrange_opt +from tinygrad.codegen.opt.postrange import pm_postrange_opt_early, pm_postrange_opt, pm_postrange_opt_merge from tinygrad.schedule.rangeify import pm_add_buffers_local, rangeify_codegen @dataclass @@ -68,6 +68,8 @@ def _get_rewrites_for_renderer(opts:Renderer, linearizer:bool, _QUANTIZE, _DEVEC ret.append(RewriteStep(sym+migrate_indexing, name="initial symbolic")) if _POSTOPT or _RANGEIFY: + ret.append(RewriteStep(pm_postrange_opt_merge, ctx=lambda _: ({}, opts), name="early range merge")) + ret.append(RewriteStep(sym, name="mid symbolic")) ret.append(RewriteStep(pm_postrange_opt_early, ctx=lambda _: ({}, opts), name="early post opt ast")) ret.append(RewriteStep(sym, name="mid symbolic")) ret.append(RewriteStep(pm_postrange_opt, ctx=lambda _: opts, name="post optimize ast")) diff --git a/tinygrad/codegen/opt/postrange.py b/tinygrad/codegen/opt/postrange.py index 258d6d6ce0..4114277b29 100644 --- a/tinygrad/codegen/opt/postrange.py +++ b/tinygrad/codegen/opt/postrange.py @@ -1,13 +1,65 @@ import math from dataclasses import replace from tinygrad.dtype import dtypes -from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, KernelInfo -from tinygrad.uop.symbolic import sym +from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, KernelInfo, graph_rewrite, _substitute +from tinygrad.uop.symbolic import symbolic from tinygrad.helpers import colored, USE_TC, DEBUG from tinygrad.codegen.opt.kernel import axis_colors, AxisType from tinygrad.renderer import Renderer from tinygrad.codegen.opt.tc import TensorCore +def count_divmod(x:UOp): return len([u for u in x.toposort() if u.op in {Ops.IDIV, Ops.MOD}]) + +def flatten_range_in_terminators(r:UOp): + off = 2 if r.op is Ops.STORE else 1 + rngs = r.src[off:] + if not len(rngs): return None + new_rngs = [x for x in UOp.sink(*rngs).toposort() if x.op is Ops.RANGE] + return r.replace(src=r.src[:off]+tuple(new_rngs)) + +pm_flatten_range = PatternMatcher([ + # flatten ranges + (UPat((Ops.REDUCE, Ops.STORE), name="r"), flatten_range_in_terminators), +]) + +# NOTE: this one is better than the one in kernel.py +def simplify_merge_adjacent(ast:UOp): + # get all ranges (sorted) + rng = sorted([u for u in ast.parents if u.op is Ops.RANGE], key=lambda x: x.arg[0:-1]) + terminators = [u for u in ast.parents if u.op in {Ops.REDUCE, Ops.STORE}] + termination = {} + for t in terminators: + for u in t.src[1 if t.op is Ops.REDUCE else 2:]: termination[u] = t + + replaces = {} + i = 0 + while i < len(rng)-1: + r0, r1 = rng[i], rng[i+1] + # same axistype and same termination + if r0.arg[1] == r1.arg[1] and termination[r0] == termination[r1]: + s0, s1 = r0.src[0], r1.src[0] + new_range = r0.replace(src=(s0*s1,)).simplify() + # this checks the legality of a merge + oidx = ast.simplify() + nidx = graph_rewrite(oidx, _substitute+symbolic+pm_flatten_range, ctx={r0:new_range//s1, r1:new_range%s1}, name=f"check_merge_{i}_{i+1}") + # it simplifies + if count_divmod(nidx) <= count_divmod(oidx): + # it is correct + midx = graph_rewrite(nidx, _substitute+symbolic+pm_flatten_range, ctx={new_range:r0*s1+r1}, name=f"correct_merge_{i}_{i+1}") + if oidx is midx: + termination[new_range] = termination[r0] + replaces[r0] = new_range//s1 + replaces[r1] = new_range%s1 + rng[i] = new_range + del rng[i+1] + continue + i += 1 + return ast.substitute(replaces, name="simplify_merge_adjacent") + +pm_postrange_opt_merge = PatternMatcher([ + (UPat(Ops.SINK, name="ast"), simplify_merge_adjacent), +]) + def apply_tensor_cores(ctx:tuple[dict, Renderer], in0:UOp, in1:UOp, r_range:UOp, reduceop:UOp): if not USE_TC: return None # tensor cores have three ranges. X, Y, and REDUCE @@ -89,7 +141,7 @@ def early_sink(ctx:tuple[dict, Renderer], s:UOp): pm_postrange_opt_early = PatternMatcher([ # TODO: this is optional (and can have internal options) and we need a way to express that - ((UPat.var("in0")*UPat.var("in1")).reduce(UPat(Ops.RANGE, name="r_range"), name="reduceop", arg=Ops.ADD), apply_tensor_cores), + #((UPat.var("in0")*UPat.var("in1")).reduce(UPat(Ops.RANGE, name="r_range"), name="reduceop", arg=Ops.ADD), apply_tensor_cores), (UPat(Ops.SINK, name="s"), early_sink), ]) @@ -103,7 +155,7 @@ def split_range(r:UOp): if r.arg[-1] not in {AxisType.LOOP, AxisType.GLOBAL, AxisType.REDUCE}: return None if r.tag is not None: return None # any divisor is an option - is_local = False # if r.arg[-1] is AxisType.REDUCE else True + is_local = False if r.arg[-1] is AxisType.REDUCE else True N = 4 rd = r.src[0].divides(N) if rd is None: return None @@ -111,13 +163,6 @@ def split_range(r:UOp): er = UOp(Ops.RANGE, dtypes.int, src=(UOp.const(dtypes.int, N),), arg=r.arg[0:-1]+(1, axis_typemap[(r.arg[-1] is AxisType.REDUCE, is_local)])) return sr*N+er -def flatten_range_in_terminators(r:UOp): - off = 2 if r.op is Ops.STORE else 1 - rngs = r.src[off:] - if not len(rngs): return None - new_rngs = [x for x in UOp.sink(*rngs).toposort() if x.op is Ops.RANGE] - return r.replace(src=r.src[:off]+tuple(new_rngs)) - def rename_sink(s:UOp): if s.arg is not None and s.arg.name != "test": return None @@ -128,13 +173,11 @@ def rename_sink(s:UOp): name = "k" + colored('_', 'BLACK').join(['']+[colored(x.src[0].render(), axis_colors[x.arg[-1]]) for x in rngs]) return s.replace(arg=KernelInfo(name=name) if s.arg is None else replace(s.arg, name=name)) -pm_postrange_opt = PatternMatcher([ +pm_postrange_opt = pm_flatten_range+PatternMatcher([ # TODO: this is optional (and can have internal options) and we need a way to express that (UPat(Ops.RANGE, name="r"), split_range), # remove axes with 1 (UPat(Ops.RANGE, name="r"), lambda r: r.const_like(0) if r.vmax == 0 else None), - # flatten ranges - (UPat((Ops.REDUCE, Ops.STORE), name="r"), flatten_range_in_terminators), # run this last (UPat(Ops.SINK, name="s"), rename_sink), ])