From 4e6e96985dbfaa7938ef0ea6fff8eb9bd08407c2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=9E=97SO?= <142557582+Linxiushen@users.noreply.github.com> Date: Mon, 28 Sep 2026 01:50:12 +0800 Subject: [PATCH] Fix last-element indexing mixed with slices or scalar indices --- onnxscript/_internal/converter.py | 4 +- onnxscript/_internal/slicing_test.py | 82 ++++++++++++++++++++++++++++ onnxscript/tensor.py | 4 +- 3 files changed, 87 insertions(+), 3 deletions(-) create mode 100644 onnxscript/_internal/slicing_test.py diff --git a/onnxscript/_internal/converter.py b/onnxscript/_internal/converter.py index c2f6b0fb63..b3cb8a7b8c 100644 --- a/onnxscript/_internal/converter.py +++ b/onnxscript/_internal/converter.py @@ -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( @@ -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)) diff --git a/onnxscript/_internal/slicing_test.py b/onnxscript/_internal/slicing_test.py new file mode 100644 index 0000000000..16ad8a63f9 --- /dev/null +++ b/onnxscript/_internal/slicing_test.py @@ -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() diff --git a/onnxscript/tensor.py b/onnxscript/tensor.py index 0c70cc57be..91090abaf6 100644 --- a/onnxscript/tensor.py +++ b/onnxscript/tensor.py @@ -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))