From 6ef3270fc80cae1db234dc040f7398b9d721bb65 Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Fri, 5 Sep 2025 18:59:54 -0700 Subject: [PATCH] fix opt gate (#12050) --- tinygrad/codegen/__init__.py | 33 +++++++++++++++++---------------- tinygrad/codegen/simplify.py | 2 +- 2 files changed, 18 insertions(+), 17 deletions(-) diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index 63b9679e16..faa0219115 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -44,28 +44,29 @@ rewrites_for_linearizer = [ RewriteStep(block_merge, name="Linearizer: Merge Blocks"), RewriteStep(pm_finalize, name="Linearizer: Finalize")] -def get_rewrites_for_renderer(opts:Renderer, linearizer:bool=True) -> list[RewriteStep]: +def get_rewrites_for_renderer(opts:Renderer, optimize:bool=True, linearizer:bool=True) -> list[RewriteStep]: # cache with the values of the context vars - return _get_rewrites_for_renderer(opts, linearizer, QUANTIZE.value, DEVECTORIZE.value, TRANSCENDENTAL.value) + return _get_rewrites_for_renderer(opts, optimize, linearizer, QUANTIZE.value, DEVECTORIZE.value, TRANSCENDENTAL.value) @functools.cache -def _get_rewrites_for_renderer(opts:Renderer, linearizer:bool, _QUANTIZE, _DEVECTORIZE, _TRANSCENDENTAL) -> list[RewriteStep]: +def _get_rewrites_for_renderer(opts:Renderer, optimize:bool, linearizer:bool, _QUANTIZE, _DEVECTORIZE, _TRANSCENDENTAL) -> list[RewriteStep]: # ** lowerer (rewrite_shapetracker_with_index) ** ret: list[RewriteStep] = [] - # view pushing - ret.extend(rewrites_for_views) + if optimize: + # view pushing + ret.extend(rewrites_for_views) - # lowerer first - if _QUANTIZE and opts.device in {"CPU", "DSP"}: ret.append(RewriteStep(pm_quant, name="quantize")) - ret.append(RewriteStep(pm_lowerer, get_index, name="lowerer", bottom_up=True)) + # lowerer first + if _QUANTIZE and opts.device in {"CPU", "DSP"}: ret.append(RewriteStep(pm_quant, name="quantize")) + ret.append(RewriteStep(pm_lowerer, get_index, name="lowerer", bottom_up=True)) - # symbolic (NOTE: this is a requirement for pm_simplify_ranges to be correct) - ret.append(RewriteStep(sym, name="initial symbolic")) + # symbolic (NOTE: this is a requirement for pm_simplify_ranges to be correct) + ret.append(RewriteStep(sym, name="initial symbolic")) - # optimize (schedule) the AST - ret.append(RewriteStep(pm_simplify_ranges, name="simplify ranges")) - ret.append(RewriteStep(pm_postrange_opt, ctx=lambda _: opts, name="post optimize ast")) + # optimize (schedule) the AST + ret.append(RewriteStep(pm_simplify_ranges, name="simplify ranges")) + ret.append(RewriteStep(pm_postrange_opt, ctx=lambda _: opts, name="post optimize ast")) # ** expander (expand_rewrite) ** ret.append(RewriteStep(sym+migrate_indexing, name="postopt symbolic")) @@ -106,8 +107,8 @@ def _get_rewrites_for_renderer(opts:Renderer, linearizer:bool, _QUANTIZE, _DEVEC # return the list (with optional linearizer) return ret + (rewrites_for_linearizer if linearizer else []) -def full_rewrite_to_sink(sink:UOp, opts:Renderer|None=None, linearizer:bool=False) -> UOp: - return apply_rewrites(sink, get_rewrites_for_renderer(opts if opts is not None else Renderer(), linearizer)) +def full_rewrite_to_sink(sink:UOp, opts:Renderer|None=None, optimize:bool=True, linearizer:bool=False) -> UOp: + return apply_rewrites(sink, get_rewrites_for_renderer(opts if opts is not None else Renderer(), optimize, linearizer)) def full_rewrite(sink:UOp, opts:Renderer|None=None) -> list[UOp]: """ @@ -121,6 +122,6 @@ def full_rewrite(sink:UOp, opts:Renderer|None=None) -> list[UOp]: Linear program in UOps. """ - lst = list(full_rewrite_to_sink(sink, opts, linearizer=True).arg.lst) + lst = list(full_rewrite_to_sink(sink, opts, optimize=sink.tag is None, linearizer=True).arg.lst) if __debug__: type_verify(lst) return lst diff --git a/tinygrad/codegen/simplify.py b/tinygrad/codegen/simplify.py index bdf6b39d1c..bc4e57e066 100644 --- a/tinygrad/codegen/simplify.py +++ b/tinygrad/codegen/simplify.py @@ -14,7 +14,7 @@ pm_flatten_range = PatternMatcher([ ]) def count_divmod(x:UOp): return len([u for u in x.toposort() if u.op in {Ops.IDIV, Ops.MOD}]) -def simplify_merge_adjacent(u:UOp) -> UOp|None: +def simplify_merge_adjacent(ctx:UOp, u:UOp) -> UOp|None: i = 2 if u.op is Ops.STORE else 1 while i < len(u.src)-1: r0, r1 = u.src[i], u.src[i+1]