From fede358811ec40bcf10e11e3fc93d9ef55ebd943 Mon Sep 17 00:00:00 2001 From: chenyu Date: Thu, 27 Aug 2026 19:59:27 -0400 Subject: [PATCH] fix fancy indexing with uint8 index (#17787) * fix fancy indexing with uint8 index * fix --- test/backend/test_ops.py | 4 ++++ tinygrad/renderer/isa/x86.py | 2 +- 2 files changed, 5 insertions(+), 1 deletion(-) diff --git a/test/backend/test_ops.py b/test/backend/test_ops.py index d0263b9f5b..30f564d6d1 100644 --- a/test/backend/test_ops.py +++ b/test/backend/test_ops.py @@ -2963,6 +2963,10 @@ class TestOps(unittest.TestCase): data = [math.inf, -math.inf, math.nan] helper_test_op((), lambda: torch.tensor(data)[torch.tensor([0, 1, 2])], lambda: Tensor(data)[Tensor([0, 1, 2])]) + def test_fancy_indexing_index_dtypes(self): + helper_test_op((), lambda: torch.tensor([10., 20., 30., 40.])[torch.tensor([1, 2, 3, 0])], + lambda: Tensor([10., 20., 30., 40.])[Tensor([1, 2, 3, 0], dtype=dtypes.uint8)]) + @slow_test def test_slice_fancy_indexing_no_dim_collapse(self): a,b,c,d,e,i,j,k,o,p = self._get_index_randoms() diff --git a/tinygrad/renderer/isa/x86.py b/tinygrad/renderer/isa/x86.py index f90c8d5686..79a8a17945 100644 --- a/tinygrad/renderer/isa/x86.py +++ b/tinygrad/renderer/isa/x86.py @@ -282,7 +282,7 @@ def shift(x:UOp, op:X86Ops) -> UOp: # it is materialized as an immediate so the address stays correct if the base register is ever spilled and refilled def fold_address(x:UOp) -> tuple[UOp, UOp, UOp, UOp]: def _disp(v:int) -> UOp: return imm(dtypes.int32 if abs(v) > dtypes.int8.max else dtypes.int8, v) - def _cast(v:UOp) -> UOp: return v.cast(dtypes.int64) if v.vmin < 0 else v + def _cast(v:UOp) -> UOp: return v.cast(dtypes.int64) if v.vmin < 0 else v.cast(dtypes.uint32) if v.dtype.itemsize < 4 else v if x.op not in {Ops.INDEX, Ops.SHRINK}: return (x, UOp(Ops.NOOP), _disp(0), imm(dtypes.uint8, x.dtype.itemsize)) base, idx = x.src[0], x.src[1] # buffers are indexed by element, everything else (the stack pointer) by byte