From 0f7e296f5b8ccdfbb0c40d6e1aee976a89a96b8e Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Thu, 30 Apr 2026 08:05:30 -0700 Subject: [PATCH] fix some indexing edge cases (#15988) --- test/null/test_uop_vmin_vmax.py | 2 +- tinygrad/codegen/late/devectorizer.py | 8 ++++++-- tinygrad/uop/symbolic.py | 5 +++-- 3 files changed, 10 insertions(+), 5 deletions(-) diff --git a/test/null/test_uop_vmin_vmax.py b/test/null/test_uop_vmin_vmax.py index 08a23cce34..a341f78f1e 100644 --- a/test/null/test_uop_vmin_vmax.py +++ b/test/null/test_uop_vmin_vmax.py @@ -297,7 +297,7 @@ class TestVminVmaxVConst(unittest.TestCase): # vmin and vmax for a vector constant of bool values d1 = UOp(Ops.PARAM, dtypes.int.ptr(), (), 1) idx = UOp.const(dtypes.int, 0) - val = UOp(Ops.LOAD, dtypes.int.vec(2), (d1.index(idx),)) + val = UOp(Ops.LOAD, dtypes.int.vec(2), (d1.index(idx).cast(dtypes.int.vec(2).ptr()),)) uop = (val // 32).gep(0) self.assertEqual(uop.vmin, -67108864) self.assertEqual(uop.vmax, 67108863) diff --git a/tinygrad/codegen/late/devectorizer.py b/tinygrad/codegen/late/devectorizer.py index 54b09f1c6c..dd4c4556e8 100644 --- a/tinygrad/codegen/late/devectorizer.py +++ b/tinygrad/codegen/late/devectorizer.py @@ -358,10 +358,14 @@ pm_reduce = PatternMatcher([ # add loads +def add_load(idx:UOp): + if isinstance(idx.dtype, PtrDType): return None + assert isinstance(idx.src[0].dtype, PtrDType), f"param is not PtrDType {idx.src[0].dtype}" + return idx.replace(dtype=idx.src[0].dtype).load(dtype=idx.dtype.base) + pm_add_loads = PatternMatcher([ # add loads to non ptr index - (UPat(Ops.INDEX, name="idx"), lambda idx: None if isinstance(idx.dtype, PtrDType) else - idx.replace(dtype=idx.src[0].dtype).load(dtype=idx.dtype.base)), + (UPat(Ops.INDEX, name="idx"), add_load), # remove loads from stores (UPat(Ops.STORE, src=(UPat(Ops.LOAD), UPat(name="val")), name="s"), lambda s,val: s.replace(src=(s.src[0].src[0], val))), ]) diff --git a/tinygrad/uop/symbolic.py b/tinygrad/uop/symbolic.py index afdc6f85fb..9fb519dfbb 100644 --- a/tinygrad/uop/symbolic.py +++ b/tinygrad/uop/symbolic.py @@ -445,8 +445,9 @@ sym = symbolic+pm_simplify_valid+PatternMatcher([ UPat.load(UPat(Ops.INDEX, name="index"))), allow_any_len=True, name="store"), lambda index, gate, alt, store: UOp.store(index.src[0].index(gate.where(index.src[1], UOp.invalid())), alt, *store.src[2:])), # fold gated LOAD/STORE - (UPat((Ops.LOAD, Ops.STORE), src=(UPat().index(UPat.const(dtypes.weakint, Invalid)).or_casted(),), allow_any_len=True, name="x"), - lambda x: UOp(Ops.NOOP) if x.op is Ops.STORE else x.const_like(0)), # invalid store does nothing. invalid load produces 0 + (UPat(Ops.STORE, src=(UPat().index(UPat.const(dtypes.weakint, Invalid)).or_casted(),), allow_any_len=True, name="x"), lambda x: UOp(Ops.NOOP)), + (UPat(Ops.LOAD, src=(UPat().index(UPat.const(dtypes.weakint, Invalid)).or_casted(),), allow_any_len=True, name="x"), + lambda x: x.src[1] if len(x.src) > 1 else x.const_like(0)), # invalid load produces 0, or the alt value if we have one (UPat(Ops.STORE, src=(UPat(), invalid_pat), allow_any_len=True), lambda i: UOp(Ops.NOOP)), # store of where with invalid -> gated store (UPat(Ops.STORE, src=(UPat(Ops.INDEX, name="index"), UPat.var("cond").where(UPat.var("val"), invalid_pat)), allow_any_len=True, name="store"),