From 2aebb6f6c40ea423f69fd14c717bc554b5a325ff Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Tue, 7 Jul 2026 14:30:32 -0700 Subject: [PATCH] move mop cleanup [pr] (#16915) * move mop cleanup [pr] * those are mop cleanups * more mop_cleanups * revert * movement.py --- tinygrad/codegen/__init__.py | 13 +++++++------ tinygrad/engine/jit.py | 2 +- tinygrad/schedule/rangeify.py | 12 +----------- tinygrad/uop/movement.py | 19 +++++++++++++++++++ tinygrad/uop/symbolic.py | 8 ++------ 5 files changed, 30 insertions(+), 24 deletions(-) create mode 100644 tinygrad/uop/movement.py diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index 013e734aa3..eeaa46cd63 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -13,6 +13,7 @@ from tinygrad.dtype import dtypes, AddrSpace # import all pattern matchers here from tinygrad.codegen.gpudims import pm_add_gpudims from tinygrad.uop.symbolic import sym, symbolic_simple, symbolic, pm_move_where_on_load, pm_clean_up_group_sink, pm_remove_invalid +from tinygrad.uop.movement import mop_cleanup from tinygrad.codegen.decomp.dtype import pm_dtype_decomps from tinygrad.codegen.decomp.op import get_late_rewrite_patterns, get_simplifying_rewrite_patterns from tinygrad.codegen.decomp.transcendental import get_transcendental_patterns @@ -20,7 +21,7 @@ from tinygrad.codegen.late.coalese import indexing_simplify from tinygrad.codegen.opt.postrange import apply_opts from tinygrad.codegen.late.gater import pm_move_gates_from_index from tinygrad.codegen.simplify import pm_simplify_ranges, pm_flatten_range, pm_split_ranges, pm_load_collapse -from tinygrad.schedule.rangeify import pm_mops, pm_syntactic_sugar, mop_cleanup +from tinygrad.schedule.rangeify import pm_mops, pm_syntactic_sugar from tinygrad.codegen.late.linearizer import CFGContext, pm_split_ends, pm_add_control_flow, linearize from tinygrad.codegen.late.regalloc import LinearScanRegallocContext, pm_regalloc_rewrite from tinygrad.codegen.late.coalese import memory_coalesing, pm_simplify_add_image @@ -152,13 +153,13 @@ ew_devectorizer = PatternMatcher([ (UPat(GroupOp.Elementwise, name="b"), do_devectorize), ]) -devectorizer2 = pm_mops+PatternMatcher([ +devectorizer2 = mop_cleanup+pm_mops+PatternMatcher([ # unpack broadcasting (UPat(GroupOp.Elementwise|{Ops.LOAD,Ops.STORE}, name="b"), do_devectorize), - # const INDEX into STACK is src (this is symbolic) + # const INDEX into STACK is src (TODO: this should be in mop_cleanup) (UPat(Ops.INDEX, src=(UPat(Ops.STACK, name="a"), UPat.cvar("i")), name="idx", allow_any_len=True), lambda a,i,idx: a.src[i.arg].index(*idx.src[2:])), - # INDEX without src is nothing + # INDEX without src is nothing (TODO: this should be in mop_cleanup) (UPat(Ops.INDEX, src=(UPat.var('x'),)), lambda x: x), # unpack WMMA (UPat(Ops.WMMA, name="u"), do_stack_wmma), @@ -249,7 +250,7 @@ pm_reduce_local = pm_wmma_add+PatternMatcher([ ])+pm_clean_up_group_sink def maybe_load(u:UOp): return u.load() if u.addrspace in (AddrSpace.GLOBAL, AddrSpace.LOCAL, AddrSpace.REG) else u -pm_move_regs = PatternMatcher([ +pm_add_loads = PatternMatcher([ # BITCAST? (UPat(GroupOp.Elementwise|{Ops.REDUCE,Ops.WMMA,Ops.STACK}, name="x"), lambda x: x.replace(src=tuple([maybe_load(u) for u in x.src]))), (UPat(Ops.STORE, name="x"), lambda x: x.replace(src=(x.src[0], maybe_load(x.src[1]))+x.src[2:])), @@ -310,7 +311,7 @@ def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp: sink = graph_rewrite(sink, symbolic_simple+unbroadcast, name="*** unbroadcast") # add loads and remove invalids - sink = graph_rewrite(sink, pm_move_regs, name="** add loads") + sink = graph_rewrite(sink, pm_add_loads, name="** add loads") # devectorize sink = graph_rewrite(sink, symbolic_simple+devectorizer2, ctx=ren, name="devectorize2") diff --git a/tinygrad/engine/jit.py b/tinygrad/engine/jit.py index de20d380d4..fca73d24dc 100644 --- a/tinygrad/engine/jit.py +++ b/tinygrad/engine/jit.py @@ -10,7 +10,7 @@ from tinygrad.engine.realize import capturing, compile_linear, link_linear, run_ from tinygrad.engine.realize import unwrap_multi, resolve_params, get_call_arg_uops, get_call_outs_ins from tinygrad.schedule.memory import memory_plan_rewrite, _collect_bufs from tinygrad.nn.state import get_parameters -from tinygrad.schedule.rangeify import mop_cleanup +from tinygrad.uop.movement import mop_cleanup from dataclasses import dataclass def prune_linear(linear:UOp, needed:set[UOp]) -> tuple[UOp, UOp]: diff --git a/tinygrad/schedule/rangeify.py b/tinygrad/schedule/rangeify.py index 5d4af5fb9b..f75e0014b3 100644 --- a/tinygrad/schedule/rangeify.py +++ b/tinygrad/schedule/rangeify.py @@ -5,6 +5,7 @@ from tinygrad.dtype import dtypes, AddrSpace, Invalid from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, resolve, GroupOp, KernelInfo, ParamArg, shape_to_shape_arg from tinygrad.uop.ops import graph_rewrite, sint, AxisType, BottomUpGate, profile_matches, identity_element from tinygrad.uop.symbolic import symbolic +from tinygrad.uop.movement import mop_cleanup from tinygrad.helpers import prod, all_same, getenv, dedup, all_int, DEBUG, SPLIT_REDUCEOP, DEBUG_RANGEIFY, VIZ, MAX_KERNEL_BUFFERS from tinygrad.helpers import PCONTIG, FLOAT16, OPENPILOT_HACKS, argsort, partition, get_single_element from tinygrad.codegen.simplify import pm_flatten_range, pm_reduce_simplify @@ -97,17 +98,6 @@ def split_reduceop(reduce:UOp, x:UOp): # reduce original axes, then split return splitted._rop(reduce.arg[0], tuple(range(reduce.arg[1]))).contiguous()._rop(reduce.arg[0], (len(reduce.shape),)).reshape(reduce.shape) -mop_cleanup = PatternMatcher([ - # merge adjacent RESHAPES - (UPat(Ops.RESHAPE, src=(UPat(Ops.RESHAPE, name="x2"), UPat()), name="x"), lambda x,x2: x.replace(src=(x2.src[0], x.src[1]))), - # remove noop RESHAPEs - (UPat(Ops.RESHAPE, src=(UPat(name="x2"), UPat()), name="x"), lambda x,x2: x2 if x2._shape is not None and x2.shape == x.shape else None), - # merge PERMUTEs - (UPat(Ops.PERMUTE, src=(UPat(Ops.PERMUTE, name="x2"),), name="x"), lambda x,x2: x2.replace(arg=tuple(x2.arg[i] for i in x.arg))), - # remove noop PERMUTEs - (UPat(Ops.PERMUTE, name="x"), lambda x: x.src[0] if list(x.arg) == list(range(len(x.arg))) else None), -]) - pm_gather_params = PatternMatcher([ (UPat(Ops.PARAM, name="p"), lambda ctx, p: ctx.append(p) if p.arg.slot >= 0 else None), ]) def resolve_function(c:UOp, allow_param_mismatch=True) -> UOp|None: if c.arg.precompile: return None diff --git a/tinygrad/uop/movement.py b/tinygrad/uop/movement.py new file mode 100644 index 0000000000..f8e15a5a65 --- /dev/null +++ b/tinygrad/uop/movement.py @@ -0,0 +1,19 @@ +from tinygrad.uop.ops import PatternMatcher, UPat, Ops + +# TODO: pm_mops from rangeify belongs here. this is all pattern matchers that strictly clean up movement ops + +mop_cleanup = PatternMatcher([ + # merge adjacent RESHAPES + (UPat(Ops.RESHAPE, src=(UPat(Ops.RESHAPE, name="x2"), UPat()), name="x"), lambda x,x2: x.replace(src=(x2.src[0], x.src[1]))), + # remove noop RESHAPEs + (UPat(Ops.RESHAPE, src=(UPat(name="x2"), UPat()), name="x"), lambda x,x2: x2 if x2._shape is not None and x2.shape == x.shape else None), + # merge PERMUTEs + (UPat(Ops.PERMUTE, src=(UPat(Ops.PERMUTE, name="x2"),), name="x"), lambda x,x2: x2.replace(arg=tuple(x2.arg[i] for i in x.arg))), + # remove noop PERMUTEs + (UPat(Ops.PERMUTE, name="x"), lambda x: x.src[0] if list(x.arg) == list(range(len(x.arg))) else None), + # STACK on INDEX CONST + (UPat(Ops.STACK, src=UPat(Ops.INDEX, src=(UPat.var("src"), UPat(Ops.CONST))), name="stk"), + lambda src,stk: src if stk.shape == src.shape and list(range(len(stk.src))) == [x.src[1].arg for x in stk.src] else None), + # INDEX on STACK (simple) + (UPat(Ops.INDEX, src=(UPat(Ops.STACK, name="stk"), UPat(Ops.CONST, name="c"))), lambda stk,c: stk.src[c.arg]), +]) diff --git a/tinygrad/uop/symbolic.py b/tinygrad/uop/symbolic.py index 79551fb2a5..2581188ef8 100644 --- a/tinygrad/uop/symbolic.py +++ b/tinygrad/uop/symbolic.py @@ -5,6 +5,7 @@ from tinygrad.uop.ops import Ops, PatternMatcher, UPat, UOp, GroupOp, exec_alu from tinygrad.dtype import PyConst, ConstType, dtypes, can_lossless_cast, Invalid from tinygrad.helpers import partition, all_same, prod, flatten, unwrap, IMAGE, dedup from tinygrad.uop.divandmod import div_and_mod_symbolic +from tinygrad.uop.movement import mop_cleanup # TODO: symbolic shouldn't be importing from codegen from tinygrad.codegen.decomp.transcendental import xpow @@ -182,12 +183,7 @@ symbolic_simple = propagate_invalid + PatternMatcher([ (UPat.cvar("gate").where(UPat.var("c0"), UPat.var("c1")), lambda gate, c0, c1: c0 if gate.arg else c1), # a.where(b.where(c, d), d) -> (a & b).where(c, d) (UPat.var("a").where(UPat.var("b").where(UPat.var("c"), UPat.var("d")), UPat.var("d")), lambda a,b,c,d: (a&b).where(c,d)), - # STACK on INDEX CONST - (UPat(Ops.STACK, src=UPat(Ops.INDEX, src=(UPat.var("src"), UPat(Ops.CONST))), name="stk"), - lambda src,stk: src if stk.shape == src.shape and list(range(len(stk.src))) == [x.src[1].arg for x in stk.src] else None), - # INDEX on STACK - (UPat(Ops.INDEX, src=(UPat(Ops.STACK, name="stk"), UPat(Ops.CONST, name="c"))), lambda stk,c: stk.src[c.arg]), -]) +])+mop_cleanup # ******** phase 2 builds on phase 1, it includes the old "symbolic", rules that match deeper ********