This is an automated email from the ASF dual-hosted git repository.
tlopex pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm.git
The following commit(s) were added to refs/heads/main by this push:
new 8526df4590 [Relax][Frontend][ONNX] Support dynamic Range bounds
(#20109)
8526df4590 is described below
commit 8526df4590149ff866b52dd3ad1f35d0e817ac84
Author: Ronald Nap <[email protected]>
AuthorDate: Thu Sep 3 12:58:32 2026 -0700
[Relax][Frontend][ONNX] Support dynamic Range bounds (#20109)
Fixes [#20064](https://github.com/apache/tvm/issues/20064).
This updates ONNX `Range` lowering to support runtime scalar bounds,
following the same strategy used by [Relax TFLite dynamic RANGE
lowering](https://github.com/apache/tvm/pull/19867). The strategy is to
compute the output length in graph, lift it to a symbolic dimension, and
rebuild the range as `arange(0, length) * delta + start`.
The existing constant path is preserved, and symbolic `PrimExpr` bounds
continue to use `relax.op.arange` directly. Tests are added for constant
`Range`, parametrized dynamic scalar `Range` inputs, and symbolic
`PrimExpr` limit.
---
python/tvm/relax/frontend/onnx/onnx_frontend.py | 88 +++++++++++--
tests/python/relax/test_frontend_onnx.py | 161 ++++++++++++++++++++++++
2 files changed, 237 insertions(+), 12 deletions(-)
diff --git a/python/tvm/relax/frontend/onnx/onnx_frontend.py
b/python/tvm/relax/frontend/onnx/onnx_frontend.py
index e8bd5eb0ae..725746c805 100644
--- a/python/tvm/relax/frontend/onnx/onnx_frontend.py
+++ b/python/tvm/relax/frontend/onnx/onnx_frontend.py
@@ -3991,24 +3991,88 @@ class Range(OnnxOpConverter):
start = get_constant(inputs[0], params)
limit = get_constant(inputs[1], params)
delta = get_constant(inputs[2], params)
- out_dtype = start.ty.dtype
- if isinstance(start, relax.Constant):
- start = start.data.numpy().tolist()
+ def get_scalar_dtype(x):
+ if tvm.ir.is_prim_expr(x):
+ return str(getattr(x, "dtype", None) or x.ty)
+ return str(x.ty.dtype)
- if isinstance(limit, relax.Constant):
- limit = limit.data.numpy().tolist()
+ out_dtype = get_scalar_dtype(start)
- assert isinstance(delta, relax.Constant), "Constant delta required for
Range."
- step = delta.data.numpy().tolist()
+ def get_scalar_value(x):
+ if isinstance(x, relax.Constant):
+ value = x.data.numpy()
+ if value.size != 1:
+ raise ValueError("Range scalar input must have exactly one
element.")
+ return value.item()
+ return x
+
+ start = get_scalar_value(start)
+ limit = get_scalar_value(limit)
+ delta = get_scalar_value(delta)
+
+ def is_dynamic_scalar(x):
+ return tvm.ir.is_prim_expr(x) or isinstance(x, relax.Expr)
- # If all inputs are constant, compute directly.
- if isinstance(start, int) and isinstance(limit, int):
- out_range = _np.arange(start=start, stop=limit, step=step)
+ out_dtype_is_float = _relax_dtype_is_floating_point(out_dtype)
+
+ if not any(is_dynamic_scalar(x) for x in [start, limit, delta]):
+ if out_dtype_is_float:
+ np_dtype = _np.dtype(out_dtype).type
+ start, limit, delta = map(np_dtype, (start, limit, delta))
+ difference = limit - start
+ count = max(int(_np.ceil(_np.float64(difference) /
_np.float64(delta))), 0)
+ out_range = _np.arange(count, dtype=out_dtype) * delta + start
+ else:
+ out_range = _np.arange(start=start, stop=limit, step=delta)
return relax.const(out_range, out_dtype)
- # Otherwise compute in graph.
- return relax.op.arange(start, limit, step, out_dtype)
+ count_dtype = "float64" if out_dtype_is_float else "int64"
+
+ def scalar_expr(x, dtype):
+ if tvm.ir.is_prim_expr(x):
+ expr_dtype = str(getattr(x, "dtype", None) or x.ty)
+ if expr_dtype != "int64":
+ x = tirx.Cast("int64", x)
+ x =
bb.normalize(relax.op.shape_to_tensor(relax.ShapeExpr([x])))
+ x = bb.normalize(relax.op.reshape(x, ()))
+ if dtype != "int64":
+ x = bb.normalize(relax.op.astype(x, dtype))
+ return x
+ if isinstance(x, relax.Expr):
+ if str(x.ty.dtype) == dtype:
+ return x
+ return bb.normalize(relax.op.astype(x, dtype))
+ return relax.const(x, dtype)
+
+ if out_dtype_is_float:
+ start_value = scalar_expr(start, out_dtype)
+ limit_value = scalar_expr(limit, out_dtype)
+ delta_value = scalar_expr(delta, out_dtype)
+
+ difference = bb.normalize(relax.op.subtract(limit_value,
start_value))
+ difference = bb.normalize(relax.op.astype(difference, count_dtype))
+ delta_count = bb.normalize(relax.op.astype(delta_value,
count_dtype))
+ count = relax.op.ceil(relax.op.divide(difference, delta_count))
+ else:
+ start_count = scalar_expr(start, count_dtype)
+ limit_count = scalar_expr(limit, count_dtype)
+ delta_count = scalar_expr(delta, count_dtype)
+ count = relax.op.negative(
+ relax.op.floor_divide(relax.op.subtract(start_count,
limit_count), delta_count)
+ )
+ start_value = scalar_expr(start, out_dtype)
+ delta_value = scalar_expr(delta, out_dtype)
+
+ count = bb.normalize(relax.op.maximum(count, relax.const(0,
count_dtype)))
+ count = bb.normalize(relax.op.astype(count, "int64"))
+ count = bb.normalize(relax.op.reshape(count, (1,)))
+ range_len = _tensor_to_shape_expr(bb, count, 1, "range_len").values[0]
+
+ positions = bb.normalize(
+ relax.op.astype(relax.op.arange(0, range_len, 1, "int64"),
out_dtype)
+ )
+ return relax.op.add(relax.op.multiply(positions, delta_value),
start_value)
class InstanceNormalization(OnnxOpConverter):
diff --git a/tests/python/relax/test_frontend_onnx.py
b/tests/python/relax/test_frontend_onnx.py
index 4c8d570279..591a2a8ee3 100644
--- a/tests/python/relax/test_frontend_onnx.py
+++ b/tests/python/relax/test_frontend_onnx.py
@@ -9202,6 +9202,167 @@ def test_range():
tvm.ir.assert_structural_equal(tvm_model, Expected)
+def test_range_constant_float():
+ range_node = helper.make_node(
+ "Range",
+ ["start", "limit", "delta"],
+ ["output"],
+ )
+
+ graph = helper.make_graph(
+ [range_node],
+ "range_constant_float_test",
+ inputs=[],
+ initializer=[
+ helper.make_tensor("start", TensorProto.FLOAT, [], [-7.8]),
+ helper.make_tensor("limit", TensorProto.FLOAT, [], [-1.4]),
+ helper.make_tensor("delta", TensorProto.FLOAT, [], [0.2]),
+ ],
+ outputs=[
+ helper.make_tensor_value_info("output", TensorProto.FLOAT,
["range_len"]),
+ ],
+ )
+
+ model = helper.make_model(graph, producer_name="range_constant_float_test")
+ check_correctness(model, opset=12, check_dtypes=True)
+
+
[email protected](
+ "start, limit, delta, tensor_dtype, np_dtype",
+ [
+ (0, 6, 2, TensorProto.INT64, np.int64),
+ (8, 0, -2, TensorProto.INT64, np.int64),
+ (5, 1, 1, TensorProto.INT64, np.int64),
+ (0, 7, 2, TensorProto.INT32, np.int32),
+ (0.0, 1.0, 0.25, TensorProto.FLOAT, np.float32),
+ (1.0, -1.0, -0.5, TensorProto.FLOAT, np.float32),
+ (0.0, 0.3, 0.1, TensorProto.FLOAT, np.float32),
+ (-7.8, -1.4, 0.2, TensorProto.FLOAT, np.float32),
+ ],
+)
+def test_range_dynamic_scalar_inputs(start, limit, delta, tensor_dtype,
np_dtype):
+ range_node = helper.make_node(
+ "Range",
+ ["start", "limit", "delta"],
+ ["output"],
+ )
+
+ graph = helper.make_graph(
+ [range_node],
+ "range_dynamic_scalar_inputs_test",
+ inputs=[
+ helper.make_tensor_value_info("start", tensor_dtype, []),
+ helper.make_tensor_value_info("limit", tensor_dtype, []),
+ helper.make_tensor_value_info("delta", tensor_dtype, []),
+ ],
+ outputs=[
+ helper.make_tensor_value_info("output", tensor_dtype,
["range_len"]),
+ ],
+ )
+
+ model = helper.make_model(graph,
producer_name="range_dynamic_scalar_inputs_test")
+ check_correctness(
+ model,
+ inputs={
+ "start": np.array(start, dtype=np_dtype),
+ "limit": np.array(limit, dtype=np_dtype),
+ "delta": np.array(delta, dtype=np_dtype),
+ },
+ opset=12,
+ check_dtypes=True,
+ )
+
+
+def test_range_mixed_tensor_and_primexpr_limit():
+ shape = helper.make_node("Shape", ["x"], ["x_shape"])
+ axis = make_constant_node("axis", TensorProto.INT64, [], [1])
+ gather = helper.make_node("Gather", ["x_shape", "axis"], ["limit_int"])
+ cast = helper.make_node("Cast", ["limit_int"], ["limit"],
to=TensorProto.FLOAT)
+ delta = make_constant_node("delta", TensorProto.FLOAT, [], [1.0])
+ range_node = helper.make_node(
+ "Range",
+ ["start", "limit", "delta"],
+ ["output"],
+ )
+
+ graph = helper.make_graph(
+ [shape, axis, gather, cast, delta, range_node],
+ "range_mixed_tensor_and_primexpr_limit_test",
+ inputs=[
+ helper.make_tensor_value_info("x", TensorProto.FLOAT, [1,
"range_len"]),
+ helper.make_tensor_value_info("start", TensorProto.FLOAT, []),
+ ],
+ outputs=[
+ helper.make_tensor_value_info("output", TensorProto.FLOAT,
["range_len"]),
+ ],
+ )
+
+ model = helper.make_model(
+ graph,
+ producer_name="range_mixed_tensor_and_primexpr_limit_test",
+ opset_imports=[helper.make_opsetid("", 17)],
+ )
+ model.ir_version = 8
+ check_correctness(
+ model,
+ inputs={
+ "x": np.ones((1, 4), dtype=np.float32),
+ "start": np.array(0.0, dtype=np.float32),
+ },
+ opset=17,
+ check_dtypes=True,
+ )
+
+
[email protected](
+ "start_from_dim, limit_from_dim, delta",
+ [
+ (True, False, -1),
+ (False, True, -1),
+ ],
+)
+def test_range_primexpr_negative_and_empty(start_from_dim, limit_from_dim,
delta):
+ shape = helper.make_node("Shape", ["x"], ["x_shape"])
+ axis = make_constant_node("axis", TensorProto.INT64, [], [1])
+ gather = helper.make_node("Gather", ["x_shape", "axis"], ["dim"])
+ start = make_constant_node("start", TensorProto.INT64, [], [0])
+ limit = make_constant_node("limit", TensorProto.INT64, [], [0])
+ delta_node = make_constant_node("delta", TensorProto.INT64, [], [delta])
+
+ range_inputs = [
+ "dim" if start_from_dim else "start",
+ "dim" if limit_from_dim else "limit",
+ "delta",
+ ]
+ range_node = helper.make_node("Range", range_inputs, ["output"])
+
+ graph = helper.make_graph(
+ [shape, axis, gather, start, limit, delta_node, range_node],
+ "range_primexpr_negative_and_empty_test",
+ inputs=[
+ helper.make_tensor_value_info("x", TensorProto.FLOAT, [1,
"range_len"]),
+ ],
+ outputs=[
+ helper.make_tensor_value_info("output", TensorProto.INT64,
["output_len"]),
+ ],
+ )
+
+ model = helper.make_model(
+ graph,
+ producer_name="range_primexpr_negative_and_empty_test",
+ opset_imports=[helper.make_opsetid("", 17)],
+ )
+ model.ir_version = 8
+ check_correctness(
+ model,
+ inputs={
+ "x": np.ones((1, 4), dtype=np.float32),
+ },
+ opset=17,
+ check_dtypes=True,
+ )
+
+
def test_batch_norm():
batch_norm_node = helper.make_node(
"BatchNormalization", ["x", "s", "bias", "mean", "var"], ["y"],
epsilon=1e-2