From de4b9bf53bc8a6af75086e34ed555f3fae56b87f Mon Sep 17 00:00:00 2001 From: qazal <77887910+Qazalin@users.noreply.github.com> Date: Tue, 24 Jun 2025 18:55:39 +0300 Subject: [PATCH] add opts_to_apply option to AST KernelInfo (#10950) * proposal: add option to override opts in the get_program API * update test_linearizer_rewrite * state in uops * update process_replay and names * empty isn't none * fix process replay --- .../external/process_replay/process_replay.py | 11 ++-------- test/unit/test_linearizer_rewrite.py | 22 +++++++++++-------- tinygrad/engine/realize.py | 2 +- tinygrad/opt/__init__.py | 3 ++- tinygrad/opt/kernel.py | 2 +- tinygrad/uop/ops.py | 1 + 6 files changed, 20 insertions(+), 21 deletions(-) diff --git a/test/external/process_replay/process_replay.py b/test/external/process_replay/process_replay.py index c70fd208d5..d26de01163 100755 --- a/test/external/process_replay/process_replay.py +++ b/test/external/process_replay/process_replay.py @@ -4,10 +4,9 @@ import os, multiprocessing, logging, pickle, sqlite3, difflib, warnings, itertoo from typing import Callable, Any from tinygrad.helpers import VERSION, Context, ContextVar, colored, db_connection, getenv, tqdm from tinygrad.kernelize.kernelize import get_kernelize_map -from tinygrad.opt.kernel import Kernel from tinygrad.renderer import Renderer, ProgramSpec from tinygrad.engine.realize import get_program -from tinygrad.uop.ops import UOp, Ops +from tinygrad.uop.ops import UOp, Ops, KernelInfo # *** process replay settings @@ -42,13 +41,7 @@ def replay_kernelize(ret:dict[UOp, UOp], big_sink:UOp) -> tuple[str, str, tuple[ return to_str(new_sink), to_str(ret[big_sink]), (big_sink,) def replay_get_program(p:ProgramSpec, ast:UOp, renderer:Renderer) -> tuple[str, str, tuple[Any, ...]]: - # only use Kernel class if captured ast isn't already optimized - if ast.arg is None: - k2 = Kernel(ast, opts=renderer) - k2.apply_opts(p.applied_opts) - optimized_ast = k2.get_optimized_ast(name_override=p.name) - else: optimized_ast = ast - p2 = get_program(optimized_ast, renderer) + p2 = get_program(ast.replace(arg=KernelInfo(opts_to_apply=p.applied_opts, name=p.name)) if ast.arg is None else ast, renderer) def to_str(ret:ProgramSpec) -> str: return ret.src return to_str(p2), to_str(p), (p.ast, renderer, p.applied_opts) diff --git a/test/unit/test_linearizer_rewrite.py b/test/unit/test_linearizer_rewrite.py index 6db29cfd3f..f427558044 100644 --- a/test/unit/test_linearizer_rewrite.py +++ b/test/unit/test_linearizer_rewrite.py @@ -1,6 +1,8 @@ import unittest from tinygrad import Tensor, Context, Device -from tinygrad.opt.kernel import Kernel, Opt, OptOps +from tinygrad.engine.realize import get_program +from tinygrad.renderer import Opt, OptOps +from tinygrad.uop.ops import KernelInfo class TestLinearizerRewrite(unittest.TestCase): def test_reduction(self): @@ -8,20 +10,22 @@ class TestLinearizerRewrite(unittest.TestCase): out = (t*2).sum(axis=1) with Context(SPLIT_REDUCEOP=0, DEVECTORIZE=0): si = out.schedule()[-1] - k = Kernel(si.ast, Device["CPU"].renderer) - k.apply_opt(Opt(OptOps.UPCAST, 0, 4)) - k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) - prg = k.to_program() + opts_to_apply = [] + opts_to_apply.append(Opt(OptOps.UPCAST, 0, 4)) + opts_to_apply.append(Opt(OptOps.UNROLL, 0, 4)) + ast = si.ast.replace(arg=KernelInfo(opts_to_apply=tuple(opts_to_apply))) + prg = get_program(ast, Device["CPU"].renderer) print(prg.src) def test_arange(self): out = Tensor.arange(32, device="NULL") with Context(SPLIT_REDUCEOP=0, DEVECTORIZE=0): si = out.schedule()[-1] - k = Kernel(si.ast, Device["CPU"].renderer) - k.apply_opt(Opt(OptOps.UPCAST, 0, 4)) - k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) - prg = k.to_program() + opts_to_apply = [] + opts_to_apply.append(Opt(OptOps.UPCAST, 0, 4)) + opts_to_apply.append(Opt(OptOps.UNROLL, 0, 4)) + ast = si.ast.replace(arg=KernelInfo(opts_to_apply=tuple(opts_to_apply))) + prg = get_program(ast, Device["CPU"].renderer) print(prg.src) if __name__ == '__main__': diff --git a/tinygrad/engine/realize.py b/tinygrad/engine/realize.py index a42c57e90f..4473dd809e 100644 --- a/tinygrad/engine/realize.py +++ b/tinygrad/engine/realize.py @@ -27,7 +27,7 @@ def get_program(ast:UOp, renderer:Renderer) -> ProgramSpec: """ if getenv("VIZ"): graph_rewrite(ast, PatternMatcher([]), name="View Base AST") - modified_ast = get_optimized_ast(ast, renderer) if ast.arg is None else ast + modified_ast = get_optimized_ast(ast, renderer) if ast.arg is None or ast.arg.opts_to_apply is not None else ast if __debug__: type_verify(list(modified_ast.toposort())) # linearize diff --git a/tinygrad/opt/__init__.py b/tinygrad/opt/__init__.py index efe6e96027..934b9f0749 100644 --- a/tinygrad/opt/__init__.py +++ b/tinygrad/opt/__init__.py @@ -19,7 +19,8 @@ def get_optimized_ast(ast:UOp, renderer:Renderer) -> UOp: """ k = Kernel(ast, opts=renderer) - if not NOOPT: + if ast.arg is not None and ast.arg.opts_to_apply is not None: k.apply_opts(ast.arg.opts_to_apply) + elif not NOOPT: if not k.apply_tensor_cores(USE_TC.value): k.apply_opts(hand_coded_optimizations(k)) if BEAM >= 1: from tinygrad.opt.search import beam_search, bufs_from_lin diff --git a/tinygrad/opt/kernel.py b/tinygrad/opt/kernel.py index c931a89782..ce0173e6fe 100644 --- a/tinygrad/opt/kernel.py +++ b/tinygrad/opt/kernel.py @@ -454,7 +454,7 @@ class Kernel: # otherwise we just replace the VIEW source return ret.replace(src=(ret.src[0].replace(arg=st),)+ret.src[1:]) if op.op is Ops.SINK: - return ret.replace(arg = KernelInfo(self.name if name_override is None else name_override, + return ret.replace(arg = KernelInfo(ret.arg.name if ret.arg is not None else self.name if name_override is None else name_override, self.local_dims, self.upcasted, self.dont_use_locals, tuple(self.applied_opts))) if op.op is Ops.REDUCE_AXIS: reduce_idx = len(self.bufs) + self.reduceops.index(op) * 2 diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index accc7ac163..3c299d53ca 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -519,6 +519,7 @@ class KernelInfo: upcasted: int = 0 # count that are upcasted (this is remapping RANGE to UNROLL) dont_use_locals: bool = False # don't use local indexing applied_opts: tuple = tuple() + opts_to_apply: tuple|None = None @property def function_name(self): return to_function_name(self.name)