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
This commit is contained in:
qazal
2025-06-24 18:55:39 +03:00
committed by GitHub
parent 18e264a449
commit de4b9bf53b
6 changed files with 20 additions and 21 deletions
+2 -9
View File
@@ -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)
+13 -9
View File
@@ -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__':
+1 -1
View File
@@ -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
+2 -1
View File
@@ -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
+1 -1
View File
@@ -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
+1
View File
@@ -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)