Skip to content
Closed
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
7 changes: 7 additions & 0 deletions tools/Polygraphy/CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
5 changes: 4 additions & 1 deletion tools/Polygraphy/polygraphy/backend/trt/loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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.
Expand Down
88 changes: 88 additions & 0 deletions tools/Polygraphy/polygraphy/backend/trt/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = {}

Expand Down
125 changes: 125 additions & 0 deletions tools/Polygraphy/tests/backend/trt/test_onnx_slice_check.py
Original file line number Diff line number Diff line change
@@ -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