diff --git a/test/null/test_uop_graph.py b/test/null/test_uop_graph.py index 6f9464b8c9..0c8616bce6 100644 --- a/test/null/test_uop_graph.py +++ b/test/null/test_uop_graph.py @@ -429,7 +429,7 @@ class TestReduceCollapse(unittest.TestCase): class TestMovementOps(unittest.TestCase): def test_pm_mops_partial_reshape_index_removes_reshape(self): - from tinygrad.schedule.rangeify import pm_mops + from tinygrad.schedule.prepare import pm_mops src = UOp.param(0, dtypes.float, shape=(32, 4)) r0, r1 = UOp.range(4, 0), UOp.range(8, 1) result = graph_rewrite(src.reshape((4, 8, 4)).index(r0, r1), pm_mops, name="test") @@ -439,7 +439,7 @@ class TestMovementOps(unittest.TestCase): self.assertNotIn(Ops.RESHAPE, [u.op for u in result.toposort()]) def test_pm_mops_partial_reshape_index_suffix_mismatch_does_nothing(self): - from tinygrad.schedule.rangeify import pm_mops + from tinygrad.schedule.prepare import pm_mops src = UOp.param(0, dtypes.float, shape=(2, 6)) result = graph_rewrite(src.reshape((2, 3, 2)).index(UOp.range(2, 0)), pm_mops, name="test") self.assertEqual(result.op, Ops.INDEX) diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index 15e9edbd9f..5ffbcf2e84 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -22,7 +22,7 @@ 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, pm_reduce_unparented from tinygrad.schedule.multi import multi_pm -from tinygrad.schedule.rangeify import pm_mops +from tinygrad.schedule.prepare import pm_mops 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.coalesce import memory_coalescing, pm_simplify_add_image diff --git a/tinygrad/schedule/__init__.py b/tinygrad/schedule/__init__.py index f00ec3f7c1..9608e032fa 100644 --- a/tinygrad/schedule/__init__.py +++ b/tinygrad/schedule/__init__.py @@ -80,6 +80,7 @@ def create_schedule(sched_sink:UOp) -> UOp: from tinygrad.schedule.memory import memory_plan_rewrite from tinygrad.engine.realize import capturing, pm_flatten_linear +from tinygrad.schedule.prepare import prepare_rangeify from tinygrad.schedule.rangeify import get_kernel_graph from tinygrad.helpers import CAPTURING from tinygrad.uop.ops import PatternMatcher, UPat, ParamArg @@ -123,7 +124,7 @@ def lower_sink_to_linear(function:UOp) -> UOp|None: if not SCACHE or (sc_ret:=schedule_cache.get(cache_key, None)) is None: if SPEC: type_verify(function, spec_tensor) # support recursive CALLs - linear = create_schedule(get_kernel_graph(function)) + linear = create_schedule(get_kernel_graph(prepare_rangeify(function))) if SCACHE: schedule_cache[cache_key] = linear else: # schedule cache hit @@ -156,7 +157,7 @@ def simplify_copy_kernel(call:UOp, ast:UOp, dst:UOp, src:UOp): # NOTE: this is a codegen for SDMA devices if dst.device == src.device and not (isinstance(dst.device, str) and dst.device.startswith("DISK")): return None from tinygrad.codegen.simplify import pm_flatten_range, pm_simplify_ranges - from tinygrad.schedule.rangeify import pm_mops + from tinygrad.schedule.prepare import pm_mops from tinygrad.uop.symbolic import sym sink = graph_rewrite(ast, sym+pm_mops+pm_flatten_range+pm_simplify_ranges, ctx={}, name="simplify ranges in copy") return call.replace(src=(sink,) + call.src[1:]) diff --git a/tinygrad/schedule/prepare.py b/tinygrad/schedule/prepare.py new file mode 100644 index 0000000000..7fac55277d --- /dev/null +++ b/tinygrad/schedule/prepare.py @@ -0,0 +1,205 @@ +import itertools +from tinygrad.dtype import dtypes, to_dtype +from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, resolve, GroupOp +from tinygrad.uop.ops import graph_rewrite, rewrite_group, shape_to_shape_arg, ParamArg, identity_element +from tinygrad.uop.movement import mop_cleanup +from tinygrad.helpers import prod, getenv, all_int, DEBUG, SPLIT_REDUCEOP, OPENPILOT_HACKS, FLOAT16, argsort +from tinygrad.schedule.indexing import apply_movement_op +from tinygrad.schedule.allreduce import create_allreduce_function +from tinygrad.schedule.multi import multi_pm + +def found_after(ctx:dict[UOp, UOp], after:UOp, src:UOp): + if (x:=src).op is Ops.CAST and x.dtype == dtypes.half and FLOAT16: x, after = x.src[0], after.cast(dtypes.float) + while True: + if x.op is Ops.PERMUTE: x, after = x.src[0], after.permute(argsort(x.marg)) + elif x.op is Ops.RESHAPE: x, after = x.src[0], after.reshape(x.src[0].shape) + elif x.op is Ops.WHERE and x.src[2].base.is_invalid and x.src[1].op is Ops.PAD: + x, after = x.src[1].src[0], after.shrink(tuple((o, s+o) for (o,_),s in zip(x.src[1].marg, x.src[1].src[0].shape))) + else: break + ctx[x] = after + +# *** fold moved AFTERs (hack for openpilot) *** +pm_fold_moved_after = PatternMatcher([ + (UPat(Ops.AFTER, src=(UPat(), UPat(Ops.STORE, src=(UPat(), UPat((*GroupOp.Movement,Ops.CAST,Ops.WHERE), name="src")))), name="after"), found_after), + # replace ALU sources with AFTER versions found above + (UPat(GroupOp.ALU, name="alu"), lambda ctx,alu: alu.replace(src=new_src) if (new_src:=tuple(ctx.get(s, s) for s in alu.src)) != alu.src else None), +]) + +# movement op on INDEX as a PatternMatcher +def _mop_index(r:UOp, idx:UOp): + idxs = idx.src[1:] + if len(idxs) == len(r.shape): + return r.src[0].index(*apply_movement_op(r.op, r.src[0].shape, r.marg, idxs), arg=idx.arg) + if r.op is Ops.RESHAPE: + src_prefix = len(r.src[0].shape) - len(r.shape[len(idxs):]) + if src_prefix >= 0 and r.src[0].shape[src_prefix:] == r.shape[len(idxs):]: + if src_prefix == 0: return r.src[0] if r.src[0].dtype == idx.dtype else None + ret = r.src[0].index(*apply_movement_op(r.op, r.src[0].shape[:src_prefix], r.shape[:len(idxs)], idxs), arg=idx.arg) + return ret if ret.shape == idx.shape else None + +pm_mops = PatternMatcher([ + # handle movement ops on INDEX + (UPat(GroupOp.Movement, name="r").f(Ops.INDEX, allow_any_len=True, name="idx"), _mop_index), + # move movement ops and INDEX after AFTER + (UPat(GroupOp.Movement|{Ops.INDEX}, name="r").after(name="a", allow_any_len=True), + lambda r,a: UOp(r.op, src=(a.replace(src=(r.src[0],)+a.src[1:]),)+r.src[1:], arg=r.arg)), + (UPat(GroupOp.Movement, name="r").end(name="a", allow_any_len=True), lambda r,a: a.replace(src=(r.src[0],)+a.src[1:])), +]) + +# ***************** +# 0. do some cleanup rewrites, mostly copied from the old stuff + +def fix_store_hazard(target:UOp, src:UOp): + if (base:=target.base) not in src.toposort(enter_calls=False): return None + # PERMUTE and FLIP reorder indices, SHRINK can have overlapping regions when dest is also shrunk + unsafe = {Ops.PERMUTE, Ops.FLIP} | ({Ops.SHRINK} if target.op_in_backward_slice_with_self(Ops.SHRINK) else set()) + reaches_base: dict[UOp, bool] = {} + for s in src.toposort(gate=lambda s: s.op is not Ops.CONTIGUOUS): + reaches_base[s] = s is base or any(reaches_base.get(c) for c in s.src) + if reaches_base[s] and s.op in unsafe and not (s is target and s.op is Ops.SHRINK): return target.store(src.contiguous()) + +def split_reduceop(reduce:UOp, x:UOp): + if prod(reduce.shape) == 0: return None + if not SPLIT_REDUCEOP or not all_int(x.shape) or (prod(x.shape)//prod(reduce.shape))1) else 0 for i,s in enumerate(x.shape)]) + range_nums = [y.arg[0] for y in indexed.substitute({x.base:UOp(Ops.NOOP)}, extra_pm=pm_mops).ranges] + is_expanded = [i not in range_nums for i in range(len(x.shape))] + + if not (split_candidates:=[(i,d) for i in range(reduce.arg[1]) + for d in range(min(256,2**getenv("REDUCEOP_SPLIT_SIZE",22)//prod(reduce.shape)),8-1,-1) + if x.shape[i]%d==0 and not is_expanded[i]]): return None + dim_to_split, divisor = split_candidates[0] + splitted_shape = x.shape[:dim_to_split]+(divisor,)+(x.shape[dim_to_split]//divisor,)+x.shape[dim_to_split+1:] + splitted = x.reshape(splitted_shape).permute(tuple([d for d in range(len(splitted_shape)) if d!=dim_to_split]+[dim_to_split])) + if DEBUG >= 3: print(f"split {divisor}: {x.shape} -> {splitted.shape} -> {reduce.shape}") + # reduce original axes, then split + return splitted._rop(reduce.arg[0], tuple(range(reduce.arg[1]))).contiguous()._rop(reduce.arg[0], (len(reduce.shape),)).reshape(reduce.shape) + +pm_gather_params = PatternMatcher([ (UPat(Ops.PARAM, name="p"), lambda ctx, p: ctx.append(p) if p.arg.slot >= 0 else None), ]) +def resolve_function(c:UOp, allow_param_mismatch=True) -> UOp|None: + if c.arg.precompile: return None + params: list[UOp] = [] + graph_rewrite(c.src[0], pm_gather_params, bottom_up=True, ctx=params, name="gather params") + params = sorted(params, key=lambda x: x.arg.slot) + args = c.src[1:] + + # NOTE: this isn't really needed. it's okay if there's unused args in the function + if not allow_param_mismatch: + if [x.arg.slot for x in params] != list(range(len(params))): raise RuntimeError(f"params not in order: {[x.arg.slot for x in params]}") + if len(params) != len(args): raise TypeError(f"expected {len(params)} args, got {len(args)}") + + dict_map = {x:args[x.arg.slot] for x in params} + for i, (p, a) in enumerate(dict_map.items()): + if p.axis != a.axis: raise TypeError(f"arg {i} axis mismatch: expected {p.axis}, got {a.axis}") + if p.max_shape != a.max_shape: raise TypeError(f"arg {i} shape mismatch: expected {p.shape}, got {a.shape}") + if p.dtype != a.dtype: raise TypeError(f"arg {i} dtype mismatch: expected {p.dtype}, got {a.dtype}") + return c.src[0].substitute(dict_map, walk=True) + +# shape-changing bitcast +def expand_bitcast(bc:UOp) -> UOp|None: + x = bc.src[0] + if (ns:=bc.dtype.itemsize) == (os:=x.dtype.itemsize) or (isinstance(x.device, str) and x.device.startswith(("DISK", "TINYFS"))): return None + new_uint, tmp = to_dtype(f"uint{8*ns}"), x.bitcast(to_dtype(f"uint{8*os}")) + if ns > os: + tmp = tmp.reshape(x.shape[:-1] + (x.shape[-1]//(rate := ns//os), rate)) + parts = [tmp.shrink((None,)*(len(tmp.shape)-1) + ((i, i+1),)).cast(new_uint)<<8*i*os for i in range(rate)] + return parts[0].usum(*parts[1:]).squeeze(-1).bitcast(bc.dtype) + parts = [tmp>>8*i*ns for i in range(os//ns)] + return parts[0].stack(*parts[1:], dim=-1).flatten(-2).cast(new_uint).bitcast(bc.dtype) + +earliest_rewrites = mop_cleanup+PatternMatcher([ + # resolve FUNCTION calls (inline the body) + (UPat(Ops.FUNCTION, name="c"), resolve_function), + + # resolve TUPLE+GETTUPLE + (UPat(Ops.GETTUPLE, src=(UPat(Ops.TUPLE, name="t"),), name="g"), lambda g,t: t.src[g.arg]), + + # resolve allreduce (must be bottom up) + (UPat(Ops.ALLREDUCE, src=(UPat.var("buf"),), name="red"), create_allreduce_function), + + # split_reduceop + (UPat(Ops.REDUCE, name="reduce", src=(UPat.var("x"),)), split_reduceop), + + # remove DETACH/CONTIGUOUS_BACKWARD (TODO: this is copied in allocations) + (UPat((Ops.DETACH, Ops.CONTIGUOUS_BACKWARD), name="x"), lambda x: x.src[0]), + + # SINK only ever references the base + (UPat(Ops.SINK, name="x"), lambda x: x.replace(src=tuple(y.unsharded_base for y in x.src))), + + # ** copy rules ** + + # COPY transfers a contiguous range, so materialize a source that's resized (shrink/pad/expand) or reordered (permute/flip) + (UPat(Ops.COPY, src=(UPat(GroupOp.Movement, name="r"),), name="c"), + lambda c,r: c.replace(src=(r.contiguous(),)) if resolve(r.numel() != r.base.numel(), False) or r.contiguous_view_offset() is None else None), + + # copy to same device is a no-op + (UPat(Ops.COPY, src=(UPat.var("x"),), name="copy"), lambda x,copy: x if x.device == copy.device else None), + + # copy on reshape is reshape on copy + (UPat(Ops.COPY, src=(UPat(Ops.RESHAPE, name="shp"),), name="cpy"), lambda shp,cpy: shp.src[0].copy_to_device(cpy.device).reshape(shp.shape)), + + # reshaping on STORE can be a NOOP + (UPat(Ops.STORE, src=(UPat(Ops.RESHAPE, src=(UPat.var("dst",),), allow_any_len=True), + UPat(Ops.RESHAPE, src=(UPat.var("src",),), allow_any_len=True))), + lambda dst,src: dst.store(src) if dst.shape == src.shape else None), + + # ** store rules ** + + # fix store hazard (dest is in used in src) by adding contiguous: TestAssign.test_post_flipped_assignment + (UPat(Ops.STORE, src=(UPat(name="target"), UPat(name="src"))), fix_store_hazard), + + # remove two STOREs that store the same thing to the same place: TestSchedule.test_dedup_Assign + (UPat.var("buf").after(UPat.var("buf").store(UPat.var("src")), name="a1").after(UPat.var("a1").store(UPat.var("src"))), lambda buf,src,a1:a1), + + # store a buffer's own current contents back into itself: TestAssign.test_nested_after_contiguous_store_no_init + (UPat.var("buf").after(UPat.var("buf").store(UPat.var("buf").after(UPat.var("buf").store(UPat.var("src")), name="a1"))), lambda buf,src,a1:a1), + + # move bitcast from store dest to source: TestAssign.test_assign_bitcast + (UPat(Ops.STORE, src=(UPat(Ops.BITCAST, src=(UPat(name="target"),)), UPat(name="src"))), + lambda target, src: target.store(src.bitcast(target.dtype))), + + (UPat(Ops.BITCAST, name="bc"), expand_bitcast), + + # ** size 0 ** + + # reduce of size 0 is the identity element + (UPat(Ops.REDUCE, name="reduce", src=(UPat.var("x"),)), + lambda reduce,x: reduce.const_like(identity_element(reduce.arg[0], reduce.dtype)) if 0 in x.shape and 0 not in reduce.shape else None), + # handle size 0 + (UPat(GroupOp.All-{Ops.SINK}, name="x"), lambda x: x.const_like(0).rtag(x.tag) if x._shape is not None and 0 in x.shape else None), +]) + +def convert_copy_to_store(ctx, copy:UOp, existing_buf:UOp|None=None): + input_src = copy.src[0] + if not input_src.has_buffer_identity(after_ok=True): input_src = input_src.contiguous() + input_src = input_src.flatten() + if existing_buf is not None: + # if the existing buffer is not a full buffer, we can't use it + if not existing_buf.has_buffer_identity(after_ok=True): return None + # if there's already a buffer, we just use it + return existing_buf.flatten().store(input_src) + # create the output buffer + buf = UOp(Ops.BUFFER, src=(shape_to_shape_arg(input_src.max_shape),), arg=ParamArg(next(ctx), copy.dtype, device=copy.device)) + # reshape back to input + return buf.after(buf.store(input_src)).reshape(copy.shape) + +pm_copy_to_store = PatternMatcher([ + (UPat(name="existing_buf").store(UPat(Ops.COPY, name="copy")), convert_copy_to_store), + (UPat(Ops.COPY, name="copy"), convert_copy_to_store), +]) + +@rewrite_group(new_ctx=False) +def prepare_rangeify(sink:UOp) -> UOp: + # prepare for rangeify + tsink = graph_rewrite(sink, multi_pm, name="multi_pm") + if OPENPILOT_HACKS: tsink = graph_rewrite(tsink, pm_fold_moved_after, ctx={}, name="fold moved afters") + tsink = graph_rewrite(tsink, pm_mops+earliest_rewrites, bottom_up=True, name="earliest rewrites") + tsink = graph_rewrite(tsink, pm_copy_to_store, ctx=itertools.count(0), bottom_up=True, name="convert copy to store") + return tsink diff --git a/tinygrad/schedule/rangeify.py b/tinygrad/schedule/rangeify.py index 89fc1b622c..0b8b152d11 100644 --- a/tinygrad/schedule/rangeify.py +++ b/tinygrad/schedule/rangeify.py @@ -1,191 +1,21 @@ from dataclasses import dataclass, field, replace from typing import cast import itertools -from tinygrad.dtype import dtypes, AddrSpace, Invalid, to_dtype, strong_dtype +from tinygrad.dtype import dtypes, AddrSpace, Invalid, strong_dtype from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, resolve, GroupOp, KernelInfo, ParamArg, shape_to_shape_arg -from tinygrad.uop.ops import graph_rewrite, sint, AxisType, BottomUpGate, rewrite_group, identity_element +from tinygrad.uop.ops import graph_rewrite, sint, AxisType, BottomUpGate, rewrite_group from tinygrad.uop.symbolic import symbolic -from tinygrad.uop.movement import mop_cleanup -from tinygrad.helpers import prod, getenv, dedup, all_int, DEBUG, SPLIT_REDUCEOP, DEBUG_RANGEIFY, VIZ, MAX_KERNEL_BUFFERS, SPEC -from tinygrad.helpers import PCONTIG, FLOAT16, OPENPILOT_HACKS, argsort, partition, get_single_element +from tinygrad.helpers import prod, dedup, DEBUG_RANGEIFY, VIZ, MAX_KERNEL_BUFFERS, SPEC +from tinygrad.helpers import PCONTIG, partition, get_single_element from tinygrad.codegen.simplify import pm_flatten_range, pm_reduce_simplify from tinygrad.codegen.opt import Opt from tinygrad.schedule.indexing import run_rangeify, BufferizeOpts, apply_movement_op -from tinygrad.schedule.multi import multi_pm -from tinygrad.schedule.allreduce import create_allreduce_function +from tinygrad.schedule.prepare import pm_mops # creation can recurse a lot import sys sys.setrecursionlimit(10000) -def found_after(ctx:dict[UOp, UOp], after:UOp, src:UOp): - if (x:=src).op is Ops.CAST and x.dtype == dtypes.half and FLOAT16: x, after = x.src[0], after.cast(dtypes.float) - while True: - if x.op is Ops.PERMUTE: x, after = x.src[0], after.permute(argsort(x.marg)) - elif x.op is Ops.RESHAPE: x, after = x.src[0], after.reshape(x.src[0].shape) - elif x.op is Ops.WHERE and x.src[2].base.is_invalid and x.src[1].op is Ops.PAD: - x, after = x.src[1].src[0], after.shrink(tuple((o, s+o) for (o,_),s in zip(x.src[1].marg, x.src[1].src[0].shape))) - else: break - ctx[x] = after - -# *** fold moved AFTERs (hack for openpilot) *** -pm_fold_moved_after = PatternMatcher([ - (UPat(Ops.AFTER, src=(UPat(), UPat(Ops.STORE, src=(UPat(), UPat((*GroupOp.Movement,Ops.CAST,Ops.WHERE), name="src")))), name="after"), found_after), - # replace ALU sources with AFTER versions found above - (UPat(GroupOp.ALU, name="alu"), lambda ctx,alu: alu.replace(src=new_src) if (new_src:=tuple(ctx.get(s, s) for s in alu.src)) != alu.src else None), -]) - -# movement op on INDEX as a PatternMatcher -def _mop_index(r:UOp, idx:UOp): - idxs = idx.src[1:] - if len(idxs) == len(r.shape): - return r.src[0].index(*apply_movement_op(r.op, r.src[0].shape, r.marg, idxs), arg=idx.arg) - if r.op is Ops.RESHAPE: - src_prefix = len(r.src[0].shape) - len(r.shape[len(idxs):]) - if src_prefix >= 0 and r.src[0].shape[src_prefix:] == r.shape[len(idxs):]: - if src_prefix == 0: return r.src[0] if r.src[0].dtype == idx.dtype else None - ret = r.src[0].index(*apply_movement_op(r.op, r.src[0].shape[:src_prefix], r.shape[:len(idxs)], idxs), arg=idx.arg) - return ret if ret.shape == idx.shape else None - -pm_mops = PatternMatcher([ - # handle movement ops on INDEX - (UPat(GroupOp.Movement, name="r").f(Ops.INDEX, allow_any_len=True, name="idx"), _mop_index), - # move movement ops and INDEX after AFTER - (UPat(GroupOp.Movement|{Ops.INDEX}, name="r").after(name="a", allow_any_len=True), - lambda r,a: UOp(r.op, src=(a.replace(src=(r.src[0],)+a.src[1:]),)+r.src[1:], arg=r.arg)), - (UPat(GroupOp.Movement, name="r").end(name="a", allow_any_len=True), lambda r,a: a.replace(src=(r.src[0],)+a.src[1:])), -]) - -# ***************** -# 0. do some cleanup rewrites, mostly copied from the old stuff - -def fix_store_hazard(target:UOp, src:UOp): - if (base:=target.base) not in src.toposort(enter_calls=False): return None - # PERMUTE and FLIP reorder indices, SHRINK can have overlapping regions when dest is also shrunk - unsafe = {Ops.PERMUTE, Ops.FLIP} | ({Ops.SHRINK} if target.op_in_backward_slice_with_self(Ops.SHRINK) else set()) - reaches_base: dict[UOp, bool] = {} - for s in src.toposort(gate=lambda s: s.op is not Ops.CONTIGUOUS): - reaches_base[s] = s is base or any(reaches_base.get(c) for c in s.src) - if reaches_base[s] and s.op in unsafe and not (s is target and s.op is Ops.SHRINK): return target.store(src.contiguous()) - -def split_reduceop(reduce:UOp, x:UOp): - if prod(reduce.shape) == 0: return None - if not SPLIT_REDUCEOP or not all_int(x.shape) or (prod(x.shape)//prod(reduce.shape))1) else 0 for i,s in enumerate(x.shape)]) - range_nums = [y.arg[0] for y in indexed.substitute({x.base:UOp(Ops.NOOP)}, extra_pm=pm_mops).ranges] - is_expanded = [i not in range_nums for i in range(len(x.shape))] - - if not (split_candidates:=[(i,d) for i in range(reduce.arg[1]) - for d in range(min(256,2**getenv("REDUCEOP_SPLIT_SIZE",22)//prod(reduce.shape)),8-1,-1) - if x.shape[i]%d==0 and not is_expanded[i]]): return None - dim_to_split, divisor = split_candidates[0] - splitted_shape = x.shape[:dim_to_split]+(divisor,)+(x.shape[dim_to_split]//divisor,)+x.shape[dim_to_split+1:] - splitted = x.reshape(splitted_shape).permute(tuple([d for d in range(len(splitted_shape)) if d!=dim_to_split]+[dim_to_split])) - if DEBUG >= 3: print(f"split {divisor}: {x.shape} -> {splitted.shape} -> {reduce.shape}") - # reduce original axes, then split - return splitted._rop(reduce.arg[0], tuple(range(reduce.arg[1]))).contiguous()._rop(reduce.arg[0], (len(reduce.shape),)).reshape(reduce.shape) - -pm_gather_params = PatternMatcher([ (UPat(Ops.PARAM, name="p"), lambda ctx, p: ctx.append(p) if p.arg.slot >= 0 else None), ]) -def resolve_function(c:UOp, allow_param_mismatch=True) -> UOp|None: - if c.arg.precompile: return None - params: list[UOp] = [] - graph_rewrite(c.src[0], pm_gather_params, bottom_up=True, ctx=params, name="gather params") - params = sorted(params, key=lambda x: x.arg.slot) - args = c.src[1:] - - # NOTE: this isn't really needed. it's okay if there's unused args in the function - if not allow_param_mismatch: - if [x.arg.slot for x in params] != list(range(len(params))): raise RuntimeError(f"params not in order: {[x.arg.slot for x in params]}") - if len(params) != len(args): raise TypeError(f"expected {len(params)} args, got {len(args)}") - - dict_map = {x:args[x.arg.slot] for x in params} - for i, (p, a) in enumerate(dict_map.items()): - if p.axis != a.axis: raise TypeError(f"arg {i} axis mismatch: expected {p.axis}, got {a.axis}") - if p.max_shape != a.max_shape: raise TypeError(f"arg {i} shape mismatch: expected {p.shape}, got {a.shape}") - if p.dtype != a.dtype: raise TypeError(f"arg {i} dtype mismatch: expected {p.dtype}, got {a.dtype}") - return c.src[0].substitute(dict_map, walk=True) - -# shape-changing bitcast -def expand_bitcast(bc:UOp) -> UOp|None: - x = bc.src[0] - if (ns:=bc.dtype.itemsize) == (os:=x.dtype.itemsize) or (isinstance(x.device, str) and x.device.startswith(("DISK", "TINYFS"))): return None - new_uint, tmp = to_dtype(f"uint{8*ns}"), x.bitcast(to_dtype(f"uint{8*os}")) - if ns > os: - tmp = tmp.reshape(x.shape[:-1] + (x.shape[-1]//(rate := ns//os), rate)) - parts = [tmp.shrink((None,)*(len(tmp.shape)-1) + ((i, i+1),)).cast(new_uint)<<8*i*os for i in range(rate)] - return parts[0].usum(*parts[1:]).squeeze(-1).bitcast(bc.dtype) - parts = [tmp>>8*i*ns for i in range(os//ns)] - return parts[0].stack(*parts[1:], dim=-1).flatten(-2).cast(new_uint).bitcast(bc.dtype) - -earliest_rewrites = mop_cleanup+PatternMatcher([ - # resolve FUNCTION calls (inline the body) - (UPat(Ops.FUNCTION, name="c"), resolve_function), - - # resolve TUPLE+GETTUPLE - (UPat(Ops.GETTUPLE, src=(UPat(Ops.TUPLE, name="t"),), name="g"), lambda g,t: t.src[g.arg]), - - # resolve allreduce (must be bottom up) - (UPat(Ops.ALLREDUCE, src=(UPat.var("buf"),), name="red"), create_allreduce_function), - - # split_reduceop - (UPat(Ops.REDUCE, name="reduce", src=(UPat.var("x"),)), split_reduceop), - - # remove DETACH/CONTIGUOUS_BACKWARD (TODO: this is copied in allocations) - (UPat((Ops.DETACH, Ops.CONTIGUOUS_BACKWARD), name="x"), lambda x: x.src[0]), - - # SINK only ever references the base - (UPat(Ops.SINK, name="x"), lambda x: x.replace(src=tuple(y.unsharded_base for y in x.src))), - - # ** copy rules ** - - # COPY transfers a contiguous range, so materialize a source that's resized (shrink/pad/expand) or reordered (permute/flip) - (UPat(Ops.COPY, src=(UPat(GroupOp.Movement, name="r"),), name="c"), - lambda c,r: c.replace(src=(r.contiguous(),)) if resolve(r.numel() != r.base.numel(), False) or r.contiguous_view_offset() is None else None), - - # copy to same device is a no-op - (UPat(Ops.COPY, src=(UPat.var("x"),), name="copy"), lambda x,copy: x if x.device == copy.device else None), - - # copy on reshape is reshape on copy - (UPat(Ops.COPY, src=(UPat(Ops.RESHAPE, name="shp"),), name="cpy"), lambda shp,cpy: shp.src[0].copy_to_device(cpy.device).reshape(shp.shape)), - - # reshaping on STORE can be a NOOP - (UPat(Ops.STORE, src=(UPat(Ops.RESHAPE, src=(UPat.var("dst",),), allow_any_len=True), - UPat(Ops.RESHAPE, src=(UPat.var("src",),), allow_any_len=True))), - lambda dst,src: dst.store(src) if dst.shape == src.shape else None), - - # ** store rules ** - - # fix store hazard (dest is in used in src) by adding contiguous: TestAssign.test_post_flipped_assignment - (UPat(Ops.STORE, src=(UPat(name="target"), UPat(name="src"))), fix_store_hazard), - - # remove two STOREs that store the same thing to the same place: TestSchedule.test_dedup_assign - (UPat.var("buf").after(UPat.var("buf").store(UPat.var("src")), name="a1").after(UPat.var("a1").store(UPat.var("src"))), lambda buf,src,a1:a1), - - # store a buffer's own current contents back into itself: TestAssign.test_nested_after_contiguous_store_no_init - (UPat.var("buf").after(UPat.var("buf").store(UPat.var("buf").after(UPat.var("buf").store(UPat.var("src")), name="a1"))), lambda buf,src,a1:a1), - - # move bitcast from store dest to source: TestAssign.test_assign_bitcast - (UPat(Ops.STORE, src=(UPat(Ops.BITCAST, src=(UPat(name="target"),)), UPat(name="src"))), - lambda target, src: target.store(src.bitcast(target.dtype))), - - (UPat(Ops.BITCAST, name="bc"), expand_bitcast), - - # ** size 0 ** - - # reduce of size 0 is the identity element - (UPat(Ops.REDUCE, name="reduce", src=(UPat.var("x"),)), - lambda reduce,x: reduce.const_like(identity_element(reduce.arg[0], reduce.dtype)) if 0 in x.shape and 0 not in reduce.shape else None), - # handle size 0 - (UPat(GroupOp.All-{Ops.SINK}, name="x"), lambda x: x.const_like(0).rtag(x.tag) if x._shape is not None and 0 in x.shape else None), -]) - # ***************** # 3.5 cleanups @@ -562,33 +392,8 @@ split_kernels = PatternMatcher([ (UPat((Ops.STORE, Ops.END), name="x"), split_store), ]) -def convert_copy_to_store(ctx, copy:UOp, existing_buf:UOp|None=None): - input_src = copy.src[0] - if not input_src.has_buffer_identity(after_ok=True): input_src = input_src.contiguous() - input_src = input_src.flatten() - if existing_buf is not None: - # if the existing buffer is not a full buffer, we can't use it - if not existing_buf.has_buffer_identity(after_ok=True): return None - # if there's already a buffer, we just use it - return existing_buf.flatten().store(input_src) - # create the output buffer - buf = UOp(Ops.BUFFER, src=(shape_to_shape_arg(input_src.max_shape),), arg=ParamArg(next(ctx), copy.dtype, device=copy.device)) - # reshape back to input - return buf.after(buf.store(input_src)).reshape(copy.shape) - -pm_copy_to_store = PatternMatcher([ - (UPat(name="existing_buf").store(UPat(Ops.COPY, name="copy")), convert_copy_to_store), - (UPat(Ops.COPY, name="copy"), convert_copy_to_store), -]) - @rewrite_group(new_ctx=False) -def get_kernel_graph(sink:UOp) -> UOp: - # prepare for rangeify - tsink = graph_rewrite(sink, multi_pm, name="multi_pm") - if OPENPILOT_HACKS: tsink = graph_rewrite(tsink, pm_fold_moved_after, ctx={}, name="fold moved afters") - tsink = graph_rewrite(tsink, pm_mops+earliest_rewrites, bottom_up=True, name="earliest rewrites") - tsink = graph_rewrite(tsink, pm_copy_to_store, ctx=itertools.count(0), bottom_up=True, name="convert copy to store") - +def get_kernel_graph(tsink:UOp) -> UOp: # convert movement ops to ranges tsink = run_rangeify(tsink, bool(DEBUG_RANGEIFY)) diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index 41358fa5cb..12bce3a8af 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -890,7 +890,7 @@ class UOp(RandMixin, metaclass=UOpMetaClass): return s def contiguous_view(self) -> tuple[UOp, int]|None: - from tinygrad.schedule.rangeify import pm_mops + from tinygrad.schedule.prepare import pm_mops from tinygrad.uop.symbolic import symbolic # WEBGPU and CL do not support views.