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
4 changes: 2 additions & 2 deletions onnxscript/_internal/converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -775,7 +775,6 @@ def translate_slice(slice_expr: ast.Slice) -> tuple[ir.Value, ir.Value, ir.Value
squeezed_axes = []
for axis, expr in scalar_indices:
# Treat a scalar index i as slice "i:i+1:1", but squeeze the axis finally.
# TODO: handle negative i
index = self._eval_constant_expr(expr)
squeezed_axes.append(axis)
kwargs = dict(
Expand All @@ -784,7 +783,8 @@ def translate_slice(slice_expr: ast.Slice) -> tuple[ir.Value, ir.Value, ir.Value
)
element = ast.Slice(
ast.Constant(index, **kwargs),
ast.Constant(index + 1, **kwargs),
# -1 selects the last element, so its stop must reach the axis end.
ast.Constant(maxint if index == -1 else index + 1, **kwargs),
ast.Constant(1, **kwargs),
)
sliced_indices.append((axis, element))
Expand Down
82 changes: 82 additions & 0 deletions onnxscript/_internal/slicing_test.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,82 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT License.
"""Parity tests for scalar indexing combined with slices."""

import itertools
import unittest

import numpy as np
import onnx
import onnxruntime as ort
import parameterized

from onnxscript import INT64, evaluator, opset15, script


@script(default_opset=opset15)
def last_row_slice(x: INT64["N", "M"]) -> INT64["K"]: # noqa: F821
return x[-1, :2]


@script(default_opset=opset15)
def last_column_slice(x: INT64["N", "M"]) -> INT64["K"]: # noqa: F821
return x[:2, -1]


@script(default_opset=opset15)
def last_element(x: INT64["N", "M"]) -> INT64: # noqa: F821
return x[-1, -1]


@script(default_opset=opset15)
def penultimate_row_last_column(x: INT64["N", "M"]) -> INT64: # noqa: F821
return x[-2, -1]


@script(default_opset=opset15)
def reversed_last_row(x: INT64["N", "M"]) -> INT64["M"]: # noqa: F821
return x[-1, ::-1]


@script(default_opset=opset15)
def first_row_slice(x: INT64["N", "M"]) -> INT64["K"]: # noqa: F821
return x[0, :2]


class TestScalarSlicing(unittest.TestCase):
@parameterized.parameterized.expand(
itertools.product(
[
(last_row_slice, (-1, slice(None, 2))),
(last_column_slice, (slice(None, 2), -1)),
(last_element, (-1, -1)),
(penultimate_row_last_column, (-2, -1)),
(reversed_last_row, (-1, slice(None, None, -1))),
(first_row_slice, (0, slice(None, 2))),
],
[(4, 3), (2, 5)],
["eager", "onnxruntime"],
)
)
def test_matches_numpy(self, function_and_index, shape, mode):
function, index = function_and_index
data = np.arange(np.prod(shape), dtype=np.int64).reshape(shape)
if mode == "eager":
with evaluator.default_as(evaluator.OnnxReferenceRuntimeEvaluator()):
actual = function(data)
else:
model = function.to_model_proto(ir_version=9)
onnx.checker.check_model(model)
options = ort.SessionOptions()
options.intra_op_num_threads = 1
options.inter_op_num_threads = 1
session = ort.InferenceSession(
model.SerializeToString(), options, providers=["CPUExecutionProvider"]
)
actual = session.run(None, {"x": data})[0]
np.testing.assert_array_equal(actual, data[index])
self.assertEqual(actual.shape, data[index].shape)


if __name__ == "__main__":
unittest.main()
4 changes: 3 additions & 1 deletion onnxscript/tensor.py
Original file line number Diff line number Diff line change
Expand Up @@ -118,7 +118,9 @@ def __getitem__(self, index):
)
elif isinstance(s, Tensor):
if s.is_scalar:
scalar_indices.append([s, s + 1, axis_, 1])
# A stop of zero would make the last-element slice empty.
stop = shape[axis_] if int(s) == -1 else s + 1
scalar_indices.append([s, stop, axis_, 1])
to_squeeze.append(axis_)
else:
non_scalar_indices.append((axis_, s))
Expand Down