Add Split op CUDA implementation (#964)

This commit is contained in:
Hector Li 2019-05-03 14:40:19 -07:00 committed by GitHub
parent f4fd36ee91
commit ab3355f6b4
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
7 changed files with 306 additions and 45 deletions

View file

@ -21,6 +21,48 @@ ONNX_CPU_OPERATOR_KERNEL(
}),
Split);
Status SplitBase::PrepareForCompute(const TensorShape& input_shape,
const int num_outputs,
int64_t& axis,
int& before_dims,
int& after_dims_including_split_axis,
int& after_dims_excluding_split,
std::vector<int64_t>& split_sizes) const {
auto& input_dims = input_shape.GetDims();
const int64_t num_dimensions = gsl::narrow_cast<int64_t>(input_shape.NumDimensions());
axis = HandleNegativeAxis(axis_, num_dimensions); // handle negative and enforce axis is valid
const int64_t split_dim_size = input_dims[axis];
before_dims = gsl::narrow<int>(input_shape.SizeToDimension(axis));
after_dims_including_split_axis = gsl::narrow<int>(input_shape.SizeFromDimension(axis));
after_dims_excluding_split = (axis + 1 == num_dimensions)
? 1 // we multiply by this value so must be 1 not 0
: gsl::narrow<int>(input_shape.SizeFromDimension(axis + 1));
if (split_sizes_.empty()) {
// equal split based on number of outputs
if (split_dim_size % num_outputs != 0) {
return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Input cannot be split evenly on selected axis. Input shape=", input_shape,
" Axis=", axis_, " NumOutputs=", num_outputs);
}
// populate split_sizes with the same size for each output
split_sizes = std::vector<int64_t>(num_outputs, split_dim_size / num_outputs);
} else {
if (split_sizes_.size() != num_outputs || split_size_sum_ != split_dim_size)
return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL,
"Cannot split using values in 'split' attribute. Axis=", axis_,
" Input shape=", input_shape,
" NumOutputs=", num_outputs,
" Num entries in 'split' (must equal number of outputs) was ", split_sizes_.size(),
" Sum of sizes in 'split' (must equal size of selected axis) was ", split_size_sum_);
split_sizes = split_sizes_;
}
return Status::OK();
}
Status Split::Compute(OpKernelContext* context) const {
const Tensor& input = *context->Input<Tensor>(0);
@ -44,45 +86,23 @@ Status Split::Compute(OpKernelContext* context) const {
template <typename T>
Status Split::ComputeImpl(OpKernelContext& context, const Tensor& input) const {
auto& input_shape = input.Shape();
auto& input_dims = input_shape.GetDims();
const int64_t num_dimensions = gsl::narrow_cast<int64_t>(input_shape.NumDimensions());
const int64_t axis = HandleNegativeAxis(axis_, num_dimensions); // handle negative and enforce axis is valid
const int64_t split_dim_size = input_dims[axis];
auto num_outputs = context.OutputCount();
std::vector<Tensor*> outputs;
outputs.reserve(num_outputs);
int before_dims = gsl::narrow<int>(input_shape.SizeToDimension(axis));
int after_dims_including_split_axis = gsl::narrow<int>(input_shape.SizeFromDimension(axis));
int after_dims_excluding_split = (axis + 1 == num_dimensions)
? 1 // we multiply by this value so must be 1 not 0
: gsl::narrow<int>(input_shape.SizeFromDimension(axis + 1));
int64_t axis = axis_;
int before_dims = 0;
int after_dims_including_split_axis = 0;
int after_dims_excluding_split = 0;
std::vector<int64_t> split_sizes;
if (split_sizes_.empty()) {
// equal split based on number of outputs
if (split_dim_size % num_outputs != 0) {
return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Input cannot be split evenly on selected axis. Input shape=", input_shape,
" Axis=", axis_, " NumOutputs=", num_outputs);
}
// populate split_sizes with the same size for each output
split_sizes = std::vector<int64_t>(num_outputs, split_dim_size / num_outputs);
} else {
if (split_sizes_.size() != num_outputs || split_size_sum_ != split_dim_size)
return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL,
"Cannot split using values in 'split' attribute. Axis=", axis_,
" Input shape=", input_shape,
" NumOutputs=", num_outputs,
" Num entries in 'split' (must equal number of outputs) was ", split_sizes_.size(),
" Sum of sizes in 'split' (must equal size of selected axis) was ", split_size_sum_);
split_sizes = split_sizes_;
}
ORT_RETURN_IF_ERROR(PrepareForCompute(input_shape,
num_outputs,
axis,
before_dims,
after_dims_including_split_axis,
after_dims_excluding_split,
split_sizes));
// copy dimensions so we can update the selected axis in place
auto& input_dims = input_shape.GetDims();
std::vector<int64_t> output_dimensions{input_dims};
int64_t input_offset = 0;

View file

@ -10,12 +10,10 @@
namespace onnxruntime {
class Split final : public OpKernel {
public:
Split(const OpKernelInfo& info) : OpKernel(info) {
// required with default of 0
if (!info.GetAttr("axis", &axis_).IsOK())
ORT_THROW("Missing 'axis' attribute value");
class SplitBase {
protected:
SplitBase(const OpKernelInfo& info) {
axis_ = info.GetAttrOrDefault<int64_t>("axis", 0);
// optional
if (info.GetAttrs("split", split_sizes_).IsOK()) {
@ -25,15 +23,28 @@ class Split final : public OpKernel {
}
}
Status Compute(OpKernelContext* context) const override;
private:
template <typename T>
Status ComputeImpl(OpKernelContext& context, const Tensor& input) const;
Status PrepareForCompute(const TensorShape& input_shape,
const int num_outputs,
int64_t& axis,
int& before_dims,
int& after_dims_including_split_axis,
int& after_dims_excluding_split,
std::vector<int64_t>& split_sizes) const;
int64_t axis_;
std::vector<int64_t> split_sizes_;
int64_t split_size_sum_ = 0;
};
class Split final : public OpKernel, public SplitBase {
public:
Split(const OpKernelInfo& info) : OpKernel(info), SplitBase(info) {}
Status Compute(OpKernelContext* context) const override;
private:
template <typename T>
Status ComputeImpl(OpKernelContext& context, const Tensor& input) const;
};
} // namespace onnxruntime

View file

@ -523,6 +523,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 10, MLFloat16, Resize);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 10, int32_t, Resize);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 10, uint8_t, Resize);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 2, Split);
static void RegisterCudaKernels(KernelRegistry& kernel_registry) {
static const BuildKernelCreateInfoFn function_table[] = {
@ -809,6 +810,7 @@ static void RegisterCudaKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 10, MLFloat16, Resize)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 10, int32_t, Resize)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 10, uint8_t, Resize)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 2, Split)>,
};
for (auto& function_table_entry : function_table) {

View file

@ -0,0 +1,83 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include "split.h"
#include "split_impl.h"
#include "core/providers/cpu/tensor/utils.h"
#include "core/providers/common.h"
namespace onnxruntime {
namespace cuda {
ONNX_OPERATOR_KERNEL_EX(
Split,
kOnnxDomain,
2,
kCudaExecutionProvider,
KernelDefBuilder()
.TypeConstraint("T", DataTypeImpl::AllFixedSizeTensorTypes()),
Split);
Status Split::ComputeInternal(OpKernelContext* ctx) const {
const Tensor* input_tensor = ctx->Input<Tensor>(0);
ORT_ENFORCE(nullptr != input_tensor);
auto& input_shape = input_tensor->Shape();
auto num_outputs = ctx->OutputCount();
int64_t axis = axis_;
int before_dims = 0;
int block_size_including_axis_dim = 0;
int block_size_inside_axis_dim = 0;
std::vector<int64_t> split_sizes;
ORT_RETURN_IF_ERROR(PrepareForCompute(input_shape,
num_outputs,
axis,
before_dims,
block_size_including_axis_dim,
block_size_inside_axis_dim,
split_sizes));
auto input_data = input_tensor->DataRaw();
auto& input_dims = input_shape.GetDims();
std::vector<int64_t> output_dimensions{input_dims};
int device_id = 0;
CudaAsyncBuffer<void*> output_ptr(this, device_id, num_outputs);
gsl::span<void*> output_ptr_span = output_ptr.CpuSpan();
for (int i = 0; i < num_outputs; ++i) {
// update size of dimension for axis we're splitting on
auto split_size = gsl::narrow<int>(split_sizes[i]);
output_dimensions[axis] = split_size;
Tensor* output = ctx->Output(i, TensorShape{output_dimensions});
auto output_data = output->MutableDataRaw();
output_ptr_span[i] = output_data;
}
output_ptr.CopyToGpu();
CudaAsyncBuffer<int64_t> split_sizes_gpu(this, device_id, split_sizes);
split_sizes_gpu.CopyToGpu();
std::vector<int64_t> split_sizes_range(split_sizes);
for (int 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, device_id, split_sizes_range);
split_sizes_range_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(),
num_outputs,
input_data,
output_ptr.GpuPtr(),
input_shape.Size()));
return Status::OK();
}
} // namespace cuda
} // namespace onnxruntime

View file

@ -0,0 +1,18 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include "core/common/common.h"
#include "core/providers/cuda/cuda_common.h"
#include "core/providers/cpu/tensor/split.h"
namespace onnxruntime {
namespace cuda {
class Split final : public CudaKernel, public SplitBase {
public:
Split(const OpKernelInfo& info) : CudaKernel(info), SplitBase(info) {}
Status ComputeInternal(OpKernelContext* context) const override;
};
} // namespace cuda
} // namespace onnxruntime

View file

@ -0,0 +1,104 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include "core/providers/cuda/cu_inc/common.cuh"
#include "core/providers/cuda/cuda_common.h"
#include "split_impl.h"
namespace onnxruntime {
namespace cuda {
template <typename T>
__global__ void _SplitKernel(const fast_divmod block_size_including_axis_dim_div,
const fast_divmod block_size_inside_axis_dim_div,
const int64_t* split_sizes,
const int64_t* split_sizes_range,
const int num_outputs,
const T* input_data,
void** output_ptr,
const CUDA_LONG N) {
CALCULATE_ELEMENTWISE_INDEX_OR_EXIT(id, N);
CUDA_LONG output_pos = 0;
int outter_block_index = 0;
int block_index = 0;
int offset = 0;
int output_index = 0;
int block_offset = 0;
block_size_including_axis_dim_div.divmod(id, outter_block_index, offset);
block_size_inside_axis_dim_div.divmod(offset, block_index, offset);
for (int i = 0; i < num_outputs; ++i) {
int64_t range_left = (i == 0) ? 0 : split_sizes_range[i - 1];
if ((range_left <= block_index) && (block_index < split_sizes_range[i])) {
output_index = i;
block_offset = block_index - range_left;
break;
}
}
output_pos = (outter_block_index * split_sizes[output_index] + block_offset) *
block_size_inside_axis_dim_div.d_ +
offset;
reinterpret_cast<T*>(output_ptr[output_index])[output_pos] = input_data[id];
}
Status SplitImpl(const size_t element_size,
const int block_size_including_axis_dim,
const int block_size_inside_axis_dim,
const int64_t* split_sizes,
const int64_t* split_sizes_range,
const int num_outputs,
const void* input_data,
void** output_ptr,
const size_t N) {
int blocksPerGrid = (int)(ceil(static_cast<float>(N) / GridDim::maxThreadsPerBlock));
fast_divmod block_size_including_axis_dim_div = fast_divmod(block_size_including_axis_dim);
fast_divmod block_size_inside_axis_dim_div = fast_divmod(block_size_inside_axis_dim);
switch (element_size) {
case sizeof(int8_t):
_SplitKernel<<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
block_size_including_axis_dim_div, block_size_inside_axis_dim_div,
split_sizes, split_sizes_range, num_outputs,
reinterpret_cast<const ToCudaType<int8_t>::MappedType*>(input_data),
output_ptr,
(CUDA_LONG)N);
break;
case sizeof(int16_t):
_SplitKernel<<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
block_size_including_axis_dim_div, block_size_inside_axis_dim_div,
split_sizes, split_sizes_range, num_outputs,
reinterpret_cast<const ToCudaType<int16_t>::MappedType*>(input_data),
output_ptr,
(CUDA_LONG)N);
break;
case sizeof(int32_t):
_SplitKernel<<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
block_size_including_axis_dim_div, block_size_inside_axis_dim_div,
split_sizes, split_sizes_range, num_outputs,
reinterpret_cast<const ToCudaType<int32_t>::MappedType*>(input_data),
output_ptr,
(CUDA_LONG)N);
break;
case sizeof(int64_t):
_SplitKernel<<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
block_size_including_axis_dim_div, block_size_inside_axis_dim_div,
split_sizes, split_sizes_range, num_outputs,
reinterpret_cast<const ToCudaType<int64_t>::MappedType*>(input_data),
output_ptr,
(CUDA_LONG)N);
break;
default:
return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Type not supported for Slice operator");
}
return Status::OK();
}
} // namespace cuda
} // namespace onnxruntime

View file

@ -0,0 +1,23 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#pragma once
#include <stdint.h>
#include "core/providers/cuda/shared_inc/cuda_utils.h"
#include "core/common/common.h"
namespace onnxruntime {
namespace cuda {
Status SplitImpl(const size_t element_size,
const int block_size_including_axis_dim,
const int block_size_inside_axis_dim,
const int64_t* split_sizes,
const int64_t* split_sizes_range,
const int num_outputs,
const void* input_data,
void** output_ptr,
const size_t N);
} // namespace cuda
} // namespace onnxruntime