diff --git a/test/test_linearizer_dumb.py b/test/test_linearizer_dumb.py index ce6d5ec144..d14d3a6ae3 100644 --- a/test/test_linearizer_dumb.py +++ b/test/test_linearizer_dumb.py @@ -23,7 +23,7 @@ class TestLinearizerFailure(unittest.TestCase): c9 = UOp(Ops.DEFINE_GLOBAL, dtypes.uchar.ptr(47040000), arg=2, src=()) c10 = c9.index((((c3*UOp.const(dtypes.index, 4704000))+c2)+(c6*UOp.const(dtypes.index, 784))).valid(UOp.const(dtypes.bool, True))).load() c11 = c5.alu(Ops.CMPNE, ((((c3*UOp.const(dtypes.index, 6000))+c6)+((c7*UOp.const(dtypes.index, 16))+c8)).alu(Ops.CMPLT, UOp.const(dtypes.index, 59999)).where(UOp.const(dtypes.int, 0), UOp.const(dtypes.int, 1)).reduce(c7, c8, arg=Ops.ADD)+UOp.const(dtypes.int, -1))).where(UOp.const(dtypes.uchar, 0), c10).reduce(c6, arg=Ops.ADD) - c12 = c0.index((((c1*UOp.const(dtypes.index, 7840))+(c2*UOp.const(dtypes.index, 10)))+c3).valid(UOp.const(dtypes.bool, True))).store(c11, c1, c2, c3) + c12 = c0.index((((c1*UOp.const(dtypes.index, 7840))+(c2*UOp.const(dtypes.index, 10)))+c3).valid(UOp.const(dtypes.bool, True))).store(c11).end(c1, c2, c3) ast = c12.sink(arg=KernelInfo(name='test', axis_types=(), dont_use_locals=False, applied_opts=(Opt(op=OptOps.GROUP, axis=1, arg=16),), opts_to_apply=None)) _ = get_program(ast, Device["METAL"].renderer) diff --git a/test/test_linearizer_failures.py b/test/test_linearizer_failures.py index 7bd0864c8e..7917fa04d5 100644 --- a/test/test_linearizer_failures.py +++ b/test/test_linearizer_failures.py @@ -16,7 +16,7 @@ class TestLinearizerFailures(unittest.TestCase): c7 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(64), arg=2, src=()) c8 = c7.index(c3).load() c9 = ((((c6+(c8*UOp.const(dtypes.float, -1.0)))*(c6+(c8*UOp.const(dtypes.float, -1.0)))).reduce(c5, arg=Ops.ADD)*UOp.const(dtypes.float, 0.000390625))+UOp.const(dtypes.float, 1e-05)).sqrt().reciprocal() - c10 = c0.index(c3).store(c9, c1, c2) + c10 = c0.index(c3).store(c9).end(c1, c2) ast = c10.sink() get_program(ast) diff --git a/test/test_uop_graph.py b/test/test_uop_graph.py index f26dcf9705..1f5a849fcf 100644 --- a/test/test_uop_graph.py +++ b/test/test_uop_graph.py @@ -453,7 +453,7 @@ class TestUOpGraph(unittest.TestCase): idx = d0.index(ridx0) ld = idx.load() val = (ridx0<50).where(5, ld) - st = idx.store(val, ridx0) + st = idx.store(val).end(ridx0) uops = to_uops_list([st]) for u in uops: assert u.op is not Ops.WHERE @@ -472,7 +472,7 @@ class TestUOpGraph(unittest.TestCase): c7 = UOp(Ops.DEFINE_GLOBAL, dtypes.uchar.ptr(60000), arg=2, src=()) c8 = c7.index(c6).load() c9 = ((c4<0).where((c4+60000), c4)!=c6.cast(dtypes.int)).where(0, c8.cast(dtypes.uint).cast(dtypes.uchar)).reduce(c5, arg=Ops.ADD) - c10 = c0.index(((c1*UOp.const(dtypes.index, 250))+c2)).store(c9, c1, c2) + c10 = c0.index(((c1*UOp.const(dtypes.index, 250))+c2)).store(c9).end(c1, c2) ast = c10.sink() uops = to_uops_list([ast]) for u in uops: diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index 1c1f8267f9..b2b6b9aa94 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -14,7 +14,7 @@ from tinygrad.codegen.late.devectorizer import load_store_folding, load_store_in from tinygrad.codegen.opt.postrange import apply_opts 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 -from tinygrad.codegen.late.control_flow import CFGContext, pm_add_ends, pm_split_ends, pm_add_control_flow, linearize +from tinygrad.codegen.late.control_flow import CFGContext, pm_split_ends, pm_add_control_flow, linearize def full_rewrite_to_sink(sink:UOp, ren:Renderer|None=None, optimize:bool=True) -> UOp: if ren is None: ren = Renderer() @@ -57,9 +57,6 @@ def full_rewrite_to_sink(sink:UOp, ren:Renderer|None=None, optimize:bool=True) - # 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") - # add ends (after reduces are removed, as long as we have reduces we can have stores) - sink = graph_rewrite(sink, pm_add_ends, name="add ends of ranges") - # devectorize (TODO: does this need opts?) if DEVECTORIZE >= 2: pm_devectorize = sym+load_store_folding+load_store_indexing elif DEVECTORIZE: pm_devectorize = sym+devectorize+load_store_folding+correct_load_store+load_store_indexing diff --git a/tinygrad/codegen/late/control_flow.py b/tinygrad/codegen/late/control_flow.py index a42ff3ac21..a2a29cf49e 100644 --- a/tinygrad/codegen/late/control_flow.py +++ b/tinygrad/codegen/late/control_flow.py @@ -100,13 +100,12 @@ pm_add_control_flow = PatternMatcher([ (UPat(Ops.RANGE, name="x"), lambda ctx,x: x.replace(src=x.src+(y,)) if (y:=ctx.edges.get(x)) is not None else None), ]) +def do_split_ends(e:UOp): + ret = e.src[0] + for r in list(UOp.sink(*e.src[1:]).ranges)[::-1]: ret = ret.end(r) + return ret + pm_split_ends = PatternMatcher([ # split the ends - (UPat(Ops.END, name="e"), lambda e: e.src[0].end(e.src[-1]).end(*e.src[1:-1]) if len(e.src) > 2 else None), + (UPat(Ops.END, name="e"), do_split_ends), ]) - -# NOTE: this can be done whenever -pm_add_ends = PatternMatcher([ - # put the end on the store - (UPat(Ops.STORE, name="s"), lambda s: s.replace(src=s.src[:2]).end(*[x for x in s.src[2:] if x.op is Ops.RANGE])), -]) \ No newline at end of file diff --git a/tinygrad/codegen/late/devectorizer.py b/tinygrad/codegen/late/devectorizer.py index f5f76a28c7..e36c056ef8 100644 --- a/tinygrad/codegen/late/devectorizer.py +++ b/tinygrad/codegen/late/devectorizer.py @@ -288,7 +288,7 @@ def reduce_to_acc(ctx:ReduceContext, red:UOp): # if we have a range if len(reduce_range) != 0: topo = inp.toposort() - ended_ranges = flatten([x.ended_ranges for x in topo if x.op is Ops.STORE]) + 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 reduce_range and x not in ended_ranges]) identity = red.const(red.dtype, identity_element(red.arg, red.dtype.scalar())) acc = UOp(Ops.DEFINE_REG, red.dtype.ptr(size=1, addrspace=AddrSpace.REG), arg=(ctx.acc_num,)) @@ -298,7 +298,7 @@ def reduce_to_acc(ctx:ReduceContext, red:UOp): ctx.acc_num += 1 ret = functools.reduce(lambda x,y: x.alu(red.arg, y), lst) if len(reduce_range) == 0: return ret - return acc.after(acc.index(UOp.const(dtypes.int, 0)).store(ret, *reduce_range)).index(UOp.const(dtypes.int, 0)).load() + return acc.after(acc.index(UOp.const(dtypes.int, 0)).store(ret).end(*reduce_range)).index(UOp.const(dtypes.int, 0)).load() pm_reduce = PatternMatcher([ # REDUCE -> DEFINE_ACC+ASSIGN diff --git a/tinygrad/codegen/opt/postrange.py b/tinygrad/codegen/opt/postrange.py index 269103134c..5f07b85101 100644 --- a/tinygrad/codegen/opt/postrange.py +++ b/tinygrad/codegen/opt/postrange.py @@ -4,7 +4,7 @@ from collections import defaultdict from typing import cast, Final from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, KernelInfo, graph_rewrite, AxisType, ssimplify, GroupOp from tinygrad.device import Buffer -from tinygrad.dtype import dtypes, ImageDType, AddrSpace +from tinygrad.dtype import dtypes, ImageDType from tinygrad.helpers import colored, BEAM, getenv, DEBUG, to_function_name, NOOPT, argsort, round_up, prod, merge_dicts, get_single_element, flatten from tinygrad.codegen.opt import axis_colors, Opt, OptOps, KernelOptError, check, axis_letters from tinygrad.codegen.simplify import pm_flatten_range @@ -64,19 +64,7 @@ class Scheduler: return self.ast.replace(arg=KernelInfo(name=name, applied_opts=tuple(self.applied_opts), dont_use_locals=self.dont_use_locals), tag=1) def _globalizable_rngs(self) -> list[UOp]: - store_rngs = self.ast.src[0].src[2:] - # filter any not in local stores - local_store_rngs = [x.ranges for x in self.ast.toposort() if (x.op is Ops.STORE and x.src[0].ptrdtype.addrspace == AddrSpace.LOCAL) \ - or (x.op is Ops.BUFFERIZE and x.arg == AddrSpace.LOCAL)] - for ls in local_store_rngs: store_rngs = tuple([x for x in store_rngs if x in ls]) - - # filter any not in reduces - # TODO: enable this - """ - reduce_rngs = [x.ranges for x in self.ast.toposort() if x.op is Ops.REDUCE] - for ls in reduce_rngs: store_rngs = tuple([x for x in store_rngs if x in ls]) - """ - return [x for x in UOp.sink(*store_rngs).toposort() if x.op is Ops.RANGE and x.arg[-1] == AxisType.LOOP] if store_rngs else [] + return flatten([list(UOp.sink(*s.src[1:]).ranges) for s in self.ast.src if s.op is Ops.END]) def convert_loop_to_global(self): if not self.ren.has_local: return None @@ -87,7 +75,7 @@ class Scheduler: self.ast = self.ast.substitute(dict(zip(self.rngs, rng))) def colors(self) -> list[str]: - output_rngs = flatten([list(UOp.sink(*s.src[2:]).ranges) for s in self.ast.src]) + output_rngs = self._globalizable_rngs() ret = [] for x,r in zip(self.axis_types, self.rngs): if self.dont_use_locals and x == AxisType.GLOBAL: ret.append("BLUE") diff --git a/tinygrad/codegen/simplify.py b/tinygrad/codegen/simplify.py index a61762c2dd..c8feeb65b2 100644 --- a/tinygrad/codegen/simplify.py +++ b/tinygrad/codegen/simplify.py @@ -12,7 +12,7 @@ def flatten_range(r:UOp): pm_flatten_range = PatternMatcher([ # real ranges only - (UPat((Ops.REDUCE, Ops.STORE), name="r"), flatten_range), + (UPat((Ops.REDUCE, Ops.STORE, Ops.END), name="r"), flatten_range), ]) def count_divmod(x:UOp): return len([u for u in x.toposort() if u.op in {Ops.IDIV, Ops.MOD}]) @@ -39,7 +39,7 @@ def simplify_merge_adjacent(u:UOp) -> UOp|None: return u pm_simplify_ranges = PatternMatcher([ - (UPat((Ops.STORE, Ops.REDUCE), name="u"), simplify_merge_adjacent), + (UPat((Ops.END, Ops.REDUCE), name="u"), simplify_merge_adjacent), ]) def mark_range_mod(ctx, r:UOp, c:UOp): diff --git a/tinygrad/schedule/rangeify.py b/tinygrad/schedule/rangeify.py index d28cfdc4af..4069b7890a 100644 --- a/tinygrad/schedule/rangeify.py +++ b/tinygrad/schedule/rangeify.py @@ -311,7 +311,7 @@ def bufferize_to_store(x:UOp, allow_locals=True): assert assign_target.op is Ops.INDEX, f"{assign_target.op} is not index" # in assign, this is the buffer size, not the bufferize size # TODO: assign_mops here - do_store = assign_target.replace(dtype=sdtype).store(assign_src, *rngs).replace(tag=x.tag) + do_store = assign_target.replace(dtype=sdtype).store(assign_src, tag=x.tag).end(*[x for x in rngs if x.op is Ops.RANGE]) ret = assign_target.src[0].after(do_store) mops = [] walk = assign_mops @@ -324,7 +324,7 @@ def bufferize_to_store(x:UOp, allow_locals=True): # NOTE: the DEFINE_LOCAL needs to be disambiguated here if sdtype.addrspace == AddrSpace.GLOBAL: buf = UOp.new_buffer(x.arg.device, size, x.dtype) - do_store = buf.reshape(shape).index(*rngs, dtype=sdtype).store(x.src[0], *rngs).replace(tag=x.tag) + do_store = buf.reshape(shape).index(*rngs, dtype=sdtype).store(x.src[0], tag=x.tag).end(*[x for x in rngs if x.op is Ops.RANGE]) ret = buf.after(do_store).forced_reshape(shape) # TODO: is this right? what if it's offset if any(r.op is Ops.RANGE and r.src[0].op is not Ops.CONST for r in rngs): @@ -337,7 +337,7 @@ def bufferize_to_store(x:UOp, allow_locals=True): tag = x.arg.device if tag is None: tag = UOp.unique().arg # TODO: hack buf = UOp(Ops.DEFINE_LOCAL, sdtype, arg=tag) - do_store = buf.reshape(shape).index(*rngs, dtype=sdtype).store(x.src[0], *rngs) + do_store = buf.reshape(shape).index(*rngs, dtype=sdtype).store(x.src[0]).end(*[x for x in rngs if x.op is Ops.RANGE]) return buf.after(do_store.barrier()).reshape(shape) pm_add_buffers = pm_mops+to_bufferview+PatternMatcher([ @@ -477,7 +477,7 @@ def split_store(ctx:list[UOp], x:UOp) -> UOp|None: return kernel split_kernels = PatternMatcher([ - (UPat(Ops.STORE, name="x"), split_store), + (UPat((Ops.STORE, Ops.END), name="x"), split_store), ]) def tag_uop(ctx:list[UOp], x:UOp): diff --git a/tinygrad/viz/serve.py b/tinygrad/viz/serve.py index 00975704fb..dc4b1498d9 100755 --- a/tinygrad/viz/serve.py +++ b/tinygrad/viz/serve.py @@ -84,7 +84,7 @@ def uop_to_json(x:UOp, ignore_indexing=False) -> dict[int, dict]: label += f"\n{shape_to_str(u.shape)}" if u.op in {Ops.INDEX, Ops.BUFFERIZE}: label += f"\n{u.render()}" - if u.op in {Ops.END, Ops.STORE, Ops.REDUCE} and len(trngs:=list(UOp.sink(*u.src[range_start[u.op]:]).ranges)): + if u.op in {Ops.END, Ops.REDUCE} and len(trngs:=list(UOp.sink(*u.src[range_start[u.op]:]).ranges)): label += "\n"+' '.join([f"{colored(s.arg[0], axis_colors[s.arg[-1]])}({s.vmax+1})" for s in trngs]) except Exception: label += "\n"