|
|
|
@@ -1,13 +1,13 @@
|
|
|
|
|
from dataclasses import replace
|
|
|
|
|
import itertools
|
|
|
|
|
import itertools, functools
|
|
|
|
|
from tinygrad.helpers import DISABLE_FAST_IDIV, TRANSCENDENTAL, SPEC, DEBUG, VIZ, IMAGE, NOOPT, EMULATED_DTYPES, NOLOCALS, USE_TC
|
|
|
|
|
from tinygrad.helpers import ALLOW_TF32, TracingKey, Context, panic
|
|
|
|
|
from tinygrad.helpers import ALLOW_TF32, TracingKey, Context, panic, all_same, flatten
|
|
|
|
|
from tinygrad.uop.ops import PatternMatcher, graph_rewrite, UOp, pm_lower_index_dtype, Ops, UPat, track_rewrites, KernelInfo, ProgramInfo, GroupOp
|
|
|
|
|
from tinygrad.uop.render import pyrender
|
|
|
|
|
from tinygrad.uop.spec import type_verify, spec_tensor, spec_program
|
|
|
|
|
from tinygrad.renderer import Renderer, Estimates
|
|
|
|
|
from tinygrad.renderer.isa import ISARenderer, IselContext, PreRegAllocContext
|
|
|
|
|
from tinygrad.dtype import dtypes, PtrDType, ImageDType
|
|
|
|
|
from tinygrad.dtype import dtypes, PtrDType, ImageDType, AddrSpace
|
|
|
|
|
|
|
|
|
|
# import all pattern matchers here
|
|
|
|
|
from tinygrad.codegen.gpudims import pm_add_gpudims
|
|
|
|
@@ -16,15 +16,15 @@ from tinygrad.codegen.decomp.dtype import pm_dtype_decomps
|
|
|
|
|
from tinygrad.codegen.decomp.op import get_late_rewrite_patterns, get_simplifying_rewrite_patterns
|
|
|
|
|
from tinygrad.codegen.decomp.transcendental import get_transcendental_patterns
|
|
|
|
|
from tinygrad.codegen.late.expander import expander, pm_pre_expander, pm_group_for_reduce
|
|
|
|
|
from tinygrad.codegen.late.devectorizer import load_store_folding, indexing_simplify, devectorize_buf_and_index, devectorize_alu, pm_reduce, \
|
|
|
|
|
ReduceContext, pm_render, pm_add_loads
|
|
|
|
|
from tinygrad.codegen.late.devectorizer import indexing_simplify, ReduceContext, pm_render, merge_reduce_ends
|
|
|
|
|
from tinygrad.codegen.opt.postrange import apply_opts
|
|
|
|
|
from tinygrad.codegen.late.gater import pm_move_gates_from_index
|
|
|
|
|
from tinygrad.codegen.simplify import pm_simplify_ranges, pm_flatten_range, pm_split_ranges, pm_load_collapse
|
|
|
|
|
from tinygrad.schedule.rangeify import pm_add_buffers_local, rangeify_codegen, pm_mops, pm_syntactic_sugar, pm_store_ranges
|
|
|
|
|
from tinygrad.schedule.rangeify import pm_mops, pm_syntactic_sugar, pm_store_ranges
|
|
|
|
|
from tinygrad.codegen.late.linearizer import CFGContext, pm_split_ends, pm_add_control_flow, linearize
|
|
|
|
|
from tinygrad.codegen.late.regalloc import LinearScanRegallocContext, pm_regalloc_rewrite
|
|
|
|
|
from tinygrad.codegen.late.coalese import memory_coalesing, pm_simplify_add_image
|
|
|
|
|
from tinygrad.uop.ops import identity_element
|
|
|
|
|
|
|
|
|
|
pm_remove_vec_dtypes = PatternMatcher([
|
|
|
|
|
# rewrite PARAM to non pointer
|
|
|
|
@@ -36,6 +36,9 @@ pm_remove_vec_dtypes = PatternMatcher([
|
|
|
|
|
lambda x: x.replace(dtype=x.dtype.base.scalar().base)),
|
|
|
|
|
# rewrite GEP to INDEX
|
|
|
|
|
(UPat(Ops.GEP, name="x"), lambda x: x.replace(op=Ops.INDEX, src=x.src+(UOp.const(dtypes.int, x.arg if len(x.arg) > 1 else x.arg[0]),), arg=None)),
|
|
|
|
|
# don't stack PARAMs/BUFFERs on INDEX. TODO: this is an expander bug
|
|
|
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.STACK, src=UPat((Ops.PARAM, Ops.BUFFER), name="b")),), allow_any_len=True, name="idx"),
|
|
|
|
|
lambda b,idx: idx.replace(src=(b,)+idx.src[1:])),
|
|
|
|
|
])+pm_clean_up_group_sink
|
|
|
|
|
|
|
|
|
|
def do_number_param(ctx:list[int], x:UOp):
|
|
|
|
@@ -51,6 +54,80 @@ pm_no_weakints = PatternMatcher([
|
|
|
|
|
(UPat(GroupOp.All, dtype=dtypes.weakint, name="x"), lambda x: x.replace(dtype=dtypes.int))
|
|
|
|
|
])
|
|
|
|
|
|
|
|
|
|
def maybe_load(u:UOp): return u.load() if u.addrspace in (AddrSpace.GLOBAL, AddrSpace.LOCAL, AddrSpace.REG) else u
|
|
|
|
|
pm_move_regs = PatternMatcher([
|
|
|
|
|
# BITCAST?
|
|
|
|
|
(UPat(GroupOp.Elementwise|{Ops.REDUCE, Ops.STACK}, name="x"), lambda x: x.replace(src=tuple([maybe_load(u) for u in x.src]))),
|
|
|
|
|
(UPat(Ops.STORE, name="x"), lambda x: x.replace(src=(x.src[0], maybe_load(x.src[1]))+x.src[2:])),
|
|
|
|
|
])
|
|
|
|
|
|
|
|
|
|
def do_devectorize(b:UOp):
|
|
|
|
|
if b.shape == (): return None
|
|
|
|
|
# broadcasting needs to be already unpacked
|
|
|
|
|
if not all_same([x.shape for x in b.src]): return None
|
|
|
|
|
src = []
|
|
|
|
|
for idx in itertools.product(*[range(x) for x in b.shape]):
|
|
|
|
|
idx_c = [UOp.const(dtypes.weakint, i) for i in idx]
|
|
|
|
|
src.append(b.replace(src=tuple([x.index(*idx_c) for x in b.src])))
|
|
|
|
|
return UOp.vectorize(*src).reshape(b.shape) if b.op is not Ops.STORE else UOp.group(*src)
|
|
|
|
|
|
|
|
|
|
devectorizer2 = pm_mops+PatternMatcher([
|
|
|
|
|
# unpack broadcasting
|
|
|
|
|
(UPat(GroupOp.Elementwise|{Ops.LOAD,Ops.STORE}, name="b"), do_devectorize),
|
|
|
|
|
# const INDEX into STACK is src
|
|
|
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.STACK, name="a"), UPat.cvar("i"))), lambda a,i: a.src[i.arg]),
|
|
|
|
|
# stacked INDEX is many INDEX
|
|
|
|
|
(UPat(Ops.INDEX, src=(UPat((Ops.PARAM, Ops.BUFFER), name="b"), UPat(Ops.STACK, name="s"))),
|
|
|
|
|
lambda b,s: UOp.vectorize(*[b.index(u) for u in s.src])),
|
|
|
|
|
# INDEX into RESHAPE moves the RESHAPE
|
|
|
|
|
(UPat(Ops.INDEX, src=(UPat((Ops.PARAM, Ops.BUFFER), name="b"), UPat(Ops.RESHAPE, name="s"))),
|
|
|
|
|
lambda b,s: b.index(s.src[0]).reshape(s.shape)),
|
|
|
|
|
# RESHAPE a void is removed (hack for AFTER)
|
|
|
|
|
(UPat(Ops.RESHAPE, dtype=dtypes.void, name="x"), lambda x: x.src[0]),
|
|
|
|
|
# reshape of a single element shaped value to scalar is an index
|
|
|
|
|
(UPat(Ops.RESHAPE, name="x"), lambda x: x.src[0].index(UOp.const(dtypes.weakint, 0)) if x.marg == () and x.src[0].shape == (1,) else None),
|
|
|
|
|
# INDEX without src is nothing
|
|
|
|
|
(UPat(Ops.INDEX, src=(UPat.var('x'),)), lambda x: x),
|
|
|
|
|
# RESHAPE+EXPAND -> STACK
|
|
|
|
|
(UPat(Ops.EXPAND, src=(UPat(Ops.RESHAPE, src=(UPat.var("x"), UPat())), UPat()), name="out"),
|
|
|
|
|
lambda x,out: UOp.vectorize(*([x]*out.max_numel())) if out.shape == (out.max_numel(),) else None),
|
|
|
|
|
# INDEX on INDEX is INDEX
|
|
|
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.INDEX, name="idx1", allow_any_len=True),), allow_any_len=True, name="idx2"),
|
|
|
|
|
lambda idx1, idx2: idx1.src[0].index(*idx1.src[1:], *idx2.src[1:])),
|
|
|
|
|
])
|
|
|
|
|
|
|
|
|
|
def reduce_ranges_to_acc(ctx:ReduceContext, r:UOp):
|
|
|
|
|
acc = UOp.placeholder_like(r, ctx.acc_num, AddrSpace.REG)
|
|
|
|
|
ctx.acc_num += 1
|
|
|
|
|
topo = r.src[0].toposort()
|
|
|
|
|
ended_ranges = flatten([x.ended_ranges for x in topo if x.op is Ops.END])
|
|
|
|
|
input_ranges = tuple(x for x in topo if x.op is Ops.RANGE and x not in r.src[1:] and x not in ended_ranges)
|
|
|
|
|
acc_init = acc.after(*input_ranges).store(identity_element(r.arg[0], r.dtype.scalar()))
|
|
|
|
|
acc_initted = acc.after(acc_init, *r.src[1:])
|
|
|
|
|
inp = r.src[0].reduce(arg=r.arg) if r.arg[1] else r.src[0]
|
|
|
|
|
acc_out = acc_initted.store(acc_initted.alu(r.arg[0], inp)).end(*r.src[1:]).rtag("mergeable")
|
|
|
|
|
return acc.after(acc_out)
|
|
|
|
|
|
|
|
|
|
def expand_horizontal_reduce(r:UOp):
|
|
|
|
|
axes = r.arg[1]
|
|
|
|
|
vals = [r.src[0].shrink(tuple((idx[axes.index(i)], idx[axes.index(i)]+1) if i in axes else None for i in range(r.src[0].ndim)))
|
|
|
|
|
for idx in itertools.product(*[range(r.src[0].max_shape[a]) for a in axes])]
|
|
|
|
|
return functools.reduce(lambda x,y: x.alu(r.arg[0], y), vals)
|
|
|
|
|
|
|
|
|
|
pm_reduce_local = PatternMatcher([
|
|
|
|
|
(UPat(Ops.REDUCE, src=(UPat(), UPat()), allow_any_len=True, name="r"), reduce_ranges_to_acc),
|
|
|
|
|
(UPat(Ops.REDUCE, src=(UPat(),), name="r"), expand_horizontal_reduce),
|
|
|
|
|
(UPat(Ops.SINK, name="sink"), merge_reduce_ends),
|
|
|
|
|
])+pm_clean_up_group_sink
|
|
|
|
|
|
|
|
|
|
def add_local_buffer(ctx, x:UOp):
|
|
|
|
|
buf = UOp.placeholder(x.max_shape, x.dtype, slot=next(ctx), addrspace=x.arg.addrspace)
|
|
|
|
|
return buf.after(buf.index(*x.src[1:]).store(x.src[0]).end(*x.src[1:]).barrier())
|
|
|
|
|
|
|
|
|
|
pm_add_local_buffers = PatternMatcher([
|
|
|
|
|
(UPat(Ops.STAGE, name="x"), add_local_buffer),
|
|
|
|
|
])+pm_mops
|
|
|
|
|
|
|
|
|
|
def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp:
|
|
|
|
|
if VIZ: graph_rewrite(ast, PatternMatcher([]), name="View Base AST")
|
|
|
|
|
if DEBUG >= 5: print(pyrender(ast))
|
|
|
|
@@ -81,13 +158,24 @@ def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp:
|
|
|
|
|
|
|
|
|
|
# expand
|
|
|
|
|
sink = graph_rewrite(sink, sym+pm_pre_expander+pm_group_for_reduce+expander, name="expander")
|
|
|
|
|
#sink = graph_rewrite(sink, gep_pushing, name="gep pushing")
|
|
|
|
|
|
|
|
|
|
# this is new style (TODO: this should all be removed)
|
|
|
|
|
sink = graph_rewrite(sink, pm_render, name="pm_render gep/stack")
|
|
|
|
|
sink = graph_rewrite(sink, pm_remove_vec_dtypes, name="transform to new style")
|
|
|
|
|
|
|
|
|
|
# ****** new style ******
|
|
|
|
|
|
|
|
|
|
# add locals
|
|
|
|
|
sink = graph_rewrite(sink, pm_add_buffers_local+rangeify_codegen, ctx=itertools.count(0), name="add local buffers")
|
|
|
|
|
sink = graph_rewrite(sink, pm_add_local_buffers, ctx=itertools.count(0), name="add local buffers")
|
|
|
|
|
sink = graph_rewrite(sink, pm_reduce_local, ctx=ReduceContext(), name="remove_reduce")
|
|
|
|
|
|
|
|
|
|
# add locals
|
|
|
|
|
#sink = graph_rewrite(sink, pm_add_buffers_local+rangeify_codegen, ctx=itertools.count(0), name="add local buffers")
|
|
|
|
|
|
|
|
|
|
# ** devectorizer (full_graph_rewrite) **
|
|
|
|
|
# remove reduce
|
|
|
|
|
sink = graph_rewrite(sink, pm_reduce+gep_pushing, ctx=ReduceContext(), name="remove_reduce")
|
|
|
|
|
#sink = graph_rewrite(sink, pm_reduce+gep_pushing, ctx=ReduceContext(), name="remove_reduce")
|
|
|
|
|
|
|
|
|
|
# add gpu dims (late). this works after devectorize, but it's faster here
|
|
|
|
|
sink = graph_rewrite(sink, pm_add_gpudims, ctx=ren, name="add gpudims")
|
|
|
|
@@ -95,14 +183,16 @@ def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp:
|
|
|
|
|
# **** optimizations are done, now we lower to actual code ****
|
|
|
|
|
|
|
|
|
|
# add loads and remove invalids
|
|
|
|
|
sink = graph_rewrite(sink, pm_add_loads+pm_remove_invalid, name="** add loads (code)")
|
|
|
|
|
#sink = graph_rewrite(sink, pm_add_loads+pm_remove_invalid, name="** add loads (code)")
|
|
|
|
|
|
|
|
|
|
# devectorize
|
|
|
|
|
sink = graph_rewrite(sink, sym+devectorize_alu+devectorize_buf_and_index+load_store_folding, ctx=ren, name="devectorize")
|
|
|
|
|
#sink = graph_rewrite(sink, sym+devectorize_alu+devectorize_buf_and_index+load_store_folding, ctx=ren, name="devectorize")
|
|
|
|
|
|
|
|
|
|
# this is new style (TODO: this should all be removed)
|
|
|
|
|
sink = graph_rewrite(sink, pm_render, name="pm_render gep/stack")
|
|
|
|
|
sink = graph_rewrite(sink, pm_remove_vec_dtypes, name="transform to new style")
|
|
|
|
|
# add loads
|
|
|
|
|
sink = graph_rewrite(sink, pm_move_regs, name="** add loads")
|
|
|
|
|
|
|
|
|
|
# devectorize
|
|
|
|
|
sink = graph_rewrite(sink, symbolic_simple+devectorizer2, ctx=ren, name="devectorize2")
|
|
|
|
|
|
|
|
|
|
# simplify indexing
|
|
|
|
|
sink = graph_rewrite(sink, indexing_simplify, name="simplify load/store indexing")
|
|
|
|
|