diff --git a/helion/_compiler/type_propagation.py b/helion/_compiler/type_propagation.py index c6dd1cd55..28560f9ca 100644 --- a/helion/_compiler/type_propagation.py +++ b/helion/_compiler/type_propagation.py @@ -758,9 +758,42 @@ def visit_NamedExpr(self, node: ast.NamedExpr) -> TypeInfo: def visit_Subscript(self, node: ast.Subscript) -> TypeInfo: value_type = self.visit(node.value) + if isinstance(value_type, TensorType): + self._expand_ellipsis_in_subscript(node, value_type.fake_value.ndim) slice_type = self.visit(node.slice) return value_type.propagate_getitem(slice_type, self.origin()) + def _expand_ellipsis_in_subscript(self, node: ast.Subscript, ndim: int) -> None: + sl = node.slice + if isinstance(sl, ast.Constant) and sl.value is ...: + slices = [ + create(ast.Slice, lower=None, upper=None, step=None) + for _ in range(ndim) + ] + node.slice = create(ast.Tuple, elts=slices, ctx=ast.Load()) + return + if not isinstance(sl, ast.Tuple): + return + ellipsis_indices = [ + i + for i, elt in enumerate(sl.elts) + if isinstance(elt, ast.Constant) and elt.value is ... + ] + if not ellipsis_indices: + return + idx = ellipsis_indices[0] + dims_consumed = sum( + 1 + for elt in sl.elts + if not (isinstance(elt, ast.Constant) and elt.value in (..., None)) + ) + n_expand = ndim - dims_consumed + slices = [ + create(ast.Slice, lower=None, upper=None, step=None) + for _ in range(n_expand) + ] + sl.elts[idx : idx + 1] = slices + def visit_Slice(self, node: ast.Slice) -> TypeInfo: lower = ( self.visit(node.lower) diff --git a/test/test_indexing.py b/test/test_indexing.py index ab90c54ea..334197e17 100644 --- a/test/test_indexing.py +++ b/test/test_indexing.py @@ -1840,9 +1840,6 @@ def kernel(x: torch.Tensor) -> torch.Tensor: expected[:, -1] = 1.0 torch.testing.assert_close(result, expected) - @skipIfNormalMode( - "RankMismatch: Cannot assign a tensor of rank 2 to a buffer of rank 3" - ) def test_ellipsis_indexing(self): """Test both setter from scalar and getter for [..., i]"""