mirror of
https://github.com/tinygrad/tinygrad.git
synced 2026-08-29 13:16:09 +00:00
simplify_merge_adjacent
This commit is contained in:
@@ -18,7 +18,7 @@ from tinygrad.codegen.late.devectorizer import load_store_folding, load_store_in
|
||||
from tinygrad.codegen.late.linearize import block_create, pm_blockend_merge, block_merge, pm_finalize, BlockContext
|
||||
from tinygrad.codegen.opt 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_early, pm_postrange_opt
|
||||
from tinygrad.codegen.opt.postrange import pm_postrange_opt_early, pm_postrange_opt, pm_postrange_opt_merge
|
||||
from tinygrad.schedule.rangeify import pm_add_buffers_local, rangeify_codegen
|
||||
|
||||
@dataclass
|
||||
@@ -68,6 +68,8 @@ def _get_rewrites_for_renderer(opts:Renderer, linearizer:bool, _QUANTIZE, _DEVEC
|
||||
ret.append(RewriteStep(sym+migrate_indexing, name="initial symbolic"))
|
||||
|
||||
if _POSTOPT or _RANGEIFY:
|
||||
ret.append(RewriteStep(pm_postrange_opt_merge, ctx=lambda _: ({}, opts), name="early range merge"))
|
||||
ret.append(RewriteStep(sym, name="mid symbolic"))
|
||||
ret.append(RewriteStep(pm_postrange_opt_early, ctx=lambda _: ({}, opts), name="early post opt ast"))
|
||||
ret.append(RewriteStep(sym, name="mid symbolic"))
|
||||
ret.append(RewriteStep(pm_postrange_opt, ctx=lambda _: opts, name="post optimize ast"))
|
||||
|
||||
@@ -1,13 +1,65 @@
|
||||
import math
|
||||
from dataclasses import replace
|
||||
from tinygrad.dtype import dtypes
|
||||
from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, KernelInfo
|
||||
from tinygrad.uop.symbolic import sym
|
||||
from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, KernelInfo, graph_rewrite, _substitute
|
||||
from tinygrad.uop.symbolic import symbolic
|
||||
from tinygrad.helpers import colored, USE_TC, DEBUG
|
||||
from tinygrad.codegen.opt.kernel import axis_colors, AxisType
|
||||
from tinygrad.renderer import Renderer
|
||||
from tinygrad.codegen.opt.tc import TensorCore
|
||||
|
||||
def count_divmod(x:UOp): return len([u for u in x.toposort() if u.op in {Ops.IDIV, Ops.MOD}])
|
||||
|
||||
def flatten_range_in_terminators(r:UOp):
|
||||
off = 2 if r.op is Ops.STORE else 1
|
||||
rngs = r.src[off:]
|
||||
if not len(rngs): return None
|
||||
new_rngs = [x for x in UOp.sink(*rngs).toposort() if x.op is Ops.RANGE]
|
||||
return r.replace(src=r.src[:off]+tuple(new_rngs))
|
||||
|
||||
pm_flatten_range = PatternMatcher([
|
||||
# flatten ranges
|
||||
(UPat((Ops.REDUCE, Ops.STORE), name="r"), flatten_range_in_terminators),
|
||||
])
|
||||
|
||||
# NOTE: this one is better than the one in kernel.py
|
||||
def simplify_merge_adjacent(ast:UOp):
|
||||
# get all ranges (sorted)
|
||||
rng = sorted([u for u in ast.parents if u.op is Ops.RANGE], key=lambda x: x.arg[0:-1])
|
||||
terminators = [u for u in ast.parents if u.op in {Ops.REDUCE, Ops.STORE}]
|
||||
termination = {}
|
||||
for t in terminators:
|
||||
for u in t.src[1 if t.op is Ops.REDUCE else 2:]: termination[u] = t
|
||||
|
||||
replaces = {}
|
||||
i = 0
|
||||
while i < len(rng)-1:
|
||||
r0, r1 = rng[i], rng[i+1]
|
||||
# same axistype and same termination
|
||||
if r0.arg[1] == r1.arg[1] and termination[r0] == termination[r1]:
|
||||
s0, s1 = r0.src[0], r1.src[0]
|
||||
new_range = r0.replace(src=(s0*s1,)).simplify()
|
||||
# this checks the legality of a merge
|
||||
oidx = ast.simplify()
|
||||
nidx = graph_rewrite(oidx, _substitute+symbolic+pm_flatten_range, ctx={r0:new_range//s1, r1:new_range%s1}, name=f"check_merge_{i}_{i+1}")
|
||||
# it simplifies
|
||||
if count_divmod(nidx) <= count_divmod(oidx):
|
||||
# it is correct
|
||||
midx = graph_rewrite(nidx, _substitute+symbolic+pm_flatten_range, ctx={new_range:r0*s1+r1}, name=f"correct_merge_{i}_{i+1}")
|
||||
if oidx is midx:
|
||||
termination[new_range] = termination[r0]
|
||||
replaces[r0] = new_range//s1
|
||||
replaces[r1] = new_range%s1
|
||||
rng[i] = new_range
|
||||
del rng[i+1]
|
||||
continue
|
||||
i += 1
|
||||
return ast.substitute(replaces, name="simplify_merge_adjacent")
|
||||
|
||||
pm_postrange_opt_merge = PatternMatcher([
|
||||
(UPat(Ops.SINK, name="ast"), simplify_merge_adjacent),
|
||||
])
|
||||
|
||||
def apply_tensor_cores(ctx:tuple[dict, Renderer], in0:UOp, in1:UOp, r_range:UOp, reduceop:UOp):
|
||||
if not USE_TC: return None
|
||||
# tensor cores have three ranges. X, Y, and REDUCE
|
||||
@@ -89,7 +141,7 @@ def early_sink(ctx:tuple[dict, Renderer], s:UOp):
|
||||
|
||||
pm_postrange_opt_early = PatternMatcher([
|
||||
# TODO: this is optional (and can have internal options) and we need a way to express that
|
||||
((UPat.var("in0")*UPat.var("in1")).reduce(UPat(Ops.RANGE, name="r_range"), name="reduceop", arg=Ops.ADD), apply_tensor_cores),
|
||||
#((UPat.var("in0")*UPat.var("in1")).reduce(UPat(Ops.RANGE, name="r_range"), name="reduceop", arg=Ops.ADD), apply_tensor_cores),
|
||||
(UPat(Ops.SINK, name="s"), early_sink),
|
||||
])
|
||||
|
||||
@@ -103,7 +155,7 @@ def split_range(r:UOp):
|
||||
if r.arg[-1] not in {AxisType.LOOP, AxisType.GLOBAL, AxisType.REDUCE}: return None
|
||||
if r.tag is not None: return None
|
||||
# any divisor is an option
|
||||
is_local = False # if r.arg[-1] is AxisType.REDUCE else True
|
||||
is_local = False if r.arg[-1] is AxisType.REDUCE else True
|
||||
N = 4
|
||||
rd = r.src[0].divides(N)
|
||||
if rd is None: return None
|
||||
@@ -111,13 +163,6 @@ def split_range(r:UOp):
|
||||
er = UOp(Ops.RANGE, dtypes.int, src=(UOp.const(dtypes.int, N),), arg=r.arg[0:-1]+(1, axis_typemap[(r.arg[-1] is AxisType.REDUCE, is_local)]))
|
||||
return sr*N+er
|
||||
|
||||
def flatten_range_in_terminators(r:UOp):
|
||||
off = 2 if r.op is Ops.STORE else 1
|
||||
rngs = r.src[off:]
|
||||
if not len(rngs): return None
|
||||
new_rngs = [x for x in UOp.sink(*rngs).toposort() if x.op is Ops.RANGE]
|
||||
return r.replace(src=r.src[:off]+tuple(new_rngs))
|
||||
|
||||
def rename_sink(s:UOp):
|
||||
if s.arg is not None and s.arg.name != "test": return None
|
||||
|
||||
@@ -128,13 +173,11 @@ def rename_sink(s:UOp):
|
||||
name = "k" + colored('_', 'BLACK').join(['']+[colored(x.src[0].render(), axis_colors[x.arg[-1]]) for x in rngs])
|
||||
return s.replace(arg=KernelInfo(name=name) if s.arg is None else replace(s.arg, name=name))
|
||||
|
||||
pm_postrange_opt = PatternMatcher([
|
||||
pm_postrange_opt = pm_flatten_range+PatternMatcher([
|
||||
# TODO: this is optional (and can have internal options) and we need a way to express that
|
||||
(UPat(Ops.RANGE, name="r"), split_range),
|
||||
# remove axes with 1
|
||||
(UPat(Ops.RANGE, name="r"), lambda r: r.const_like(0) if r.vmax == 0 else None),
|
||||
# flatten ranges
|
||||
(UPat((Ops.REDUCE, Ops.STORE), name="r"), flatten_range_in_terminators),
|
||||
# run this last
|
||||
(UPat(Ops.SINK, name="s"), rename_sink),
|
||||
])
|
||||
|
||||
Reference in New Issue
Block a user