mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-24 19:43:35 +00:00
Allow zero in split op (#3389)
Allow zero in split op (A change in onnx 1.7 without bumping up the op version)
This commit is contained in:
parent
1671072b6b
commit
aefa466334
3 changed files with 6 additions and 9 deletions
|
|
@ -472,17 +472,14 @@ Status SplitToSequence::ComputeImpl(OpKernelContext& context, const Tensor& inpu
|
|||
is_split_input_scalar = true;
|
||||
} else {
|
||||
GetSplitSizesInput(*p_split_input, split_sizes);
|
||||
ORT_ENFORCE(std::all_of(split_sizes.cbegin(), split_sizes.cend(), [](int64_t value) { return value > 0; }),
|
||||
"Invalid value in 'split' input. All values must be > 0");
|
||||
ORT_ENFORCE(std::all_of(split_sizes.cbegin(), split_sizes.cend(), [](int64_t value) { return value >= 0; }),
|
||||
"Invalid value in 'split' input. All values must be >= 0");
|
||||
}
|
||||
}
|
||||
|
||||
// Keep the split dimension or not. Default 1, which means we keep split dimension.
|
||||
// If input 'split' is specified, this attribute is ignored.
|
||||
bool use_keep_dims = false;
|
||||
if (split_sizes.empty()) {
|
||||
use_keep_dims = true;
|
||||
}
|
||||
bool use_keep_dims = split_sizes.empty();
|
||||
|
||||
ORT_RETURN_IF_ERROR(PrepareForCompute(input_shape,
|
||||
split_scalar,
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ class SplitBase {
|
|||
// 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; }),
|
||||
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");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -429,7 +429,7 @@ TEST(SplitOperatorTest, SplitAttributeSumTooSmall) {
|
|||
}
|
||||
|
||||
TEST(SplitOperatorTest, InvalidValueInSplitAttribute) {
|
||||
const int64_t axis = 0;
|
||||
const int64_t axis = -1;
|
||||
std::vector<ShapeAndFloatData> outputs;
|
||||
|
||||
// input shape and data
|
||||
|
|
@ -443,7 +443,7 @@ TEST(SplitOperatorTest, InvalidValueInSplitAttribute) {
|
|||
outputs.push_back({{1, 2}, {1.f, 2.f}});
|
||||
outputs.push_back({{3, 2}, {3.f, 4.f, 5.f, 6.f, 7.f, 8.f}});
|
||||
|
||||
RunTest<float>(axis, splits, input, outputs, false, true, "Invalid value in 'split' attribute"); //TensorRT parser: Assertion failed: axis != BATCH_DIM
|
||||
RunTest<float>(axis, splits, input, outputs, false, true, "in 'split' attribute"); //TensorRT parser: Assertion failed: axis != BATCH_DIM
|
||||
}
|
||||
|
||||
/*
|
||||
|
|
|
|||
Loading…
Reference in a new issue