Fix Split CUDA implementation for zero sized input (#2942)

* Fix Split CUDA implementation for zero sized input

* resolve comments

* add case

* test case update: split into 2 tensors
This commit is contained in:
Yulong Wang 2020-04-07 14:44:20 -07:00 committed by GitHub
parent 48e96ea65f
commit aabf47b107
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
2 changed files with 37 additions and 26 deletions

View file

@ -64,34 +64,36 @@ Status Split::ComputeInternal(OpKernelContext* ctx) const {
}
}
output_ptr.CopyToGpu();
if (input_tensor->Shape().Size() > 0) {
output_ptr.CopyToGpu();
CudaAsyncBuffer<int64_t> split_sizes_gpu(this, split_sizes);
split_sizes_gpu.CopyToGpu();
CudaAsyncBuffer<int64_t> split_sizes_gpu(this, split_sizes);
split_sizes_gpu.CopyToGpu();
std::vector<int64_t> split_sizes_range(split_sizes);
for (size_t i = 1; i < split_sizes_range.size(); ++i) {
split_sizes_range[i] += split_sizes_range[i - 1];
std::vector<int64_t> split_sizes_range(split_sizes);
for (size_t i = 1; i < split_sizes_range.size(); ++i) {
split_sizes_range[i] += split_sizes_range[i - 1];
}
CudaAsyncBuffer<int64_t> split_sizes_range_gpu(this, split_sizes_range);
split_sizes_range_gpu.CopyToGpu();
CudaAsyncBuffer<int64_t> axis_dimension_input_output_mapping_gpu(this, axis_dimension_input_output_mapping);
axis_dimension_input_output_mapping_gpu.CopyToGpu();
size_t element_size = input_tensor->DataType()->Size();
ORT_RETURN_IF_ERROR(SplitImpl(element_size,
block_size_including_axis_dim,
block_size_inside_axis_dim,
split_sizes_gpu.GpuPtr(),
split_sizes_range_gpu.GpuPtr(),
axis_dimension_input_output_mapping_gpu.GpuPtr(),
num_outputs,
input_data,
output_ptr.GpuPtr(),
input_shape.Size()));
}
CudaAsyncBuffer<int64_t> split_sizes_range_gpu(this, split_sizes_range);
split_sizes_range_gpu.CopyToGpu();
CudaAsyncBuffer<int64_t> axis_dimension_input_output_mapping_gpu(this, axis_dimension_input_output_mapping);
axis_dimension_input_output_mapping_gpu.CopyToGpu();
size_t element_size = input_tensor->DataType()->Size();
ORT_RETURN_IF_ERROR(SplitImpl(element_size,
block_size_including_axis_dim,
block_size_inside_axis_dim,
split_sizes_gpu.GpuPtr(),
split_sizes_range_gpu.GpuPtr(),
axis_dimension_input_output_mapping_gpu.GpuPtr(),
num_outputs,
input_data,
output_ptr.GpuPtr(),
input_shape.Size()));
return Status::OK();
}

View file

@ -88,11 +88,11 @@ static void SplitTestInt() {
}
TEST(SplitOperatorTest, Axis0EqualSplitInt32) {
SplitTestInt<int32_t>();
SplitTestInt<int32_t>();
}
TEST(SplitOperatorTest, Axis0EqualSplitInt64) {
SplitTestInt<int64_t>();
SplitTestInt<int64_t>();
}
TEST(SplitOperatorTest, Axis0EqualSplitString) {
@ -322,6 +322,15 @@ TEST(SplitOperatorTest, Axis2UnequalSplit) {
RunTest<float>(axis, splits, input, outputs, false);
}
TEST(SplitOperatorTest, ZeroSizeInput) {
const int64_t axis = -1;
std::vector<ShapeAndFloatData> outputs{{{0, 1}, {}}, {{0, 1}, {}}};
ShapeAndFloatData input = CreateInput({0, 2});
RunTest<float>(axis, {}, input, outputs, false);
}
// test a split of a dimension that has leading and trailing dimensions
TEST(SplitOperatorTest, Axis1SplitMiddleDimensionEqually) {
const int64_t axis = 1;