mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-20 19:12:24 +00:00
Add Split op CUDA implementation (#964)
This commit is contained in:
parent
f4fd36ee91
commit
ab3355f6b4
7 changed files with 306 additions and 45 deletions
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
83
onnxruntime/core/providers/cuda/tensor/split.cc
Normal file
83
onnxruntime/core/providers/cuda/tensor/split.cc
Normal 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
|
||||
18
onnxruntime/core/providers/cuda/tensor/split.h
Normal file
18
onnxruntime/core/providers/cuda/tensor/split.h
Normal 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
|
||||
104
onnxruntime/core/providers/cuda/tensor/split_impl.cu
Normal file
104
onnxruntime/core/providers/cuda/tensor/split_impl.cu
Normal 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
|
||||
23
onnxruntime/core/providers/cuda/tensor/split_impl.h
Normal file
23
onnxruntime/core/providers/cuda/tensor/split_impl.h
Normal 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
|
||||
Loading…
Reference in a new issue