From 8305a5804cdbb74cfbb758cc0edad12db22ac44a Mon Sep 17 00:00:00 2001 From: George Hotz Date: Tue, 22 Jul 2025 18:52:26 -0700 Subject: [PATCH] just the range thing --- tinygrad/codegen/devectorizer.py | 11 +++-------- tinygrad/codegen/linearize.py | 10 +--------- tinygrad/renderer/cstyle.py | 3 +-- tinygrad/uop/ops.py | 4 ++-- tinygrad/uop/spec.py | 9 ++++++--- tinygrad/uop/symbolic.py | 2 +- 6 files changed, 14 insertions(+), 25 deletions(-) diff --git a/tinygrad/codegen/devectorizer.py b/tinygrad/codegen/devectorizer.py index 58a9ff0b6e..2675bb2e7a 100644 --- a/tinygrad/codegen/devectorizer.py +++ b/tinygrad/codegen/devectorizer.py @@ -113,7 +113,6 @@ def cat_after_store(cat:UOp, data:UOp, sto:UOp): for s in cat.src: ret.append(s.store(data.gep(tuple(range(offset, offset+s.dtype.count))), *sto.src[2:])) offset += s.dtype.count - return UOp.sink(*ret) # dtype CAT dtypes: list[PtrDType] = [x.dtype for x in ret if isinstance(x.dtype, PtrDType)] assert len(dtypes) == len(ret) and all_same([(x.size, x.addrspace) for x in dtypes]) @@ -335,15 +334,11 @@ def reduce_to_acc(ctx:ReduceContext, red:UOp): assert all(x.dtype == red.dtype for x in lst), f"horizontal reduction mismatch {lst[0].dtype} != {red.dtype}" # if we have a range if len(reduce_range) != 0: - #acc = UOp(Ops.DEFINE_REG, red.dtype.ptr(size=1, addrspace=AddrSpace.REG), - # (red.const_like(identity_element(red.arg, red.dtype.scalar())),) + tuple(reduce_range), (ctx.acc_num,)).index(UOp.const(dtypes.int, 0)) - reduce_start = red.const_like(identity_element(red.arg, red.dtype.scalar())) - is_start = functools.reduce(lambda x,y: x&y, [x.eq(x.const_like(0)) for x in reduce_range]).broadcast(red.dtype.count) - acc = UOp(Ops.DEFINE_REG, red.dtype.ptr(size=1, addrspace=AddrSpace.REG), (reduce_start,), (ctx.acc_num,)).index(UOp.const(dtypes.int, 0)) - lst = [is_start.where(reduce_start, acc.load(*reduce_range))] + lst # put acc as the first element + acc = UOp(Ops.DEFINE_REG, red.dtype.ptr(size=1, addrspace=AddrSpace.REG), + (red.const_like(identity_element(red.arg, red.dtype.scalar())),) + tuple(reduce_range), (ctx.acc_num,)).index(UOp.const(dtypes.int, 0)) ctx.acc_num += 1 ret = functools.reduce(lambda x,y: x.alu(red.arg, y), lst) - return acc.load(acc.store(ret, *reduce_range)) if len(reduce_range) != 0 else ret + return acc.store(ret, *reduce_range).load() if len(reduce_range) != 0 else ret def no_vectorized_reduce(inp:UOp, red:UOp): if inp.dtype != red.dtype: diff --git a/tinygrad/codegen/linearize.py b/tinygrad/codegen/linearize.py index 1b84818012..0581e7b386 100644 --- a/tinygrad/codegen/linearize.py +++ b/tinygrad/codegen/linearize.py @@ -92,15 +92,7 @@ class BlockContext: # RANGE/IF add to the next ctx # STORE/ASSIGN subtract from the next ctx if u.op in {Ops.RANGE, Ops.IF}: ctx.child_ctxs[u] = _sort_ctx(ctx.block_ctxs[u] + (u,)) - elif u.op is Ops.STORE: - if len(definereg:=[x for x in u.src[0].toposort() if x.op is Ops.DEFINE_REG]): - # old assign logic - ctx.child_ctxs[u] = tuple([y for y in ctx.last_ctx(u.src[1]) if y not in definereg[0].src[1:]]) - elif any(x.op is Ops.DEFINE_LOCAL for x in u.src[0].toposort()): - # deal with non-reduce locals. probably wrong - idx_context, store_context = ctx.last_ctx(u.src[0]), ctx.last_ctx(u.src[1]) - ctx.child_ctxs[u] = tuple([y for y in store_context if y not in idx_context and y.op is Ops.RANGE]) - else: ctx.child_ctxs[u] = () + elif u.op is Ops.STORE: ctx.child_ctxs[u] = tuple([x for x in ctx.block_ctxs[u] if x not in u.src]) return ctx # ***** make blocks ***** diff --git a/tinygrad/renderer/cstyle.py b/tinygrad/renderer/cstyle.py index 1c4e9ada10..5027faeac5 100644 --- a/tinygrad/renderer/cstyle.py +++ b/tinygrad/renderer/cstyle.py @@ -167,8 +167,7 @@ class CStyleLanguage(Renderer): r[u] = l else: if u.op in {Ops.RANGE, Ops.DEFINE_LOCAL, Ops.STORE, Ops.DEFINE_REG} or u.dtype == dtypes.void: - #if u.op is Ops.STORE: r[u] = r[u.src[0]] - pass + if u.op is Ops.STORE: r[u] = r[u.src[0]] else: l = f"{self.render_dtype(u.dtype)} {r[u]} = {l}" + (";" if u.op is not Ops.SPECIAL else "") kernel.append(" "*depth + l) diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index c705d745ae..8d86938f0a 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -233,7 +233,7 @@ class UOp(MathTrait, metaclass=UOpMetaClass): i = (i,) return UOp(Ops.GEP, self.dtype.scalar().vec(len(i)) if len(i) > 1 else self.dtype.scalar(), (self,), i) def load(self, *src:UOp, **kwargs): return UOp(Ops.LOAD, dtype=kwargs.pop("dtype", self.dtype.base), src=(self,)+src, **kwargs) - def store(self, *src:UOp, **kwargs): return UOp(Ops.STORE, dtypes.void, (self,)+src, **kwargs) + def store(self, *src:UOp, **kwargs): return UOp(Ops.STORE, self.dtype, (self,)+src, **kwargs) def assign(self, x:UOp): return UOp(Ops.ASSIGN, self.dtype, (self, x)) def alu(self, arg, *src:UOp): out_dtype = (self, *src)[-1].dtype @@ -645,7 +645,7 @@ class UPat(MathTrait): def bitcast(self, dtype=None): return UPat(Ops.BITCAST, dtype, (self,)) def gep(self, i:int|None=None, **kwargs): return UPat(Ops.GEP, None, (self,), (i,) if i is not None else None, **kwargs) def load(self, *src:UPat, **kwargs): return UPat(Ops.LOAD, src=(self,)+src, **kwargs) - def store(self, *src:UPat, **kwargs): return UPat(Ops.STORE, dtypes.void, (self,)+src, **kwargs) + def store(self, *src:UPat, **kwargs): return UPat(Ops.STORE, self.dtype, (self,)+src, **kwargs) def assign(self, x:UPat, **kwargs): return UPat(Ops.ASSIGN, self.dtype, (self,x), **kwargs) def reduce(self, *src:UPat, **kwargs): return UPat(Ops.REDUCE, self.dtype, src=(self,)+src, **kwargs) def fuse(self): return self.alu(Ops.FUSE) diff --git a/tinygrad/uop/spec.py b/tinygrad/uop/spec.py index 16f8c7a22b..5d9750815f 100644 --- a/tinygrad/uop/spec.py +++ b/tinygrad/uop/spec.py @@ -158,15 +158,18 @@ spec = PatternMatcher([ (UPat(Ops.INDEX, src=(UPat((Ops.DEFINE_GLOBAL, Ops.DEFINE_LOCAL, Ops.DEFINE_REG)), UPat())), lambda: True), (UPat(Ops.INDEX, src=(UPat((Ops.DEFINE_GLOBAL, Ops.DEFINE_LOCAL, Ops.DEFINE_REG)), UPat(), UPat(dtype=dtypes.bool))), lambda: True), + # LOAD on STORE + (UPat(Ops.LOAD, src=(UPat(Ops.STORE),)), lambda: True), + # LOAD takes a (UPat(Ops.LOAD, src=(index_pat,)), validate_index), - (UPat(Ops.LOAD, src=(index_pat, UPat((Ops.BARRIER, Ops.STORE)))), validate_index), + (UPat(Ops.LOAD, src=(index_pat, UPat(Ops.BARRIER))), validate_index), (UPat(Ops.LOAD, src=(index_pat, UPat(Ops.IF, name="cond"))), lambda idx,cond: validate_index(idx,cond.src[0])), (UPat(Ops.LOAD, src=(index_pat, UPat.var("alt")), name="ld"), lambda ld,alt,idx: ld.dtype == alt.dtype and validate_index(idx)), # STORE takes a - (UPat(Ops.STORE, dtypes.void, src=(index_pat, UPat(name="val"), UPat(Ops.IF, name="gate")), allow_any_len=True), validate_store), - (UPat(Ops.STORE, dtypes.void, src=(index_pat, UPat(name="val")), allow_any_len=True), validate_store), + (UPat(Ops.STORE, src=(index_pat, UPat(name="val"), UPat(Ops.IF, name="gate")), allow_any_len=True), validate_store), + (UPat(Ops.STORE, src=(index_pat, UPat(name="val")), allow_any_len=True), validate_store), # most ALUs have all matching dtypes, except CMPLT, CMPNE, and WHERE (UPat(Ops.WHERE, name="w", src=(UPat(dtype=dtypes.bool), UPat.var("x"), UPat.var("y"))), lambda w,x,y: w.dtype == x.dtype == y.dtype), diff --git a/tinygrad/uop/symbolic.py b/tinygrad/uop/symbolic.py index 8469da0055..e63fe5e464 100644 --- a/tinygrad/uop/symbolic.py +++ b/tinygrad/uop/symbolic.py @@ -437,7 +437,7 @@ sym = symbolic_flat+PatternMatcher([ ((UPat.var('x', dtypes.uint64)&(UPat.var('y').where(UPat.const(dtypes.uint64, 0xFFFFFFFF), UPat.const(dtypes.uint64, 0)))).cast(dtypes.uint32), lambda x,y: y.where(x.cast(dtypes.uint32), UOp.const(dtypes.uint32, 0))), # ** self folding ** - #(UPat(Ops.DEFINE_REG, src=(UPat.var("x"),)), lambda x: x), # a DEFINE_ACC without ranges is a CONST + (UPat(Ops.DEFINE_REG, src=(UPat.var("x"),)), lambda x: x), # a DEFINE_ACC without ranges is a CONST # x!=0 -> (bool)x (UPat.var("x")!=0, lambda x: x.cast(dtypes.bool.vec(x.dtype.count))), # ** where **