diff --git a/tinygrad/opt/swizzler.py b/tinygrad/opt/swizzler.py index f4f82ad075..9c7047b27a 100644 --- a/tinygrad/opt/swizzler.py +++ b/tinygrad/opt/swizzler.py @@ -95,7 +95,7 @@ view_right = merge_views+PatternMatcher([ # apply view after reduceops (UPat(Ops.REDUCE_AXIS, src=(UPat(Ops.VIEW, src=(UPat(GroupOp.All-ALWAYS_CONTIGUOUS, name="src"),), name="v"),), name="r"), reduceop_view_right), # apply view after elementwise ops - (UPat(GroupOp.All-{Ops.SINK, Ops.REDUCE_AXIS}, name="root"), elementwise_view_right), + (UPat(GroupOp.All-{Ops.SINK, Ops.REDUCE_AXIS, Ops.LOAD, Ops.STORE}, name="root"), elementwise_view_right), # merge axes for double reduce (invert of SPLIT_REDUCEOP=1) (UPat(Ops.REDUCE_AXIS, src=(UPat(Ops.REDUCE_AXIS, name="r1"),), name="r2"), lambda r1,r2: r1.replace(arg=(r1.arg[0], r2.arg[1]+r1.arg[1])) if r1.arg[0] is r2.arg[0] else None), diff --git a/tinygrad/schedule/kernelize.py b/tinygrad/schedule/kernelize.py index d35200adb0..5ef58bd6ec 100644 --- a/tinygrad/schedule/kernelize.py +++ b/tinygrad/schedule/kernelize.py @@ -193,10 +193,6 @@ replace_globals = PatternMatcher([ def fix_kernel_ast(k:UOp) -> UOp|None: if k.arg.ast.op in GroupOp.Meta or all(s.op is Ops.STORE for s in k.arg.ast.src): return None - # replace global memory ops with the BUFFER they write to - ast = graph_rewrite(k.arg.ast, replace_globals, bottom_up=True, name="replace globals") - # push views to edges - ast = graph_rewrite(graph_rewrite(ast, view_left, name="Main View Left"), view_right, name="Main View Right") # replace buffer with define_global + add load/store last bufs = [] for s in k.src: @@ -204,6 +200,11 @@ def fix_kernel_ast(k:UOp) -> UOp|None: # traverse back through MSELECT and MSTACK. HACK: 0 branch of MSTACK only while s.op in {Ops.MSELECT, Ops.MSTACK}: s = s.src[0] bufs.append(s) + # replace global memory ops with the BUFFER they write to + ast = graph_rewrite(k.arg.ast, replace_globals, bufs, bottom_up=True, name="replace globals") + # TODO: move these to codegen + ast = graph_rewrite(ast, view_left, name="Main View Left") + ast = graph_rewrite(ast, view_right, name="Main View Right") ast = graph_rewrite(ast, view_left+add_buffer_ops+fix_kernel_ops, bufs, bottom_up=True, name="replace buffer") if ast.op is Ops.SINK and not all_same([x.device for x in k.src]): raise RuntimeError(f"all buffers must be on the same device: {tuple(b.buf_uop.buffer for b in k.src)}")