forked from tinygrad/tinygrad
376 lines
16 KiB
Python
376 lines
16 KiB
Python
from typing import Any
|
|
from dataclasses import dataclass, field
|
|
from tinygrad.dtype import dtypes, AddrSpace, PtrDType
|
|
from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, resolve, GroupOp, RewriteNotReady
|
|
from tinygrad.helpers import argsort, prod, all_same, pluralize, getenv
|
|
|
|
from tinygrad.schedule.kernelize import Kernel
|
|
from tinygrad.uop.ops import track_rewrites, graph_rewrite_map, graph_rewrite, KernelInfo, identity_element
|
|
|
|
imported_rewrites = PatternMatcher([
|
|
# UOp with size 0 is zero
|
|
(UPat(GroupOp.All-{Ops.SINK}, name="root"), lambda root: root.const_like(0) if root.base.st is not None and root.size == 0 else None),
|
|
# DETACH and CONTIGUOUS_BACKWARD are NOOPs here
|
|
(UPat((Ops.DETACH, Ops.CONTIGUOUS_BACKWARD), name="x"), lambda x: x.src[0]),
|
|
# reduce of size 0 is the identity element
|
|
(UPat(Ops.REDUCE_AXIS, name="reduce", src=(UPat.var("x"),)),
|
|
lambda reduce,x: reduce.const_like(identity_element(reduce.arg[0], reduce.dtype)) if x.size == 0 and reduce.size != 0 else None),
|
|
])
|
|
|
|
earliest_rewrites = imported_rewrites+PatternMatcher([
|
|
# RESHAPE on RESHAPE is the second reshape
|
|
(UPat(Ops.RESHAPE, src=(UPat(Ops.RESHAPE),), name="x"), lambda x: x.replace(src=(x.src[0].src[0],))),
|
|
# non shape changing RESHAPE is NOOP
|
|
(UPat(Ops.RESHAPE, name="x"), lambda x: x.src[0] if x.src[0].shape == x.arg else None),
|
|
# RESHAPE after COPY
|
|
(UPat(Ops.COPY, src=(UPat(Ops.RESHAPE, name="r"),UPat(name="d")), name="c"), lambda c,r,d: c.replace(src=(r.src[0],d)).reshape(r.arg)),
|
|
# TODO: this should be BUFFER_VIEW
|
|
(UPat(Ops.COPY, src=(UPat(Ops.SHRINK, name="r"),UPat(name="d")), name="c"), lambda c,r,d: c.replace(src=(r.src[0],d)).shrink(r.arg)),
|
|
# const hacks
|
|
(UPat(Ops.CONST, name="x"), lambda x:
|
|
x.replace(src=(x.src[0].src[0],)).reshape((1,)*len(x.shape)).expand(x.shape) if \
|
|
len(x.src) and x.src[0].op is Ops.VIEW and not any(s == 0 for s in x.shape) else None),
|
|
])
|
|
|
|
# 1. add contiguous where we have to
|
|
|
|
ALWAYS_CONTIGUOUS: set[Ops] = {Ops.CONTIGUOUS, Ops.ASSIGN, Ops.COPY, Ops.BUFFER, Ops.BUFFER_VIEW,
|
|
Ops.CONST, Ops.BIND, Ops.DEVICE, Ops.MSELECT, Ops.MSTACK, Ops.DEFINE_GLOBAL,
|
|
Ops.DEFINE_LOCAL, Ops.DEFINE_REG, Ops.LOAD}
|
|
|
|
def realize(ctx:dict[UOp, None], tr:UOp) -> None: ctx[tr] = None
|
|
|
|
def realize_parents(ctx:dict[UOp, None], rb:UOp) -> None:
|
|
for s in rb.src:
|
|
if s.op not in ALWAYS_CONTIGUOUS: ctx[s] = None
|
|
|
|
do_realize = PatternMatcher([
|
|
# always realize SINK parents
|
|
(UPat(Ops.SINK, name="s"), lambda ctx,s: ctx.update((x.base, None) for x in s.src if x.base.op not in ALWAYS_CONTIGUOUS)),
|
|
# always realize ASSIGN/CONTIGUOUS/COPY/BUFFER_VIEW
|
|
(UPat({Ops.ASSIGN, Ops.CONTIGUOUS, Ops.COPY, Ops.BUFFER_VIEW}, name="tr"), realize),
|
|
# realize parents of COPY, MSELECT, MSTACK
|
|
(UPat((Ops.COPY, Ops.MSELECT, Ops.MSTACK), name="rb"), realize_parents),
|
|
])
|
|
|
|
add_contiguous = PatternMatcher([(UPat(GroupOp.All-{Ops.CONTIGUOUS, Ops.ASSIGN}, name="x"),
|
|
lambda ctx,x: x.replace(tag=1).contiguous() if x in ctx and x.tag is None else None)])
|
|
remove_tags = PatternMatcher([(UPat(GroupOp.All, name="x"), lambda x: x.replace(tag=None) if x.tag is not None else None)])
|
|
early_cleanups = PatternMatcher([(UPat().contiguous(name="c").contiguous(), lambda c: c),])
|
|
|
|
# 2. mark all children
|
|
|
|
@dataclass
|
|
class ChildrenContext: children: dict[UOp, list[UOp]]|None = None
|
|
def extract_children(ctx:ChildrenContext, x:UOp):
|
|
if ctx.children is not None: return
|
|
# REDUCE_AXIS is fine here, should go to contig only (gate)
|
|
ctx.children = {k:list(v.keys()) for k,v in x.get_children_map().items() if len(v) > 1 and any(x.op is Ops.REDUCE_AXIS for x in k.toposort())}
|
|
def mark_children(ctx:ChildrenContext, x:UOp):
|
|
new_srcs = [(UOp(Ops.CHILD, s.dtype, src=(s,), arg=(ctx.children[s].index(x), len(ctx.children[s]))) if s in ctx.children else s) for s in x.src]
|
|
return x.replace(src=tuple(new_srcs))
|
|
pm_children = PatternMatcher([
|
|
(UPat(Ops.SINK, name="x"), extract_children),
|
|
(UPat(GroupOp.All-{Ops.CHILD}, name="x"), mark_children),
|
|
])
|
|
|
|
# 3. rangeify
|
|
|
|
@dataclass
|
|
class RangeifyContext:
|
|
idx: int = 0
|
|
regs: int = 0
|
|
seen_children: dict[UOp, dict[int, UOp]] = field(default_factory=dict)
|
|
seen_child: dict[UOp, Any] = field(default_factory=dict)
|
|
progress_children: dict[UOp, int] = field(default_factory=dict)
|
|
|
|
def map_reshape(idx:UOp, r:UOp):
|
|
acc = 1
|
|
to_sum = []
|
|
for s,src in list(zip(idx.shape, idx.src[1:]))[::-1]:
|
|
to_sum.append(acc*src)
|
|
acc *= s
|
|
mish = sum(to_sum)
|
|
ret = []
|
|
for s in r.src[0].shape[::-1]:
|
|
if resolve(s!=1):
|
|
# this MOD should limit any ranges outside s
|
|
ret.append(mish % s)
|
|
mish //= s
|
|
else:
|
|
ret.append(UOp.const(dtypes.int, 0))
|
|
ret = UOp.sink(*ret).simplify().src[::-1] if len(ret) else ()
|
|
return r.src[0].index(*ret, dtype=idx.dtype, arg=idx.arg)
|
|
|
|
def map_pad(idx:UOp, r:UOp):
|
|
ret = list(idx.src[1:])
|
|
bigwhere = UOp.const(dtypes.bool, True)
|
|
for i,(sh,(s,e)) in enumerate(zip(r.shape, r.arg)):
|
|
if s == 0 and e == 0: continue
|
|
where = UOp.const(dtypes.bool, True)
|
|
if e > 0: where = where & (ret[i] < (sh-e))
|
|
if s > 0: where = where & (ret[i] >= s)
|
|
bigwhere = bigwhere & where
|
|
# this is safe but dumb
|
|
ret[i] = (ret[i] - s).maximum(0).minimum(r.src[0].shape[i]-1)
|
|
# PAD is with 0
|
|
return bigwhere.simplify().where(r.src[0].index(*ret, dtype=idx.dtype, arg=idx.arg), UOp.const(r.dtype, 0))
|
|
|
|
def map_expand(r:UOp, idx:UOp):
|
|
new_rngs = []
|
|
ending_ranges = []
|
|
non_ending_ranges = []
|
|
for a,x,y in zip(idx.src[1:], r.src[0].shape, r.shape):
|
|
axis_to_range = [u for u in a.toposort() if u.op is Ops.RANGE]
|
|
if resolve(x!=y, False):
|
|
ending_ranges.extend(axis_to_range)
|
|
new_rngs.append(a.const_like(0))
|
|
else:
|
|
non_ending_ranges.extend(axis_to_range)
|
|
new_rngs.append(a)
|
|
ending_ranges = [x.arg for x in ending_ranges if x not in non_ending_ranges]
|
|
if idx.arg is not None: ending_ranges.append(idx.arg)
|
|
return r.src[0].index(*new_rngs, arg=min([x for x in ending_ranges]) if ending_ranges else None)
|
|
|
|
pm_mops = PatternMatcher([
|
|
# this is like the definitions of these
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.SHRINK, name="r"),), allow_any_len=True, name="idx"),
|
|
lambda r,idx: r.src[0].index(*[a+ss if resolve(ss != 0) else a for a,(ss,_) in zip(idx.src[1:], r.arg)], dtype=idx.dtype, arg=idx.arg)),
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.PERMUTE, name="r"),), allow_any_len=True, name="idx"),
|
|
lambda r,idx: r.src[0].index(*[idx.src[1+p] for p in argsort(idx.src[0].arg)], dtype=idx.dtype, arg=idx.arg)),
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.FLIP, name="r"),), allow_any_len=True, name="idx"),
|
|
lambda r,idx: r.src[0].index(*[((s-1)-a) if f else a for a,s,f in zip(idx.src[1:], r.shape, r.arg)], dtype=idx.dtype, arg=idx.arg)),
|
|
# expand needs to end ranges
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.EXPAND, name="r"),), allow_any_len=True, name="idx"), map_expand),
|
|
# reshape does a lot of symbolic stuff
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.RESHAPE, name="r"),), allow_any_len=True, name="idx"), map_reshape),
|
|
# pad adds min and max
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.PAD, name="r"),), allow_any_len=True, name="idx"), map_pad),
|
|
])
|
|
|
|
def map_contiguous(ctx:RangeifyContext, x:UOp, idx:UOp|None=None):
|
|
if x.arg is None and idx is not None: return None
|
|
if x.arg is not None and idx is None: return None
|
|
ranges = []
|
|
new_ranges = []
|
|
passthrough_idx = []
|
|
for i,s in enumerate(x.shape):
|
|
if x.arg is not None and i not in x.arg:
|
|
assert idx is not None, "partial contig requires index"
|
|
ranges.append(idx.src[1+i])
|
|
continue
|
|
if idx is not None: passthrough_idx.append(idx.src[1+i])
|
|
if resolve(s!=1):
|
|
ranges.append(UOp.range(dtypes.int, s, ctx.idx))
|
|
new_ranges.append(ranges[-1])
|
|
ctx.idx += 1
|
|
else:
|
|
ranges.append(UOp.const(dtypes.int, 0))
|
|
ret = x.src[0].index(*ranges).bufferize(*new_ranges, arg=x.device)
|
|
ret = ret.index(*passthrough_idx) if len(passthrough_idx) else ret.reshape(x.shape)
|
|
return ret
|
|
|
|
def map_reduce(ctx:RangeifyContext, idx:UOp, red:UOp):
|
|
rngs = list(idx.src[1:])
|
|
new_ranges = []
|
|
for i,s in enumerate(red.src[0].shape):
|
|
if i in red.arg[1]:
|
|
rngs[i] = UOp.range(dtypes.int, s, ctx.idx)
|
|
ctx.idx += 1
|
|
new_ranges.append(rngs[i])
|
|
return UOp(Ops.REDUCE, red.dtype, src=(red.src[0].index(*rngs),)+tuple(new_ranges), arg=red.arg[0])
|
|
|
|
def index_child(ctx:RangeifyContext, c:UOp, x:UOp, idx:UOp):
|
|
if c not in ctx.seen_children:
|
|
ctx.seen_children[c] = {}
|
|
ctx.progress_children[c] = 0
|
|
ctx.seen_children[c][x.arg[0]] = idx
|
|
# wait here until we have seen all the children
|
|
if len(ctx.seen_children[c]) != x.arg[1]:
|
|
if ctx.progress_children[c] == len(ctx.seen_children[c]): raise RuntimeError("revisited child before visiting all children")
|
|
ctx.progress_children[c] = len(ctx.seen_children[c])
|
|
raise RewriteNotReady
|
|
|
|
if c not in ctx.seen_child:
|
|
all_rngs = zip(*[ch.src[1:] for ch in ctx.seen_children[c].values()])
|
|
out_rngs = []
|
|
end_ranges = []
|
|
idx_ranges = []
|
|
for i,r in enumerate(all_rngs):
|
|
if all_same(r):
|
|
out_rngs.append(r[0])
|
|
else:
|
|
out_rngs.append(UOp.range(dtypes.int, c.shape[i], ctx.idx))
|
|
ctx.idx += 1
|
|
end_ranges.append(out_rngs[-1])
|
|
idx_ranges.append(i)
|
|
ctx.seen_child[c] = (idx_ranges, end_ranges)
|
|
else:
|
|
out_rngs = list(idx.src[1:])
|
|
idx_ranges, end_ranges = ctx.seen_child[c]
|
|
for i,nr in zip(idx_ranges, end_ranges): out_rngs[i] = nr
|
|
if len(idx_ranges) == 0: return c.index(*out_rngs)
|
|
return c.index(*out_rngs).bufferize(*end_ranges, arg=x.device).index(*[idx.src[1+i] for i in idx_ranges])
|
|
|
|
def might_end_axis(idx:UOp):
|
|
if idx.arg is None: return None
|
|
to_end_axis = []
|
|
for i,a in enumerate(idx.src[1:]):
|
|
if any(x.arg > idx.arg for x in a.toposort() if x.op is Ops.RANGE):
|
|
to_end_axis.append(i)
|
|
if to_end_axis: return idx.replace(src=(idx.src[0].contiguous(arg=tuple(to_end_axis)),)+idx.src[1:], arg=None)
|
|
return idx.replace(arg=None)
|
|
|
|
pm_rangeify = pm_mops+PatternMatcher([
|
|
# if there are new ended children, tag the SINK
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.CHILD, src=(UPat(name="c"), ), name="x"),), allow_any_len=True, name="idx"), index_child),
|
|
|
|
# if there's an INDEX it can support partial contig
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.CONTIGUOUS, src=(UPat(),), name="x"),), allow_any_len=True, name="idx"), map_contiguous),
|
|
|
|
# CONST can't have axes. remove srcs when we idx
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.CONST, name="c"),)), lambda c: c.replace(src=())),
|
|
|
|
# sink contigs to kick it off
|
|
(UPat(Ops.CONTIGUOUS, src=(UPat(),), name="x"), lambda ctx,x: map_contiguous(ctx, x)),
|
|
|
|
# handle arg on any op with weight. old endrange stuff
|
|
(UPat(Ops.INDEX,src=(UPat(GroupOp.Elementwise.union({Ops.REDUCE_AXIS})),), allow_any_len=True, name="idx"), might_end_axis),
|
|
|
|
# move MAP through elementwise ALU / reduce. these are the items with cost
|
|
(UPat(Ops.INDEX, src=(UPat(GroupOp.Elementwise.union({Ops.STORE, Ops.ASSIGN, Ops.COPY, Ops.DEVICE})),), allow_any_len=True, name="x"),
|
|
lambda x: x.src[0].replace(src=tuple([s.index(*x.src[1:]) for s in x.src[0].src]))),
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.REDUCE_AXIS, name="red"),), allow_any_len=True, name="idx"), map_reduce),
|
|
|
|
# CONTIGUOUS on ASSIGN is STORE
|
|
# TODO: tag in UPat?
|
|
(UPat(Ops.CONTIGUOUS, src=(UPat(Ops.ASSIGN, name="a"),), name="c", allow_any_len=True),
|
|
lambda c,a: UOp(Ops.STORE, src=a.src+c.src[1:]) if c.tag == 1 else None),
|
|
])
|
|
|
|
# 4. remove bufferize
|
|
|
|
def bufferize_to_store(x:UOp):
|
|
rngs = x.src[1:]
|
|
shape = tuple([r.vmax+1 for r in rngs])
|
|
assert prod(shape) > 0, f"no zero sized buffers {shape}"
|
|
buf = UOp.new_buffer(x.arg, prod(shape), x.dtype)
|
|
return buf.reshape(shape).index(*rngs, dtype=x.dtype.ptr(size=prod(shape))).store(x.src[0], *rngs)
|
|
|
|
def add_load_on_buffer(x:UOp, b:UOp):
|
|
if isinstance(x.dtype, PtrDType): return None
|
|
return x.replace(dtype=x.dtype.ptr(b.size)).load()
|
|
|
|
def add_load_on_store(x:UOp, st:UOp):
|
|
if isinstance(x.dtype, PtrDType): return None
|
|
rngs = x.src[1:]
|
|
shape = tuple([r.vmax+1 for r in rngs])
|
|
b = st.src[0].src[0]
|
|
assert b.op is Ops.BUFFER
|
|
return b.shrink(((0,prod(shape)),)).reshape(shape).index(*rngs, dtype=x.dtype.ptr(size=b.size)).load(st)
|
|
|
|
pm_add_buffers = pm_mops+PatternMatcher([
|
|
(UPat(Ops.BUFFERIZE, name="x"), bufferize_to_store),
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.BUFFER, name="b"), UPat()), name="x"), add_load_on_buffer),
|
|
(UPat(Ops.INDEX, src=(UPat(Ops.STORE, name="st"),), allow_any_len=True, name="x"), add_load_on_store),
|
|
])
|
|
|
|
# 5 (alt). create pointers
|
|
|
|
def debuf(ctx, b:UOp):
|
|
ret = UOp(Ops.DEFINE_GLOBAL, b.dtype.ptr(b.arg), arg=ctx[0])
|
|
ctx[0] += 1
|
|
return ret
|
|
|
|
pm_debuf = PatternMatcher([
|
|
(UPat(Ops.BUFFER, name="b"), debuf),
|
|
# HACK: consts shouldn't have srcs by here
|
|
(UPat(Ops.CONST, name="x"), lambda x: x.replace(src=()) if len(x.src) else None),
|
|
# no movement ops
|
|
(UPat(GroupOp.Movement, name="x"), lambda x: x.src[0]),
|
|
# HACK: no copy
|
|
(UPat(Ops.COPY, name="x"), lambda x: x.src[0]),
|
|
])
|
|
|
|
# 5. split into kernels
|
|
|
|
@dataclass
|
|
class LocalAddBufferContext:
|
|
dg:int = 0
|
|
map:dict = field(default_factory=dict)
|
|
|
|
def debuf(ctx:LocalAddBufferContext, b:UOp): return UOp(Ops.DEFINE_GLOBAL, b.dtype.ptr(b.arg), arg=ctx.map[b][1])
|
|
|
|
def split_load(ctx:LocalAddBufferContext, s:UOp):
|
|
b = s.src[0].src[0]
|
|
if b.op is not Ops.BUFFER: return None
|
|
|
|
if len(s.src) == 2 and s.src[1].op is Ops.ASSIGN:
|
|
assert len(s.src) == 2
|
|
lb = s.src[1]
|
|
assert b not in ctx.map or ctx.map[b][0] == lb
|
|
else:
|
|
lb = b
|
|
if b not in ctx.map:
|
|
ctx.map[b] = (lb, ctx.dg)
|
|
ctx.dg += 1
|
|
return s.replace(src=s.src[0:1]) if b is not lb else None
|
|
|
|
def handle_store(ctx:LocalAddBufferContext, s:UOp):
|
|
b = s.src[0].src[0]
|
|
if b.op is not Ops.BUFFER: return None
|
|
if b not in ctx.map:
|
|
ctx.map[b] = (b, ctx.dg)
|
|
ctx.dg += 1
|
|
if s.src[1].op is Ops.COPY: return s.src[1]
|
|
return None
|
|
|
|
to_define_global = PatternMatcher([
|
|
(UPat(Ops.BUFFER, name="b"), debuf),
|
|
(UPat(Ops.LOAD, name="s"), split_load),
|
|
(UPat(Ops.STORE, name="s"), handle_store),
|
|
])
|
|
|
|
def split_store(x:UOp):
|
|
if len(x.ranges): return None
|
|
shape = tuple([r.vmax+1 for r in x.src[2:]])
|
|
name = "k_"+'_'.join([str(s) for s in shape])
|
|
ctx = LocalAddBufferContext()
|
|
ret = graph_rewrite(x, to_define_global, ctx=ctx, name="kernel split", bottom_up=True)
|
|
ret = ret.sink(arg=KernelInfo(name=name)) if ret.op is Ops.STORE else ret
|
|
kernel = UOp(Ops.KERNEL, src=tuple([x[0] for x in ctx.map.values()]), arg=Kernel(ret, ()))
|
|
return kernel.src[0].assign(kernel)
|
|
|
|
split_kernels = PatternMatcher([
|
|
(UPat(Ops.STORE, name="x"), split_store),
|
|
])
|
|
|
|
@track_rewrites(name=lambda sink,ret: f"Schedule {pluralize('Kernel',len([u for u in ret[sink].toposort() if u.op is Ops.KERNEL]))}", replay=True)
|
|
def get_kernelize_map(sink:UOp) -> dict[UOp, UOp]:
|
|
tensor_map = graph_rewrite_map(sink, earliest_rewrites, name="earliest")
|
|
realize_map = {}
|
|
graph_rewrite(tensor_map[sink], do_realize, ctx=realize_map, name="Input Graph")
|
|
tensor_map = graph_rewrite_map(tensor_map[sink], add_contiguous, ctx=realize_map, bottom_up=True, input_map=tensor_map, name="add contiguous")
|
|
tensor_map = graph_rewrite_map(tensor_map[sink], early_cleanups+remove_tags, input_map=tensor_map, name="cleanup")
|
|
tensor_map = graph_rewrite_map(tensor_map[sink], pm_children, ctx=ChildrenContext(), bottom_up=True, input_map=tensor_map, name="children")
|
|
tensor_map = graph_rewrite_map(tensor_map[sink], pm_rangeify, ctx=RangeifyContext(), bottom_up=True, input_map=tensor_map, name="rangeify")
|
|
tensor_map = graph_rewrite_map(tensor_map[sink], pm_add_buffers, bottom_up=True, input_map=tensor_map, name="add buffers")
|
|
if getenv("VIZ"): graph_rewrite(tensor_map[sink], PatternMatcher([]), name="View Rangeify Graph")
|
|
|
|
# render
|
|
if getenv("SRC"):
|
|
rsink = tensor_map[sink]
|
|
from tinygrad.codegen.devectorizer import pm_reduce, ReduceContext
|
|
rsink = graph_rewrite(rsink, pm_reduce, ctx=ReduceContext(), name="remove reduce")
|
|
rsink = graph_rewrite(rsink, pm_debuf, ctx=[0], name="debuf", bottom_up=True)
|
|
from tinygrad.codegen import rewrites_for_linearizer, apply_rewrites
|
|
rsink = apply_rewrites(rsink, rewrites_for_linearizer)
|
|
from tinygrad.renderer.cstyle import CStyleLanguage
|
|
src = CStyleLanguage().render(rsink.arg.lst)
|
|
print(src)
|
|
return {sink:sink}
|
|
|
|
tensor_map = graph_rewrite_map(tensor_map[sink], split_kernels, input_map=tensor_map, name="split kernels")
|
|
if getenv("VIZ"): graph_rewrite(tensor_map[sink], PatternMatcher([]), name="View Kernel Graph")
|
|
return tensor_map
|