From dcc6ddf0eb29fca24afa4f1baf21618634cb45cf Mon Sep 17 00:00:00 2001 From: George Hotz Date: Tue, 5 Aug 2025 18:24:58 -0700 Subject: [PATCH] that hack broke things --- tinygrad/schedule/grouper.py | 2 +- tinygrad/schedule/kernelize.py | 4 +++- tinygrad/uop/ops.py | 2 +- 3 files changed, 5 insertions(+), 3 deletions(-) diff --git a/tinygrad/schedule/grouper.py b/tinygrad/schedule/grouper.py index 2477da5263..66063c0246 100644 --- a/tinygrad/schedule/grouper.py +++ b/tinygrad/schedule/grouper.py @@ -3,7 +3,7 @@ from tinygrad.helpers import all_int, prod, unwrap, dedup, DONT_REALIZE_EXPAND, from tinygrad.shape.shapetracker import ShapeTracker ALWAYS_CONTIGUOUS = {Ops.CONTIGUOUS, Ops.ASSIGN, Ops.COPY, Ops.BUFFER, Ops.BUFFER_VIEW, - Ops.CONST, Ops.BIND, Ops.DEVICE, Ops.MSELECT, Ops.MSTACK} + Ops.CONST, Ops.BIND, Ops.DEVICE, Ops.MSELECT, Ops.MSTACK, Ops.DEFINE_GLOBAL} # **** Grouper decides which of the UOps realize diff --git a/tinygrad/schedule/kernelize.py b/tinygrad/schedule/kernelize.py index f185120db5..5a7ade255c 100644 --- a/tinygrad/schedule/kernelize.py +++ b/tinygrad/schedule/kernelize.py @@ -152,7 +152,7 @@ create_kernels = PatternMatcher([ early_buffer_ops = PatternMatcher([ # LOAD - (UPat(Ops.BUFFER, name="x"), lambda ctx,x: UOp.load(UOp(Ops.DEFINE_GLOBAL, x.dtype.ptr(x.size), (), ctx.index(x)).view(x.st),)), + (UPat(Ops.BUFFER, name="x"), lambda ctx,x: UOp(Ops.DEFINE_GLOBAL, x.dtype.ptr(x.size), (), ctx.index(x), tag=1)), # no SINK for meta ops (UPat(Ops.SINK, src=(UPat(Ops.CONTIGUOUS, src=(UPat(GroupOp.Meta, name="x"),),))), lambda x:x), ]) @@ -168,6 +168,8 @@ def check_load_st(glbl:UOp, view:UOp): +colored(" - a += a.T\n", "red")+colored(" + a += a.T.contiguous()", "green")) fix_kernel_ops = PatternMatcher([ + # add the LOAD + (UPat(Ops.DEFINE_GLOBAL, name="x"), lambda x: x.replace(tag=None).view(x.st).load() if x.tag is not None else None), # 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/uop/ops.py b/tinygrad/uop/ops.py index 1dfdaeb857..a26e3ef4c8 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -155,7 +155,7 @@ class UOp(MathTrait, metaclass=UOpMetaClass): return ShapeTracker.from_shape((sz,)) if sz > 0 else None # hack for PTX, CASTing the ptr loses the shape - if self.op is Ops.CAST and self.src[0].op is Ops.DEFINE_GLOBAL: return None + #if self.op is Ops.CAST and self.src[0].op is Ops.DEFINE_GLOBAL: return None # otherwise we get the shape from sources if not (src_sts := [x.st for x in self.src if x.st is not None]): return None