mirror of
https://github.com/tinygrad/tinygrad.git
synced 2026-08-29 12:36:07 +00:00
175 lines
8.5 KiB
Python
175 lines
8.5 KiB
Python
from dataclasses import dataclass, field
|
|
from tinygrad.uop.ops import UOp, UPat, PatternMatcher, Ops, GroupOp, graph_rewrite, track_rewrites
|
|
from tinygrad.dtype import dtypes, ImageDType
|
|
from tinygrad.helpers import prod, DEBUG, VIZ, pluralize
|
|
|
|
@dataclass
|
|
class AllocCtx:
|
|
uop_list: list[UOp] = field(default_factory=list)
|
|
buffer_map: dict[UOp, UOp] = field(default_factory=dict)
|
|
bases: set[UOp] = field(default_factory=set)
|
|
assigns: list[UOp] = field(default_factory=list)
|
|
replacements: list[UOp] = field(default_factory=list)
|
|
|
|
def tag_uop(ctx:AllocCtx, x:UOp):
|
|
if x.tag is not None: return None
|
|
ctx.uop_list.append(x)
|
|
return x.replace(tag=(len(ctx.uop_list)-1,))
|
|
|
|
def disk_copy_is_buffer(ctx:AllocCtx, u:UOp):
|
|
# copies to disk are replaced with the disk buffer
|
|
to_disk = isinstance(u._device, str) and u._device.startswith(("DISK", "TINYFS"))
|
|
if to_disk: ctx.buffer_map[u] = UOp.new_buffer(u.device, u.shard_size, u.dtype).reshape(u.max_shard_shape)
|
|
# all copies from disk/numpy are realized into a real buffer
|
|
from_creation = isinstance(u.src[0]._device, str) and any(u.src[0]._device.startswith(x) for x in ["NPY", "DISK", "PYTHON", "TINYFS"])
|
|
if from_creation: return tag_uop(ctx, u)
|
|
|
|
def apply_after(ctx:AllocCtx, u:UOp):
|
|
base = u.src[0]
|
|
while base.op is Ops.AFTER: base = base.src[0]
|
|
ctx.buffer_map[u] = base
|
|
|
|
# CONTIGUOUS and ASSIGN + parents are the only nodes that get updated
|
|
add_tags = PatternMatcher([
|
|
(UPat(Ops.COPY, name="u"), disk_copy_is_buffer),
|
|
# no tag on copies that are assigned
|
|
(UPat(Ops.ASSIGN, src=(UPat(), UPat(Ops.COPY, name="c")), name="a"),
|
|
lambda a,c: a.replace(src=(a.src[0], c.rtag(())), tag=a.tag+c.tag) if a.tag and c.tag else None),
|
|
(UPat(Ops.AFTER, name="u"), apply_after),
|
|
(UPat({Ops.CONTIGUOUS, Ops.ASSIGN}, name="x"), tag_uop),
|
|
(UPat(GroupOp.All, name="x"), lambda ctx,x: tag_uop(ctx,x) if x in ctx.bases else None),
|
|
])
|
|
|
|
def _buffer_like(u:UOp) -> UOp:
|
|
dtype = u.dtype
|
|
if isinstance(dtype, ImageDType):
|
|
if prod(dtype.shape) != prod(u.max_shard_shape) or ([x for x in u.max_shard_shape if x != 1] or [1])[-1] % 4 != 0:
|
|
if DEBUG >= 1: print(f"demoting Image {dtype} with shape {u.max_shard_shape}")
|
|
dtype = dtype.base
|
|
buffer = UOp.new_buffer(u.device, u.shard_size, dtype).reshape(u.max_shard_shape)
|
|
if isinstance(u.device, tuple) and u.axis is not None: buffer = buffer.multi(u.axis)
|
|
return buffer
|
|
|
|
def replace_contig_with_assign(u:UOp):
|
|
# if size is 0, remove the contig
|
|
if u.size == 0: return u.src[0]
|
|
# no real contig for DISK/TINYFS tensors, they are left alone
|
|
if isinstance(u._device, str) and u._device.startswith(("DISK", "TINYFS")): return u.rtag(None)
|
|
return _buffer_like(u).assign(u.src[0]).rtag(u.tag)
|
|
|
|
def replace_assign_with_contig(u:UOp):
|
|
assigned_to = u
|
|
while assigned_to.op in {Ops.ASSIGN, Ops.BITCAST, Ops.AFTER}: assigned_to = assigned_to.src[0].base
|
|
if assigned_to.op is not Ops.BUFFER:
|
|
return u.src[1].contiguous(tag=u.tag)
|
|
|
|
def contiguous_mops_to_view(c:UOp):
|
|
"""CONTIGUOUS(MOPS(BUFFER)) → CONTIGUOUS(BUFFER_VIEW) when movement ops collapse to a contiguous range."""
|
|
src = c.src[0]
|
|
buf = src.base
|
|
if buf.op not in {Ops.BUFFER, Ops.BUFFER_VIEW}: return None
|
|
if src.op is Ops.RESHAPE and src.src[0].op in {Ops.BUFFER, Ops.BUFFER_VIEW}: return None
|
|
|
|
# no symbolic shape
|
|
if not all(isinstance(x, int) for x in c.shape): return None
|
|
|
|
# check if view is supported
|
|
if not isinstance(c.device, str): return None
|
|
from tinygrad.device import Device
|
|
if not hasattr(Device[c.device].allocator, "_offset"): return None
|
|
|
|
# see if this can be a view
|
|
offset = src.contiguous_view_offset()
|
|
if offset is None: return None
|
|
|
|
# merge BUFFER_VIEWs
|
|
if buf.op is Ops.BUFFER_VIEW: offset, buf = offset + buf.arg[1], buf.src[0]
|
|
|
|
# NOTE: this contiguous is removed because this BUFFER_VIEW/RESHAPE has_buffer_identity
|
|
return UOp(Ops.BUFFER_VIEW, src.dtype, (buf,), (src.size, offset)).reshape(src.shape).contiguous(tag=c.tag)
|
|
|
|
def transform_precompiled_call(c:UOp) -> UOp|None:
|
|
if not c.arg.precompile: return None
|
|
if c.src[0].op is Ops.SINK: return None
|
|
out = _buffer_like(c)
|
|
fxn = out.param_like(len(c.src)-1).assign(c.src[0]).sink()
|
|
return out.after(c.replace(src=(fxn,)+tuple(x.contiguous() if x.op is not Ops.AFTER else x for x in c.src[1:])+(out,), dtype=dtypes.void, tag=None))
|
|
|
|
# NOTE: adding rules to here is bad. these all need to run before the schedule cache
|
|
pm_early_transform_tensor_graph = PatternMatcher([
|
|
# transform precompiled CALLs
|
|
(UPat(Ops.CALL, name="c"), transform_precompiled_call),
|
|
|
|
# CONTIGUOUS(MOPS(BUFFER/BUFFER_VIEW)) → CONTIGUOUS(BUFFER_VIEW) when movement ops collapse to contiguous range
|
|
(UPat(Ops.CONTIGUOUS, src=(UPat(GroupOp.Movement),), name="c"), contiguous_mops_to_view),
|
|
|
|
# add CONTIGUOUS to tagged UOps
|
|
(UPat(GroupOp.All-{Ops.CONTIGUOUS, Ops.ASSIGN}, name="x"), lambda x: x.rtag(None).contiguous(tag=x.tag) if x.tag else x.replace(tag=None)),
|
|
# remove extra CONTIGUOUS on ASSIGN (only when assign target is contiguous)
|
|
(UPat(Ops.CONTIGUOUS, src=(UPat(Ops.ASSIGN, name="a"),), name="c"),
|
|
lambda a,c: a.replace(tag=(a.tag or ())+(c.tag or ())) if a.src[0].has_buffer_identity() else None),
|
|
# replace ASSIGN with CONTIGUOUS
|
|
(UPat(Ops.ASSIGN, name="u"), replace_assign_with_contig),
|
|
# replace CONTIGUOUS with ASSIGNs
|
|
(UPat(Ops.CONTIGUOUS, name="u"), replace_contig_with_assign),
|
|
# remove DETACH/CONTIGUOUS_BACKWARD (allows more contiguous removal)
|
|
(UPat((Ops.DETACH, Ops.CONTIGUOUS_BACKWARD), name="x"), lambda x: x.src[0]),
|
|
])
|
|
|
|
def untag_and_append(ctx:AllocCtx, x:UOp):
|
|
if x.tag is None: return None
|
|
ret = x.replace(tag=None)
|
|
for t in x.tag:
|
|
original_uop: UOp = ctx.uop_list[t]
|
|
replace_uop = ret
|
|
while replace_uop.op is Ops.ASSIGN: replace_uop = replace_uop.src[0]
|
|
ctx.buffer_map[original_uop] = replace_uop.shrink_to(original_uop.shape)
|
|
ctx.assigns.append(ret)
|
|
return ret
|
|
|
|
def append_after(ctx:AllocCtx, x:UOp):
|
|
ctx.assigns.append(x)
|
|
|
|
def replace_input_buffer(ctx:AllocCtx, b:UOp):
|
|
ctx.replacements.append(b)
|
|
return UOp.param(len(ctx.replacements)-1, b.dtype, b.shape, b._device,
|
|
b._min_max if b.op is Ops.BIND else None, b.src[0].arg[0] if b.op is Ops.BIND else None)
|
|
|
|
pm_finalize_call = PatternMatcher([
|
|
(UPat(Ops.ASSIGN, name="x"), untag_and_append),
|
|
(UPat(Ops.AFTER, name="x"), append_after),
|
|
(UPat(Ops.COPY, name="x"), lambda ctx,x: append_after(ctx,x) if isinstance(x.device, str) and x.device.startswith(("DISK", "TINYFS")) else None),
|
|
# replace UNIQUE with LUNIQUE for CONST cache key normalization
|
|
(UPat(Ops.CONST, src=(UPat(Ops.UNIQUE), UPat(Ops.DEVICE, name="d")), name="b"), lambda b,d: b.replace(src=(d,))),
|
|
])
|
|
|
|
pm_replace_buf = PatternMatcher([
|
|
# replace BUFFER with PARAM for cache key normalization
|
|
(UPat(Ops.BUFFER, src=(UPat(Ops.UNIQUE), UPat(Ops.DEVICE)), name="b"), replace_input_buffer),
|
|
# replace BUFFER_VIEW with PARAM. this rewrite is bottom up so BUFFERs we don't need won't be in the input
|
|
(UPat(Ops.BUFFER_VIEW, src=(UPat(Ops.BUFFER),), name="b"), replace_input_buffer),
|
|
# strip value from BIND for cache key normalization, so different values hit same cache
|
|
(UPat(Ops.BIND, src=(UPat(Ops.DEFINE_VAR), UPat(Ops.CONST)), name="b"), replace_input_buffer),
|
|
])
|
|
|
|
@track_rewrites(lambda _,ret: f"Process {pluralize('Buffer', len(ret[1]))}")
|
|
def transform_to_call(big_sink:UOp) -> tuple[UOp, dict[UOp, UOp]]:
|
|
if VIZ: graph_rewrite(big_sink, PatternMatcher([]), name="View Tensor Graph")
|
|
# uop list is a list in the original_sink graph and we can map to the tags later
|
|
# here we build buffer map
|
|
dont_realize = {Ops.CONST, Ops.BUFFER, Ops.BIND, Ops.DEFINE_VAR, Ops.AFTER}
|
|
ctx = AllocCtx(bases=set([x.multibase for x in big_sink.src if x.base.op not in dont_realize]))
|
|
|
|
# this rewrite is "read-only", it adds simple things to buffer_map and may sink things on big_sink, bottom_up
|
|
# this is the only one where we have to be careful to not break the tensor graph
|
|
big_sink = graph_rewrite(big_sink, add_tags, ctx=ctx, bottom_up=True, name="number the uops")
|
|
|
|
# here we can break the tensor graph. this is the only place you need to maintain numbered tags
|
|
big_sink = graph_rewrite(big_sink, pm_early_transform_tensor_graph, name="early transform tensor graph")
|
|
|
|
# here we construct the final buffer_map. this is everything that will go into the tensor map
|
|
graph_rewrite(big_sink, pm_finalize_call, ctx=ctx, name="finalize call")
|
|
ret = graph_rewrite(UOp.sink(*ctx.assigns), pm_replace_buf, ctx=ctx, bottom_up=True, name="replace bufs").call(*ctx.replacements)
|
|
if VIZ: graph_rewrite(ret, PatternMatcher([]), name="View Call")
|
|
return ret, ctx.buffer_map
|