diff --git a/tinygrad/opt/swizzler.py b/tinygrad/opt/swizzler.py index 44e50549f4..6e3012af78 100644 --- a/tinygrad/opt/swizzler.py +++ b/tinygrad/opt/swizzler.py @@ -119,6 +119,8 @@ def check_load_st(glbl:UOp, view:UOp): +colored(" - a += a.T\n", "red")+colored(" + a += a.T.contiguous()", "green")) fix_kernel_ops = view_left_through_load+PatternMatcher([ + # add view to LOAD + (UPat(Ops.DEFINE_GLOBAL, name="g").load(), lambda g: g.view(g.st).load()), # STORE (except for meta ops) (UPat(Ops.SINK, src=UPat(GroupOp.All-{Ops.STORE}), name="sink"), lambda sink: UOp.sink(*[UOp.store(UOp(Ops.DEFINE_GLOBAL, (s:=x.base).dtype.ptr(s.st.real_size()), (), i).view(s.st), s) for i,x in enumerate(sink.src)])), diff --git a/tinygrad/schedule/kernelize.py b/tinygrad/schedule/kernelize.py index 8ad6db77dc..7c2468ae31 100644 --- a/tinygrad/schedule/kernelize.py +++ b/tinygrad/schedule/kernelize.py @@ -151,7 +151,7 @@ create_kernels = PatternMatcher([ early_buffer_ops = PatternMatcher([ # LOAD - (UPat(Ops.BUFFER, name="x"), lambda ctx,x: UOp(Ops.DEFINE_GLOBAL, x.dtype.ptr(x.size), (), ctx.index(x)).view(x.st).load()), + (UPat(Ops.BUFFER, name="x"), lambda ctx,x: UOp(Ops.DEFINE_GLOBAL, x.dtype.ptr(x.size), (), ctx.index(x)).load()), # no SINK for meta ops (UPat(Ops.SINK, src=(UPat(Ops.CONTIGUOUS, src=(UPat(GroupOp.Meta, name="x"),),))), lambda x:x), ])