diff --git a/tools/Polygraphy/CHANGELOG.md b/tools/Polygraphy/CHANGELOG.md index cea3567e..4562906b 100644 --- a/tools/Polygraphy/CHANGELOG.md +++ b/tools/Polygraphy/CHANGELOG.md @@ -2,6 +2,13 @@ Dates are in YYYY-MM-DD format. +## Unreleased +### Added +- `NetworkFromOnnx` / `NetworkFromOnnxPath` now validate `Slice` input lengths before + invoking the TensorRT ONNX parser and raise a clear, actionable error for mismatched + `starts`/`ends`/`axes`/`steps` lengths instead of TensorRT's cryptic + `Assertion failed: (starts.size() == axes.size())`. + ## v0.53.6 (2026-09-22) ### Added diff --git a/tools/Polygraphy/polygraphy/backend/trt/loader.py b/tools/Polygraphy/polygraphy/backend/trt/loader.py index 26177bdc..233a697c 100644 --- a/tools/Polygraphy/polygraphy/backend/trt/loader.py +++ b/tools/Polygraphy/polygraphy/backend/trt/loader.py @@ -335,7 +335,9 @@ def call_impl(self): used to populate it. """ builder, network, parser = super().call_impl() - success = parser.parse(util.invoke_if_callable(self._model_bytes)[0]) + model = util.invoke_if_callable(self._model_bytes)[0] + trt_util.check_onnx_slice_input_lengths(model) + success = parser.parse(model) trt_util.check_onnx_parser_errors(parser, success) return builder, network, parser @@ -407,6 +409,7 @@ def call_impl(self): used to populate it. """ path = util.invoke_if_callable(self.path)[0] + trt_util.check_onnx_slice_input_lengths(path) builder, network, parser = super().call_impl() # We need to use parse_from_file for the ONNX parser to keep track of the location of the ONNX file for # potentially parsing any external weights. diff --git a/tools/Polygraphy/polygraphy/backend/trt/util.py b/tools/Polygraphy/polygraphy/backend/trt/util.py index f09c6d80..49b0cf3d 100644 --- a/tools/Polygraphy/polygraphy/backend/trt/util.py +++ b/tools/Polygraphy/polygraphy/backend/trt/util.py @@ -96,6 +96,94 @@ def check_onnx_parser_errors(parser, success): ) +def check_onnx_slice_input_lengths(model): + """Validates Slice input lengths before invoking the TensorRT ONNX parser. + + TensorRT's ONNX parser rejects models whose ``Slice`` nodes have mismatched + ``starts``/``ends``/``axes``/``steps`` input lengths with a low-level assertion + (e.g. ``Assertion failed: (starts.size() == axes.size())``), even when the + offending node is easy to identify from the model. This check inspects the ONNX + graph up front and raises a clear, actionable error that names the offending node. + + Args: + model (Union[str, bytes]): Path to an ONNX model, or serialized ONNX model bytes. + + Raises: + PolygraphyException: If any ``Slice`` node has inconsistent input lengths. + """ + onnx = mod.lazy_import("onnx") + + if isinstance(model, str): + model_proto = onnx.load(model) + else: + model_proto = onnx.ModelProto() + model_proto.ParseFromString(model) + + graph = model_proto.graph + + def shape_of(name): + # Constant weights (initializers) + for tensor in graph.initializer: + if tensor.name == name: + return list(tensor.dims) + # Constant nodes + for node in graph.node: + if node.op_type == "Constant" and name in node.output: + for attr in node.attribute: + if attr.name == "value": + return list(attr.t.dims) + # Graph inputs / value_info (dynamic tensors with static shapes) + for value_info in list(graph.input) + list(graph.value_info): + if value_info.name == name: + tensor_type = value_info.type.tensor_type + if tensor_type.HasField("shape"): + return [dim.dim_value for dim in tensor_type.shape.dim] + return None + + def num_elements(shape): + if shape is None: + return None + n = 1 + for dim in shape: + if dim < 0: # symbolic dimension + return None + n *= dim + if dim == 0: + return 0 + return n + + slice_inputs = {"starts": 1, "ends": 2, "axes": 3, "steps": 4} + for node in graph.node: + if node.op_type != "Slice": + continue + + known = {} + for label, index in slice_inputs.items(): + if len(node.input) > index and node.input[index]: + n = num_elements(shape_of(node.input[index])) + if n is not None: + known[label] = n + + if "starts" not in known: + continue + + expected = known["starts"] + mismatched = {label: n for label, n in known.items() if n != expected} + if not mismatched: + continue + + node_name = node.name or f"(unnamed, output: {node.output[0]})" + lengths = ", ".join(f"{label}={n} element(s)" for label, n in known.items()) + raise PolygraphyException( + f"Could not import ONNX model into TensorRT: Slice node '{node_name}' has " + f"mismatched input lengths ({lengths}). TensorRT's ONNX parser requires " + f"starts, ends, axes, and steps to contain the same number of elements (see " + f"the ONNX Slice spec: https://onnx.ai/onnx/operators/onnx__Slice.html). " + f"This usually indicates the model was exported incorrectly; provide a " + f"matching start/end (and axis/step, if used) for every sliced axis." + ) + + def get_layer_class_mapping(): layer_class_mapping = {} diff --git a/tools/Polygraphy/tests/backend/trt/test_onnx_slice_check.py b/tools/Polygraphy/tests/backend/trt/test_onnx_slice_check.py new file mode 100644 index 00000000..7fd67f3b --- /dev/null +++ b/tools/Polygraphy/tests/backend/trt/test_onnx_slice_check.py @@ -0,0 +1,125 @@ +# +# SPDX-FileCopyrightText: Copyright (c) 1993-2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +""" +Tests for ``trt_util.check_onnx_slice_input_lengths``. + +These tests only require ``onnx`` and can run without a TensorRT installation. +""" + +from __future__ import annotations + +import pytest + +from polygraphy import mod +from polygraphy.backend.trt import util as trt_util +from polygraphy.exception import PolygraphyException + +onnx = mod.lazy_import("onnx") + + +def make_slice_model(start_len, end_len, axes_len=None, as_inputs=False): + """Build an opset-13 model with a single Slice node.""" + x = onnx.helper.make_tensor_value_info("x", onnx.TensorProto.FLOAT, [2, 3]) + y = onnx.helper.make_tensor_value_info("y", onnx.TensorProto.FLOAT, [None, None]) + + node_inputs = ["x", "starts", "ends"] + initializers = [ + onnx.helper.make_tensor( + "starts", onnx.TensorProto.INT64, [start_len], [0] * start_len + ), + onnx.helper.make_tensor( + "ends", onnx.TensorProto.INT64, [end_len], [2] * end_len + ), + ] + graph_inputs = [x] + if axes_len is not None: + node_inputs.append("axes") + initializers.append( + onnx.helper.make_tensor( + "axes", onnx.TensorProto.INT64, [axes_len], list(range(axes_len)) + ) + ) + + if as_inputs: + # Keep the Slice parameters as graph inputs with static shapes instead of initializers. + graph_inputs += [ + onnx.helper.make_tensor_value_info(name, onnx.TensorProto.INT64, [n]) + for name, n in (("starts", start_len), ("ends", end_len)) + ] + if axes_len is not None: + graph_inputs.append( + onnx.helper.make_tensor_value_info( + "axes", onnx.TensorProto.INT64, [axes_len] + ) + ) + initializers = [] + + node = onnx.helper.make_node("Slice", node_inputs, ["y"], name="slicer") + graph = onnx.helper.make_graph( + [node], "g", graph_inputs, [y], initializer=initializers + ) + model = onnx.helper.make_model( + graph, opset_imports=[onnx.helper.make_opsetid("", 13)] + ) + return model + + +def test_slice_mismatched_lengths_raises(): + model = make_slice_model(start_len=1, end_len=2, axes_len=2) + with pytest.raises(PolygraphyException, match="slicer"): + trt_util.check_onnx_slice_input_lengths(model.SerializeToString()) + + +def test_slice_mismatched_lengths_from_path(tmp_path): + model = make_slice_model(start_len=1, end_len=2, axes_len=2) + path = tmp_path / "slice_bad.onnx" + onnx.save(model, path) + with pytest.raises(PolygraphyException, match="starts=1"): + trt_util.check_onnx_slice_input_lengths(str(path)) + + +def test_slice_dynamic_inputs_with_static_shapes_raises(): + model = make_slice_model(start_len=1, end_len=2, axes_len=2, as_inputs=True) + with pytest.raises(PolygraphyException, match="slicer"): + trt_util.check_onnx_slice_input_lengths(model.SerializeToString()) + + +def test_slice_matching_lengths_passes(): + model = make_slice_model(start_len=2, end_len=2, axes_len=2) + assert trt_util.check_onnx_slice_input_lengths(model.SerializeToString()) is None + + +def test_slice_no_axes_passes(): + model = make_slice_model(start_len=2, end_len=2) + assert trt_util.check_onnx_slice_input_lengths(model.SerializeToString()) is None + + +def test_slice_steps_omitted_passes(): + model = make_slice_model(start_len=1, end_len=1) + assert trt_util.check_onnx_slice_input_lengths(model.SerializeToString()) is None + + +def test_model_without_slice_passes(): + x = onnx.helper.make_tensor_value_info("x", onnx.TensorProto.FLOAT, [2, 2]) + y = onnx.helper.make_tensor_value_info("y", onnx.TensorProto.FLOAT, [2, 2]) + node = onnx.helper.make_node("Relu", ["x"], ["y"]) + graph = onnx.helper.make_graph([node], "g", [x], [y]) + model = onnx.helper.make_model( + graph, opset_imports=[onnx.helper.make_opsetid("", 13)] + ) + assert trt_util.check_onnx_slice_input_lengths(model.SerializeToString()) is None