Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 33 additions & 0 deletions helion/_compiler/type_propagation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
3 changes: 0 additions & 3 deletions test/test_indexing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]"""

Expand Down
Loading