diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index 0783c7bc76..e4803b70ff 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -1,7 +1,7 @@ from typing import Any, Callable import functools from dataclasses import dataclass -from tinygrad.helpers import QUANTIZE, DEVECTORIZE, TRANSCENDENTAL, RANGEIFY, POSTOPT +from tinygrad.helpers import QUANTIZE, DEVECTORIZE, TRANSCENDENTAL from tinygrad.uop.ops import PatternMatcher, graph_rewrite, UOp from tinygrad.uop.spec import type_verify from tinygrad.renderer import Renderer @@ -16,7 +16,6 @@ from tinygrad.codegen.late.expander import migrate_indexing, expander, pm_pre_ex from tinygrad.codegen.late.devectorizer import load_store_folding, load_store_indexing, devectorize, pm_reduce, \ ReduceContext, correct_load_store, pm_render from tinygrad.codegen.late.linearize import block_create, pm_blockend_merge, block_merge, pm_finalize, BlockContext -from tinygrad.codegen.opt.kernel 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 from tinygrad.schedule.rangeify import pm_add_buffers_local, rangeify_codegen @@ -46,24 +45,22 @@ rewrites_for_linearizer = [ def get_rewrites_for_renderer(opts:Renderer, 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, RANGEIFY.value, POSTOPT.value) + return _get_rewrites_for_renderer(opts, linearizer, QUANTIZE.value, DEVECTORIZE.value, TRANSCENDENTAL.value) @functools.cache -def _get_rewrites_for_renderer(opts:Renderer, linearizer:bool, _QUANTIZE, _DEVECTORIZE, _TRANSCENDENTAL, _RANGEIFY, _POSTOPT) -> list[RewriteStep]: +def _get_rewrites_for_renderer(opts:Renderer, linearizer:bool, _QUANTIZE, _DEVECTORIZE, _TRANSCENDENTAL) -> list[RewriteStep]: # ** lowerer (rewrite_shapetracker_with_index) ** ret: list[RewriteStep] = [] # view pushing ret.extend(rewrites_for_views) - # this is kernel.py - if _POSTOPT <= 1 and not _RANGEIFY: ret.append(RewriteStep(pm_get_optimization, ctx=lambda _: opts, name="get optimization")) - if not _POSTOPT and not _RANGEIFY: ret.append(RewriteStep(pm_do_optimize, ctx=lambda _: opts, name="optimize ast")) - + # 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)) - if _POSTOPT or _RANGEIFY: ret.append(RewriteStep(pm_postrange_opt, ctx=lambda _: opts, name="post optimize ast")) + # optimize (schedule) the AST + ret.append(RewriteStep(pm_postrange_opt, ctx=lambda _: opts, name="post optimize ast")) # ** expander (expand_rewrite) ** ret.append(RewriteStep(sym+migrate_indexing, name="initial symbolic")) diff --git a/tinygrad/codegen/opt/kernel.py b/tinygrad/codegen/opt/kernel.py index 02f522dce4..832a8e5fa6 100644 --- a/tinygrad/codegen/opt/kernel.py +++ b/tinygrad/codegen/opt/kernel.py @@ -4,13 +4,13 @@ from dataclasses import dataclass from collections import defaultdict from typing import cast, Final, Callable, Sequence from tinygrad.codegen.opt import OptOps, Opt, KernelOptError, check, axis_letters, axis_colors -from tinygrad.uop.ops import GroupOp, KernelInfo, UOp, Ops, can_pad, resolve, Variable, sint, graph_rewrite, AxisType, PatternMatcher, UPat +from tinygrad.uop.ops import GroupOp, KernelInfo, UOp, Ops, can_pad, resolve, Variable, sint, graph_rewrite, AxisType from tinygrad.uop.spec import type_verify, ast_spec from tinygrad.device import Device from tinygrad.codegen.opt.tc import TensorCore from tinygrad.renderer import Renderer from tinygrad.dtype import ImageDType -from tinygrad.helpers import all_same, colored, ansilen, dedup, prod, round_up, to_function_name, unwrap, argfix, DEBUG, NOOPT, BEAM, getenv, POSTOPT +from tinygrad.helpers import all_same, colored, ansilen, dedup, prod, round_up, to_function_name, unwrap, argfix, DEBUG from tinygrad.shape.shapetracker import ShapeTracker from tinygrad.shape.view import strides_for_shape, get_contraction from tinygrad.codegen.opt.swizzler import view_left, view_left_through_load @@ -433,47 +433,3 @@ class Kernel: fixed_ast = fixup_ast(self.ast) del fixup_ast return graph_rewrite(fixed_ast, view_left+view_left_through_load, name="fixup optimized AST") - -def get_optimized_ast(ast:UOp, renderer:Renderer) -> UOp|None: - """ - Optimize an AST based on heuristics or BEAM search. - - Args: - ast: The Ops.SINK rooted AST - renderer: The renderer used to generate the code - - Returns: - The Ops.SINK rooted AST transformed to apply the opts and with a KernelInfo in the arg. - """ - - # no shape, no opt - if ast.src[0].st is None: return None - new_arg = ast.arg - if new_arg is None: - k = Kernel(ast, opts=renderer) - if not NOOPT: - from tinygrad.codegen.opt.heuristic import hand_coded_optimizations - k.apply_opts(hand_coded_optimizations(k)) - if not POSTOPT and BEAM >= 1: - from tinygrad.codegen.opt.search import beam_search, bufs_from_lin - kb = Kernel(ast, opts=renderer) - rawbufs = bufs_from_lin(kb, allocate=False) - k = beam_search(kb, rawbufs, BEAM.value, bool(getenv("BEAM_ESTIMATE", 1))) - new_arg = KernelInfo(opts_to_apply=tuple(k.applied_opts)) - elif len(new_arg.applied_opts): return None - return Kernel(ast.replace(arg=None), opts=renderer).get_optimized_ast().replace(arg=new_arg) - -pm_get_optimization = PatternMatcher([ - (UPat(Ops.SINK, name="ast"), lambda ctx,ast: get_optimized_ast(ast, ctx)), -]) - -def apply_opt(ast:UOp, renderer:Renderer): - k = Kernel(ast, opts=renderer) - k.apply_opts(ast.arg.opts_to_apply) - ret = k.get_optimized_ast() - if __debug__: type_verify(list(ret.toposort())) - return ret - -pm_do_optimize = PatternMatcher([ - (UPat(Ops.SINK, name="ast"), lambda ctx,ast: apply_opt(ast, ctx) if ast.arg is not None and ast.arg.opts_to_apply is not None else None), -]) \ No newline at end of file diff --git a/tinygrad/codegen/opt/postrange.py b/tinygrad/codegen/opt/postrange.py index f143e5beba..0df88ffef3 100644 --- a/tinygrad/codegen/opt/postrange.py +++ b/tinygrad/codegen/opt/postrange.py @@ -6,7 +6,7 @@ from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, KernelInfo, graph_r from tinygrad.uop.symbolic import symbolic_flat from tinygrad.device import Buffer from tinygrad.dtype import AddrSpace, dtypes, ImageDType -from tinygrad.helpers import colored, BEAM, getenv, DEBUG, to_function_name, NOOPT, argsort, round_up, POSTOPT, prod +from tinygrad.helpers import colored, BEAM, getenv, DEBUG, to_function_name, NOOPT, argsort, round_up, prod from tinygrad.codegen.opt import axis_colors, Opt, OptOps, KernelOptError, check, axis_letters from tinygrad.renderer import Renderer from tinygrad.schedule.rangeify import remove_tags @@ -347,16 +347,14 @@ def apply_opts(ctx:Renderer, ast:UOp): if ast.tag is not None: return None k = Scheduler(ast, ctx) k.convert_loop_to_global() + k.simplify_merge_adjacent() if BEAM >= 1: - k.simplify_merge_adjacent() from tinygrad.codegen.opt.search import beam_search rawbufs = bufs_from_ast(ast, ctx.device) k = beam_search(k, rawbufs, BEAM.value, bool(getenv("BEAM_ESTIMATE", 1))) elif ast.arg is not None and ast.arg.opts_to_apply is not None: - if POSTOPT >= 2: k.simplify_merge_adjacent() for opt in ast.arg.opts_to_apply: k.apply_opt(opt) elif not NOOPT and (ast.arg is None or ast.arg.applied_opts == ()): - k.simplify_merge_adjacent() from tinygrad.codegen.opt.heuristic import hand_coded_optimizations # NOTE: hand_coded_optimizations doesn't support multiblock opts yet if all(len(u.src) == 1 for u in ast.parents if u.op is Ops.LOAD): diff --git a/tinygrad/helpers.py b/tinygrad/helpers.py index 088b42ef1e..37467c5eea 100644 --- a/tinygrad/helpers.py +++ b/tinygrad/helpers.py @@ -140,7 +140,7 @@ DONT_REALIZE_EXPAND, DONT_GROUP_REDUCES = ContextVar("DONT_REALIZE_EXPAND", 0), QUANTIZE, VALIDATE_WITH_CPU, DISABLE_FAST_IDIV = ContextVar("QUANTIZE", 0), ContextVar("VALIDATE_WITH_CPU", 0), ContextVar("DISABLE_FAST_IDIV", 0) CORRECT_DIVMOD_FOLDING, FUSE_OPTIM = ContextVar("CORRECT_DIVMOD_FOLDING", 0), ContextVar("FUSE_OPTIM", 0) ALLOW_DEVICE_USAGE, MAX_BUFFER_SIZE, AMD_LLVM = ContextVar("ALLOW_DEVICE_USAGE", 1), ContextVar("MAX_BUFFER_SIZE", 0), ContextVar("AMD_LLVM", 1) -RANGEIFY, POSTOPT, FUSE_ATTENTION = ContextVar("RANGEIFY", 0), ContextVar("POSTOPT", 2), ContextVar("FUSE_ATTENTION", 0) +RANGEIFY, FUSE_ATTENTION = ContextVar("RANGEIFY", 0), ContextVar("FUSE_ATTENTION", 0) EMULATE = ContextVar("EMULATE", "") @dataclass(frozen=True)