diff --git a/test/imported/test_indexing.py b/test/imported/test_indexing.py index 51f69455bf..3bed0c47c7 100644 --- a/test/imported/test_indexing.py +++ b/test/imported/test_indexing.py @@ -110,19 +110,17 @@ class TestIndexing(unittest.TestCase): # pick a random valid indexer type def ri(indices): - # choice = random.randint(0, 2) - # TODO: we only support Tensor index now - choice = 0 + choice = random.randint(0, 1) + # TODO: we do not support tuple of list for index now if choice == 0: return Tensor(indices) if choice == 1: return list(indices) return tuple(indices) def validate_indexing(x): - # TODO: we only support Tensor index now - # numpy_testing_assert_equal_helper(x[[0]], consec((1,))) + numpy_testing_assert_equal_helper(x[[0]], consec((1,))) numpy_testing_assert_equal_helper(x[ri([0]),], consec((1,))) numpy_testing_assert_equal_helper(x[ri([3]),], consec((1,), 4)) - # numpy_testing_assert_equal_helper(x[[2, 3, 4]], consec((3,), 3)) + numpy_testing_assert_equal_helper(x[[2, 3, 4]], consec((3,), 3)) numpy_testing_assert_equal_helper(x[ri([2, 3, 4]),], consec((3,), 3)) numpy_testing_assert_equal_helper(x[ri([0, 2, 4]),], np.array([1, 3, 5])) diff --git a/test/test_ops.py b/test/test_ops.py index a378cac03d..082eefa863 100644 --- a/test/test_ops.py +++ b/test/test_ops.py @@ -1240,6 +1240,14 @@ class TestOps(unittest.TestCase): helper_test_op([(2,3)], lambda x: x[torch.tensor([[0,0,0],[0,0,0]]), torch.tensor(1)], lambda x: x[Tensor([[0,0,0],[0,0,0]]), Tensor(1)]) helper_test_op([(2,3)], lambda x: x[torch.tensor([1]), torch.tensor([[0,0,0],[0,0,0]])], lambda x: x[Tensor([1]), Tensor([[0,0,0],[0,0,0]])]) + def test_slice_fancy_indexing_list_indices(self): + a,b,c,d,e,i,j,k,o,p = self._get_index_randoms() + helper_test_op([(2,5,6,5,3,4)], lambda x: x[[0],b,c,d,:], lambda x: x[[0],j,k,o,:]) + helper_test_op([(2,5,6,5,3,4)], lambda x: x[[1],b,c,d,:], lambda x: x[[1],j,k,o,:]) + helper_test_op([(2,5,6,5,3,4)], lambda x: x[[1,0],b,c,d,:], lambda x: x[[1,0],j,k,o,:]) + helper_test_op([(2,5,6,5,3,4)], lambda x: x[a,b,c,[1,2,3],...], lambda x: x[i,j,k,[1,2,3],...]) + helper_test_op([(2,5,6,5,3,4)], lambda x: x[a,[2,1,0],c,[2,1,0],e], lambda x: x[i,[2,1,0],k,[2,1,0],p]) + def test_gather(self): # indices cannot have gradient # indices cannot be negative (torch gather) diff --git a/tinygrad/tensor.py b/tinygrad/tensor.py index 03824fdda6..8f428fa162 100644 --- a/tinygrad/tensor.py +++ b/tinygrad/tensor.py @@ -300,18 +300,22 @@ class Tensor: # - There's a special case where a permute is needed at the end: # - if first Tensor passed in (expand dims) is not at dim 0 # - and following Tensors does not follow consecutively to the end of fancy indexing's dims - def __getitem__(self, val) -> Tensor: # val: Union[int, slice, Tensor, None, Ellipsis, Tuple[Union[int, slice, Tensor, None, Ellipsis], ...]] + def __getitem__(self, indices) -> Tensor: # indices: Union[int, slice, Tensor, None, Ellipsis, List, Tuple[Union[int, slice, Tensor, None, Ellipsis], ...]] def normalize_int(e, i, dim_sz): if -dim_sz <= e < dim_sz: return e if e != -1 else dim_sz-1 raise IndexError(f"index {e} is out of bounds for dimension {i} with size {self.shape[i]}") - orig_slices = list(val) if isinstance(val, tuple) else [val] + # TODO: if indices is a tuple of any sequence, or if indices is a list, it's for advanced indexing + orig_slices = list(indices) if isinstance(indices, tuple) else [indices] count = defaultdict(list) for i,v in enumerate(orig_slices): count[type(v)].append(i) - if (num_slices := len(count[int]) + len(count[slice]) + len(count[Tensor])) > len(self.shape): raise IndexError(f"too many indices for tensor of dimension {len(self.shape)}") + # TODO: boolean indices + if (num_slices := len(count[int]) + len(count[slice]) + len(count[Tensor]) + len(count[list])) > len(self.shape): raise IndexError(f"too many indices for tensor of dimension {len(self.shape)}") if len(ellipsis_found := count[type(Ellipsis)]) > 1: raise IndexError("an index can only have a single ellipsis ('...')") + # replace ellipsis with equivalent number of slice(None) + # TODO: move all slice(None) to the end and transpose non-None to the front ellipsis_idx = ellipsis_found[0] if ellipsis_found else len(orig_slices) orig_slices[ellipsis_idx:ellipsis_idx+1] = [slice(None)] * (len(self.shape) - num_slices) @@ -337,8 +341,8 @@ class Tensor: if s is None: final_shape.append(1) else: # s is int or slice or Tensor dim_shape = next(it_shape) - if isinstance(s, int): - dim_collapsed += 1 + if isinstance(s, list): s = Tensor(s) + if isinstance(s, int): dim_collapsed += 1 else: assert isinstance(dim_shape, int), f"does not support symbolic shape {dim_shape}" final_shape.append(dim_shape)