From e255506bcd495116b23fd3f8f74ca7499d48be9d Mon Sep 17 00:00:00 2001 From: Scott McKay Date: Fri, 30 Apr 2021 07:24:32 +1000 Subject: [PATCH] Add another input validation to ReverseSequence (#7445) * Add another input validation to ReverseSequence * Limit the bad length test to the CPU EP --- .../providers/cpu/tensor/reverse_sequence.cc | 8 ++++- .../cpu/tensor/reverse_sequence_test.cc | 30 +++++++++++++++++++ 2 files changed, 37 insertions(+), 1 deletion(-) diff --git a/onnxruntime/core/providers/cpu/tensor/reverse_sequence.cc b/onnxruntime/core/providers/cpu/tensor/reverse_sequence.cc index 4022f1460e..493ff62b1d 100644 --- a/onnxruntime/core/providers/cpu/tensor/reverse_sequence.cc +++ b/onnxruntime/core/providers/cpu/tensor/reverse_sequence.cc @@ -139,8 +139,14 @@ static Status ReverseSequenceImpl(const Tensor& X, for (int i = 0; i < batch_size; i++) { int64_t seq_len = sequence_lengths[i]; - if (seq_len == 0) + if (seq_len == 0) { continue; + } + + if (seq_len > max_seq_len || seq_len < 0) { + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Invalid sequence length: ", seq_len, + ". Value must be in range [0,", max_seq_len, "]"); + } for (int64_t j = 0; j < seq_len; j++) { gsl::span src = inputs.subspan(input_offset(max_seq_len, batch_size, input_size, i, j), input_size); diff --git a/onnxruntime/test/providers/cpu/tensor/reverse_sequence_test.cc b/onnxruntime/test/providers/cpu/tensor/reverse_sequence_test.cc index a5961f5042..6da1766e10 100644 --- a/onnxruntime/test/providers/cpu/tensor/reverse_sequence_test.cc +++ b/onnxruntime/test/providers/cpu/tensor/reverse_sequence_test.cc @@ -3,6 +3,7 @@ #include "gtest/gtest.h" #include "test/providers/provider_test_utils.h" +#include "test/util/include/default_providers.h" namespace onnxruntime { namespace test { @@ -144,5 +145,34 @@ TEST(ReverseSequenceTest, InvalidInput) { } } +TEST(ReverseSequenceTest, BadLength) { + auto run_test = [](bool use_negative) { + OpTester test("ReverseSequence", 10); + std::vector input = {0, 1, 2, 3, + 4, 5, 6, 7}; + + std::vector sequence_lens = {4, 3}; + + // make sequence_lens invalid for the input + sequence_lens[1] = use_negative ? -2 : 6; + + test.AddAttribute("batch_axis", int64_t(0)); + test.AddAttribute("time_axis", int64_t(1)); + + test.AddInput("input", {2, 4, 1}, input); + test.AddInput("sequence_lens", {2}, sequence_lens); + test.AddOutput("Y", {0}, {}); + + // the bad length check is just in the CPU EP + std::vector> eps; + eps.push_back(DefaultCpuExecutionProvider()); + + test.Run(OpTester::ExpectResult::kExpectFailure, "Invalid sequence length", {}, nullptr, &eps); + }; + + run_test(true); + run_test(false); +} + } // namespace test } // namespace onnxruntime