From ea5caefef0918f26ce5ead8f3eaf186727aaa767 Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Wed, 2 Apr 2025 18:10:57 +0800 Subject: [PATCH] gep should look at count, not vcount (#9698) * gep should look at count, not vcount * gep in order is a rule * min change * gep on void --- tinygrad/codegen/symbolic.py | 4 ++++ tinygrad/ops.py | 1 - 2 files changed, 4 insertions(+), 1 deletion(-) diff --git a/tinygrad/codegen/symbolic.py b/tinygrad/codegen/symbolic.py index df95f8d8ec..70461cda81 100644 --- a/tinygrad/codegen/symbolic.py +++ b/tinygrad/codegen/symbolic.py @@ -180,6 +180,10 @@ gep_pushing = PatternMatcher([ lambda gep, vec: UOp(Ops.VECTORIZE, gep.dtype, tuple(vec.src[i] for i in gep.arg)) if len(gep.arg) > 1 else vec.src[gep.arg[0]]), (UPat(Ops.GEP, src=(UPat.cvar("c", vec=False),), name="gep"), lambda gep, c: gep.const_like(c.arg)), (UPat(Ops.GEP, src=(UPat(Ops.VCONST, name="c"),), name="gep"), lambda gep, c: gep.const_like(tuple(c.arg[x] for x in gep.arg))), + # GEP on void is skipped + (UPat(Ops.GEP, src=(UPat(dtype=dtypes.void, name="x"),)), lambda x: x), + # GEP in order is removed + (UPat(Ops.GEP, name="g"), lambda g: g.src[0] if not isinstance(g.dtype, PtrDType) and g.arg == tuple(range(g.src[0].dtype.count)) else None), # push all GEPs through ALUs (fix arange stuff) (UPat(Ops.GEP, src=(UPat((*GroupOp.ALU, Ops.CAST, Ops.BITCAST), name='alu'),), name='gep'), lambda gep,alu: UOp(alu.op, alu.dtype.scalar().vec(gep.dtype.count), tuple(x.gep(gep.arg) for x in alu.src), alu.arg) \ diff --git a/tinygrad/ops.py b/tinygrad/ops.py index ed007dc290..c3541e4130 100644 --- a/tinygrad/ops.py +++ b/tinygrad/ops.py @@ -387,7 +387,6 @@ class UOp(MathTrait, metaclass=UOpMetaClass): if self.op is Ops.VCONST: return UOp.const(self.dtype.scalar(), self.arg[i]) if self.op is Ops.CONST: return UOp.const(self.dtype.scalar(), self.arg) i = (i,) - if (self.dtype.vcount == len(i) and i == tuple(range(len(i)))) or self.dtype == dtypes.void: return self 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, src=(self,)+src, **kwargs) def store(self, *src:UOp, **kwargs): return UOp(Ops.STORE, dtypes.void, (self,)+src, **kwargs)