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 fa8789b388 [Fix][Relax] Canonicalize strided slice begin indices
(#20284)
fa8789b388 is described below
commit fa8789b388da0f642b0a210e17ad5ddcd11eb309
Author: Ronald Nap <[email protected]>
AuthorDate: Thu Sep 10 10:55:04 2026 -0700
[Fix][Relax] Canonicalize strided slice begin indices (#20284)
## Summary
This fixes #20262 and #20263.
Currently, `strided_slice` lowering can fail when `begin` falls outside
the input bounds. The output shape already canonicalizes `begin` and
`end` when `assume_inbound=False`, but the generated source index could
still use the original `begin` value. This can produces the correct
output shape while reading from the wrong location or outside the input
tensor.
This PR updates the affected lowering paths to use the existing
`CanonicalizeIndex` logic when constructing source indices and handles
`begin == extent` with a negative stride, where the canonical start is
`extent - 1`. Existing `assume_inbound=True` logic is unchanged.
---
include/tvm/topi/transform.h | 22 ++--
tests/python/relax/test_e2e_op_dynamic.py | 37 ++++++-
tests/python/relax/test_op_index.py | 4 +-
..._transform_legalize_ops_index_linear_algebra.py | 114 +++++++++++++++++----
4 files changed, 148 insertions(+), 29 deletions(-)
diff --git a/include/tvm/topi/transform.h b/include/tvm/topi/transform.h
index 8ffa28c43a..7092eaa42c 100644
--- a/include/tvm/topi/transform.h
+++ b/include/tvm/topi/transform.h
@@ -650,11 +650,8 @@ inline ffi::Array<Tensor> split_indices_array(const
Tensor& x, ffi::Array<PrimEx
}
inline PrimExpr DynamicCanonicalizeIndex(PrimExpr index, PrimExpr extent,
PrimExpr stride) {
- auto idx_var = index.as<tvm::tirx::PrimVar>();
- auto extent_var = extent.as<tvm::tirx::PrimVar>();
-
- if (idx_var && extent_var && (*idx_var)->name == (*extent_var)->name) {
- return index;
+ if (index.same_as(extent)) {
+ return tvm::if_then_else(stride < 0, extent - 1, extent);
}
PrimExpr begin_range = tvm::if_then_else(stride < 0, -1, 0);
@@ -747,7 +744,11 @@ inline te::Tensor dynamic_strided_slice_with_axes(
for (size_t i = 0; i < begin.size(); i++) {
int axis = static_cast<int>(axes[i]);
- PrimExpr new_index = indices[axis] * strides[i] + begin[i];
+ PrimExpr begin_index = begin[i];
+ if (!assume_inbound) {
+ begin_index = CanonicalizeIndex(begin_index, x->shape[axis],
strides[i]);
+ }
+ PrimExpr new_index = indices[axis] * strides[i] + begin_index;
real_indices.Set(axis, new_index);
}
@@ -805,7 +806,11 @@ inline Tensor dynamic_strided_slice(const Tensor& x, const
ffi::Array<PrimExpr>&
[&](const ffi::Array<tvm::tirx::PrimVar>& indices) {
ffi::Array<PrimExpr> real_indices;
for (size_t i = 0; i < num_slice_axes; ++i) {
- real_indices.push_back(indices[i] * strides[i] + tvm::min(begin[i],
x->shape[i] - 1));
+ PrimExpr begin_index = tvm::min(begin[i], x->shape[i] - 1);
+ if (!assume_inbound) {
+ begin_index = CanonicalizeIndex(begin[i], x->shape[i], strides[i]);
+ }
+ real_indices.push_back(indices[i] * strides[i] + begin_index);
}
// keep input dim
for (size_t i = num_slice_axes; i < src_tensor_dim; ++i) {
@@ -2319,7 +2324,8 @@ inline te::Tensor dynamic_strided_slice(const te::Tensor&
x, const te::Tensor& b
ffi::Array<PrimExpr> real_indices;
for (size_t i = 0; i < num_dynamic_axes; ++i) {
auto ind = IntImm::Int64(i);
- real_indices.push_back(indices[i] * strides(ind) +
tvm::min(begin(ind), x->shape[i] - 1));
+ PrimExpr begin_index = CanonicalizeIndex(begin(ind), x->shape[i],
strides(ind));
+ real_indices.push_back(indices[i] * strides(ind) + begin_index);
}
return x(real_indices);
},
diff --git a/tests/python/relax/test_e2e_op_dynamic.py
b/tests/python/relax/test_e2e_op_dynamic.py
index 25ebbea2c2..56628c7270 100644
--- a/tests/python/relax/test_e2e_op_dynamic.py
+++ b/tests/python/relax/test_e2e_op_dynamic.py
@@ -24,7 +24,7 @@ import tvm.testing
pytest.importorskip("scipy") # tvm.topi.testing imports scipy
import tvm.topi.testing
-from tvm import relax
+from tvm import relax, tirx
from tvm.relax.transform import LegalizeOps
from tvm.script import relax as R
from tvm.script import tirx as T
@@ -44,6 +44,9 @@ def build(mod):
([0, 2, 4, 4], [5, 5, 7, 8], [1, 1, 2, 3]),
([0, 2, 4, 4], [5, 5, 11, 10], [1, 1, 1, 1]),
([0, 2, 10, 14], [0, 5, 1, 1], [1, 1, -1, -2]),
+ ([-2, 0, 0, 0], [8, 9, 10, 10], [1, 1, 1, 1]),
+ ([-58, 0, 0, 0], [8, 9, 10, 10], [1, 1, 1, 1]),
+ ([-1, 0, 0, 0], [-9, 9, 10, 10], [-1, 1, 1, 1]),
],
)
def test_dynamic_strided_slice(begin, end, strides):
@@ -75,6 +78,9 @@ def test_dynamic_strided_slice(begin, end, strides):
([0, 2, 4, 4], [5, 5, 7, 8], [1, 1, 2, 3]),
([0, 2, 4, 4], [5, 5, 11, 10], [1, 1, 1, 1]),
([0, 2, 10, 14], [0, 5, 1, 1], [1, 1, -1, -2]),
+ ([-2, 0, 0, 0], [8, 9, 10, 10], [1, 1, 1, 1]),
+ ([-58, 0, 0, 0], [8, 9, 10, 10], [1, 1, 1, 1]),
+ ([-1, 0, 0, 0], [-9, 9, 10, 10], [-1, 1, 1, 1]),
],
)
def test_dynamic_strided_slice_symbolic(begin, end, strides):
@@ -102,5 +108,34 @@ def test_dynamic_strided_slice_symbolic(begin, end,
strides):
tvm.testing.assert_allclose(out_nd.numpy(), out_npy)
[email protected]("begin_offset", [0, 1])
+def test_strided_slice_symbolic_out_of_bounds(begin_offset):
+ dim = tirx.Var("m", "int64")
+
+ bb = relax.BlockBuilder()
+ x = relax.Var("x", relax.TensorType([dim], "float32"))
+
+ with bb.function("main", params=[x]):
+ with bb.dataflow():
+ y = bb.emit(
+ relax.op.strided_slice(
+ x,
+ axes=[0],
+ begin=[dim + begin_offset],
+ end=[-dim - 1],
+ strides=[-1],
+ )
+ )
+ gv = bb.emit_output(y)
+ bb.emit_func_output(gv)
+
+ vm = build(bb.get())
+
+ x_np = np.arange(8, dtype=np.float32)
+ out = vm["main"](tvm.runtime.tensor(x_np, dev)).numpy()
+
+ tvm.testing.assert_allclose(out, x_np[::-1])
+
+
if __name__ == "__main__":
tvm.testing.main()
diff --git a/tests/python/relax/test_op_index.py
b/tests/python/relax/test_op_index.py
index f70e542b95..2894eb98e7 100644
--- a/tests/python/relax/test_op_index.py
+++ b/tests/python/relax/test_op_index.py
@@ -954,9 +954,9 @@ def test_legalize_dynamic_begin_inf_end():
for ax0, ax1 in T.grid(T.max(T.int64(16) -
T.max(T.if_then_else(index < T.int64(0), index + T.int64(16), index),
T.int64(0)), T.int64(0)), T.int64(16)):
with T.sblock("T_dynamic_strided_slice_with_axes"):
v_ax0, v_ax1 = T.axis.remap("SS", [ax0, ax1])
- T.reads(A[v_ax0 + index, v_ax1])
+ T.reads(A[v_ax0 : v_ax0 + T.int64(17), v_ax1])
T.writes(T_dynamic_strided_slice_with_axes[v_ax0, v_ax1])
- T_dynamic_strided_slice_with_axes[v_ax0, v_ax1] = A[v_ax0
+ index, v_ax1]
+ T_dynamic_strided_slice_with_axes[v_ax0, v_ax1] =
A[T.min(T.max(T.if_then_else(index < T.int64(0), index + T.int64(16), index),
T.int64(0)), T.int64(16)) + v_ax0, v_ax1]
@R.function
def main(A: R.Tensor((16, 16), dtype="float32"), B:
R.Shape(["index"])) -> R.Tensor(("T.max(16 - T.max(T.if_then_else(index < 0,
index + 16, index), 0), 0)", 16), dtype="float32"):
diff --git
a/tests/python/relax/test_transform_legalize_ops_index_linear_algebra.py
b/tests/python/relax/test_transform_legalize_ops_index_linear_algebra.py
index a026507ae9..55b52393d5 100644
--- a/tests/python/relax/test_transform_legalize_ops_index_linear_algebra.py
+++ b/tests/python/relax/test_transform_legalize_ops_index_linear_algebra.py
@@ -460,27 +460,79 @@ def test_dynamic_strided_slice():
v_ax0, v_ax1, v_ax2, v_ax3 = T.axis.remap("SSSS", [ax0,
ax1, ax2, ax3])
T.reads(
rxplaceholder[
- T.min(rxplaceholder_1[T.int64(0)], T.int64(7))
- + v_ax0 * rxplaceholder_3[T.int64(0)],
- T.min(rxplaceholder_1[T.int64(1)], T.int64(8))
- + v_ax1 * rxplaceholder_3[T.int64(1)],
- T.min(rxplaceholder_1[T.int64(2)], T.int64(9))
- + v_ax2 * rxplaceholder_3[T.int64(2)],
- T.min(rxplaceholder_1[T.int64(3)], T.int64(9))
- + v_ax3 * rxplaceholder_3[T.int64(3)],
+ T.int64(0) : T.int64(8),
+ T.int64(0) : T.int64(9),
+ T.int64(0) : T.int64(10),
+ T.int64(0) : T.int64(10),
],
rxplaceholder_1[T.int64(0) : T.int64(4)],
rxplaceholder_3[T.int64(0) : T.int64(4)],
)
T.writes(T_strided_slice_dynamic[v_ax0, v_ax1, v_ax2,
v_ax3])
T_strided_slice_dynamic[v_ax0, v_ax1, v_ax2, v_ax3] =
rxplaceholder[
- T.min(rxplaceholder_1[T.int64(0)], T.int64(7))
+ T.min(
+ T.max(
+ T.if_then_else(
+ rxplaceholder_1[T.int64(0)] < T.int64(0),
+ rxplaceholder_1[T.int64(0)] + T.int64(8),
+ rxplaceholder_1[T.int64(0)],
+ ),
+ T.if_then_else(
+ rxplaceholder_3[T.int64(0)] < T.int64(0),
T.int64(-1), T.int64(0)
+ ),
+ ),
+ T.if_then_else(
+ rxplaceholder_3[T.int64(0)] < T.int64(0),
T.int64(7), T.int64(8)
+ ),
+ )
+ v_ax0 * rxplaceholder_3[T.int64(0)],
- T.min(rxplaceholder_1[T.int64(1)], T.int64(8))
+ T.min(
+ T.max(
+ T.if_then_else(
+ rxplaceholder_1[T.int64(1)] < T.int64(0),
+ rxplaceholder_1[T.int64(1)] + T.int64(9),
+ rxplaceholder_1[T.int64(1)],
+ ),
+ T.if_then_else(
+ rxplaceholder_3[T.int64(1)] < T.int64(0),
T.int64(-1), T.int64(0)
+ ),
+ ),
+ T.if_then_else(
+ rxplaceholder_3[T.int64(1)] < T.int64(0),
T.int64(8), T.int64(9)
+ ),
+ )
+ v_ax1 * rxplaceholder_3[T.int64(1)],
- T.min(rxplaceholder_1[T.int64(2)], T.int64(9))
+ T.min(
+ T.max(
+ T.if_then_else(
+ rxplaceholder_1[T.int64(2)] < T.int64(0),
+ rxplaceholder_1[T.int64(2)] + T.int64(10),
+ rxplaceholder_1[T.int64(2)],
+ ),
+ T.if_then_else(
+ rxplaceholder_3[T.int64(2)] < T.int64(0),
T.int64(-1), T.int64(0)
+ ),
+ ),
+ T.if_then_else(
+ rxplaceholder_3[T.int64(2)] < T.int64(0),
T.int64(9), T.int64(10)
+ ),
+ )
+ v_ax2 * rxplaceholder_3[T.int64(2)],
- T.min(rxplaceholder_1[T.int64(3)], T.int64(9))
+ T.min(
+ T.max(
+ T.if_then_else(
+ rxplaceholder_1[T.int64(3)] < T.int64(0),
+ rxplaceholder_1[T.int64(3)] + T.int64(10),
+ rxplaceholder_1[T.int64(3)],
+ ),
+ T.if_then_else(
+ rxplaceholder_3[T.int64(3)] < T.int64(0),
T.int64(-1), T.int64(0)
+ ),
+ ),
+ T.if_then_else(
+ rxplaceholder_3[T.int64(3)] < T.int64(0),
T.int64(9), T.int64(10)
+ ),
+ )
+ v_ax3 * rxplaceholder_3[T.int64(3)],
]
@@ -748,19 +800,45 @@ def test_dynamic_strided_slice_symbolic():
v_ax0, v_ax1 = T.axis.remap("SS", [ax0, ax1])
T.reads(
rxplaceholder_3[
- T.min(rxplaceholder[T.int64(0)], T.int64(9))
- + v_ax0 * rxplaceholder_2[T.int64(0)],
- T.min(rxplaceholder[T.int64(1)], n - T.int64(1))
- + v_ax1 * rxplaceholder_2[T.int64(1)],
+ T.int64(0) : T.int64(10),
+ T.int64(0) : n,
],
rxplaceholder[T.int64(0) : T.int64(2)],
rxplaceholder_2[T.int64(0) : T.int64(2)],
)
T.writes(T_strided_slice_dynamic[v_ax0, v_ax1])
T_strided_slice_dynamic[v_ax0, v_ax1] = rxplaceholder_3[
- T.min(rxplaceholder[T.int64(0)], T.int64(9))
+ T.min(
+ T.max(
+ T.if_then_else(
+ rxplaceholder[T.int64(0)] < T.int64(0),
+ rxplaceholder[T.int64(0)] + T.int64(10),
+ rxplaceholder[T.int64(0)],
+ ),
+ T.if_then_else(
+ rxplaceholder_2[T.int64(0)] < T.int64(0),
T.int64(-1), T.int64(0)
+ ),
+ ),
+ T.if_then_else(
+ rxplaceholder_2[T.int64(0)] < T.int64(0),
T.int64(9), T.int64(10)
+ ),
+ )
+ v_ax0 * rxplaceholder_2[T.int64(0)],
- T.min(rxplaceholder[T.int64(1)], n - T.int64(1))
+ T.min(
+ T.max(
+ T.if_then_else(
+ rxplaceholder[T.int64(1)] < T.int64(0),
+ rxplaceholder[T.int64(1)] + n,
+ rxplaceholder[T.int64(1)],
+ ),
+ T.if_then_else(
+ rxplaceholder_2[T.int64(1)] < T.int64(0),
T.int64(-1), T.int64(0)
+ ),
+ ),
+ T.if_then_else(
+ rxplaceholder_2[T.int64(1)] < T.int64(0), n -
T.int64(1), n
+ ),
+ )
+ v_ax1 * rxplaceholder_2[T.int64(1)],
]