Skip to content
Merged
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
22 changes: 17 additions & 5 deletions onnxruntime/core/providers/cpu/tensor/split.h
Original file line number Diff line number Diff line change
Expand Up @@ -61,11 +61,23 @@ class SplitBase {
}
split_sizes = std::vector<int64_t>(static_cast<size_t>(num_outputs), split_dim_size / num_outputs);
} else {
int64_t split_size_sum = split_size_sum_;
if (split_size_sum == -1) {
split_size_sum = std::accumulate(split_sizes.cbegin(), split_sizes.cend(), 0LL);
int64_t remaining_split_size = split_dim_size;
for (int64_t s : split_sizes) {
if (s < 0) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Invalid negative value in 'split'. All split sizes must be >= 0.");
Comment thread
chilo-ms marked this conversation as resolved.
}
if (s > remaining_split_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Invalid value in 'split'. Split size ", s,
" exceeds the remaining size of the selected axis, ", remaining_split_size, ".");
}
remaining_split_size -= s;
}
if (split_sizes.size() != static_cast<size_t>(num_outputs) || split_size_sum != split_dim_size)

const int64_t split_size_sum =
split_size_sum_ == -1 ? split_dim_size - remaining_split_size : split_size_sum_;
if (split_sizes.size() != static_cast<size_t>(num_outputs) || remaining_split_size != 0)
return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL,
"Cannot split using values in 'split' attribute. Axis=", axis_,
" Input shape=", input_shape,
Expand All @@ -86,9 +98,9 @@ class SplitBase {
if (num_inputs == 1) {
// optional
if (info.GetAttrs("split", split_sizes_).IsOK()) {
split_size_sum_ = std::accumulate(split_sizes_.cbegin(), split_sizes_.cend(), 0LL);
ORT_ENFORCE(std::all_of(split_sizes_.cbegin(), split_sizes_.cend(), [](int64_t value) { return value >= 0; }),
"Invalid value in 'split' attribute. All values must be > 0");
split_size_sum_ = std::accumulate(split_sizes_.cbegin(), split_sizes_.cend(), SafeInt<int64_t>{0});
}
}

Expand Down
19 changes: 15 additions & 4 deletions onnxruntime/core/providers/cuda/tensor/split.cc
Original file line number Diff line number Diff line change
Expand Up @@ -85,11 +85,22 @@ Status SplitKernel::PrepareForComputeLocal(const TensorShape& input_shape,
}
split_sizes = std::vector<int64_t>(static_cast<size_t>(num_outputs), split_dim_size / num_outputs);
} else {
int64_t split_size_sum = split_size_sum_;
if (split_size_sum == -1) {
split_size_sum = std::accumulate(split_sizes.cbegin(), split_sizes.cend(), 0LL);
int64_t remaining_split_size = split_dim_size;
for (int64_t s : split_sizes) {
if (s < 0) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Invalid negative value in 'split'. All split sizes must be >= 0.");
Comment thread
chilo-ms marked this conversation as resolved.
}
if (s > remaining_split_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Invalid value in 'split'. Split size ", s,
" exceeds the remaining size of the selected axis, ", remaining_split_size, ".");
}
remaining_split_size -= s;
}
if (split_sizes.size() != static_cast<size_t>(num_outputs) || split_size_sum != split_dim_size) {

const int64_t split_size_sum = split_dim_size - remaining_split_size;
if (split_sizes.size() != static_cast<size_t>(num_outputs) || remaining_split_size != 0) {
return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL,
"Cannot split using values in 'split' attribute. Axis=", axis_,
" Input shape=", input_shape,
Expand Down
19 changes: 19 additions & 0 deletions onnxruntime/test/providers/cpu/tensor/split_op_test.cc
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.

#include <limits>

#include "gtest/gtest.h"
#include "core/framework/to_tensor_proto_element_type.h"
#include "test/providers/provider_test_utils.h"
Expand Down Expand Up @@ -949,5 +951,22 @@ TEST(SplitOperatorTest, InvalidValueInSplitInput_NegativeEntry_NegativeAxis) {
{}, nullptr, &execution_providers);
}

TEST(SplitOperatorTest, InvalidValueInSplitInput_Overflow) {
OpTester test("Split", 13, onnxruntime::kOnnxDomain);
test.AddAttribute<int64_t>("axis", 0);
test.AddInput<float>("input", {4, 2}, {1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f});
test.AddInput<int64_t>("split", {3}, {6, std::numeric_limits<int64_t>::max(), std::numeric_limits<int64_t>::max()},
/*is_initializer=*/false);
test.AddOutput<float>("output0", {1, 2}, {0.f, 0.f});
test.AddOutput<float>("output1", {1, 2}, {0.f, 0.f});
test.AddOutput<float>("output2", {1, 2}, {0.f, 0.f});

std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultCpuExecutionProvider());
test.Run(OpTester::ExpectResult::kExpectFailure,
"exceeds the remaining size of the selected axis",
{}, nullptr, &execution_providers);
}

} // namespace test
} // namespace onnxruntime
Loading