SplitTraining op to support split as input (#4597)

* SplitTraining op to support split as input

* on comments and minor refactor

Co-authored-by: Ethan Tao <ettao@microsoft.com>
This commit is contained in:
ytaous 2020-07-24 12:49:19 -07:00 committed by GitHub
parent aa328c2c20
commit 9888c9e944
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
9 changed files with 521 additions and 2 deletions

View file

@ -1,8 +1,8 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include "split.h"
#include "split_impl.h"
#include "core/providers/cuda/tensor/split.h"
#include "core/providers/cuda/tensor/split_impl.h"
#include "core/providers/cpu/tensor/utils.h"
#include "core/providers/common.h"

View file

@ -1128,6 +1128,86 @@ Example 4:
}
}
});
ONNX_CONTRIB_OPERATOR_SCHEMA(SplitTraining)
.SetDomain(kMSDomain)
.SinceVersion(1)
.SetSupportLevel(OpSchema::SupportType::EXPERIMENTAL)
.SetDoc("SplitTraining")
.Attr("axis",
"Which axis to split on. "
"A negative value means counting dimensions from the back. Accepted range is [-rank, rank-1] "
"where r = rank(input).",
AttributeProto::INT,
static_cast<int64_t>(0))
.AllowUncheckedAttributes()
.Input(0, "input", "The tensor to split", "T")
.Input(1, "split", "length of each output", "tensor(int64)")
.Output(0,
"outputs",
"One or more outputs forming list of tensors after splitting",
"T",
OpSchema::Variadic)
.TypeConstraint(
"T",
OpSchema::all_tensor_types(),
"Constrain input and output types to all tensor types.")
.TypeAndShapeInferenceFunction([](ONNX_NAMESPACE::InferenceContext& ctx) {
for (int i = 0; i < static_cast<int>(ctx.getNumOutputs()); ++i) {
propagateElemTypeFromInputToOutput(ctx, 0, i);
}
if (!hasNInputShapes(ctx, 1)) {
return;
}
// skip if split is not an initializer
auto split_proto = ctx.getInputData(1);
if (split_proto == nullptr) {
return;
}
std::vector<int64_t> split = ParseData<int64_t>(split_proto);
if (!ctx.getInputType(0)->tensor_type().has_shape()) {
return;
}
const auto& shape = ctx.getInputType(0)->tensor_type().shape();
int rank = shape.dim_size();
int axis = static_cast<int>(getAttribute(ctx, "axis", 0));
if (axis < -rank || axis >= rank) {
fail_type_inference(
"Invalid value of attribute 'axis'. Rank=",
rank,
" Value=",
axis);
}
if (axis < 0) {
axis += rank;
}
const auto& splitDim = shape.dim(axis);
if (!splitDim.has_dim_value()) {
return;
}
int splitDimValue = static_cast<int>(splitDim.dim_value());
if (split.empty()) {
int chunkSize =
splitDimValue / static_cast<int>(ctx.getNumOutputs());
int leftOver = splitDimValue -
(chunkSize * static_cast<int>(ctx.getNumOutputs()));
for (int i = 0; i < static_cast<int>(ctx.getNumOutputs()); i++) {
split.push_back(i < leftOver ? chunkSize + 1 : chunkSize);
}
}
for (size_t i = 0; i < ctx.getNumOutputs(); i++) {
*ctx.getOutputType(i)->mutable_tensor_type()->mutable_shape() =
shape;
ctx.getOutputType(i)
->mutable_tensor_type()
->mutable_shape()
->mutable_dim(axis)
->set_dim_value(split[i]);
}
});
ONNX_CONTRIB_OPERATOR_SCHEMA(ConcatTraining)
.SetDomain(kMSDomain)

View file

@ -0,0 +1,132 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include "gtest/gtest.h"
#include "test/providers/provider_test_utils.h"
namespace onnxruntime {
namespace test {
template <class T>
using ShapeAndData = std::pair<const std::vector<int64_t>, const std::vector<T>>;
using ShapeAndFloatData = ShapeAndData<float>;
using ShapeAndStringData = ShapeAndData<std::string>;
using ExpectResult = OpTester::ExpectResult;
template <typename T>
void SplitTrainingOpTester(int64_t axis, const std::vector<int64_t> split_sizes, const ShapeAndData<T>& input,
const std::vector<ShapeAndData<T>>& outputs, bool is_initializer = true,
bool expect_failure = false, const std::string& err_msg = {}) {
OpTester test("SplitTraining", 1, onnxruntime::kMSDomain);
test.AddAttribute("axis", axis);
test.AddInput<T>("input", input.first, input.second);
test.AddInput<int64_t>("split", {static_cast<int64_t>(split_sizes.size())}, split_sizes, is_initializer);
int i = 0;
for (auto& output : outputs) {
auto& shape = output.first;
auto& data = output.second;
std::ostringstream oss;
oss << "output" << i++;
test.AddOutput<T>(oss.str().c_str(), shape, data);
}
test.Run(expect_failure ? ExpectResult::kExpectFailure : ExpectResult::kExpectSuccess, err_msg);
}
TEST(SplitTrainingOpTest, Axis0EqualSplitFloat) {
const int64_t axis = 0;
std::vector<ShapeAndFloatData> outputs;
// input shape and data
ShapeAndFloatData input = {{4, 2}, // shape
{1.f, 2.f,
3.f, 4.f,
5.f, 6.f,
7.f, 8.f}};
outputs.push_back({{2, 2},
{1.f, 2.f,
3.f, 4.f}});
outputs.push_back({{2, 2},
{5.f, 6.f,
7.f, 8.f}});
SplitTrainingOpTester<float>(axis, {}, input, outputs);
}
TEST(SplitTrainingOpTest, Axis0UnequalSplitFloat) {
const int64_t axis = 0;
std::vector<ShapeAndFloatData> outputs;
// input shape and data
ShapeAndFloatData input = {{4, 2}, // shape
{1.f, 2.f,
3.f, 4.f,
5.f, 6.f,
7.f, 8.f}};
std::vector<int64_t> splits{1, 3};
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}});
SplitTrainingOpTester<float>(axis, splits, input, outputs);
}
TEST(SplitTrainingOpTest, Axis0EqualSplitFloat_not_initializer) {
const int64_t axis = 0;
std::vector<ShapeAndFloatData> outputs;
// input shape and data
ShapeAndFloatData input = {{4, 2}, // shape
{1.f, 2.f,
3.f, 4.f,
5.f, 6.f,
7.f, 8.f}};
outputs.push_back({{2, 2},
{1.f, 2.f,
3.f, 4.f}});
outputs.push_back({{2, 2},
{5.f, 6.f,
7.f, 8.f}});
SplitTrainingOpTester<float>(axis, {}, input, outputs, false);
}
TEST(SplitTrainingOpTest, Axis0UnequalSplitFloat_not_initializer) {
const int64_t axis = 0;
std::vector<ShapeAndFloatData> outputs;
// input shape and data
ShapeAndFloatData input = {{4, 2}, // shape
{1.f, 2.f,
3.f, 4.f,
5.f, 6.f,
7.f, 8.f}};
std::vector<int64_t> splits{1, 3};
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}});
SplitTrainingOpTester<float>(axis, splits, input, outputs, false);
}
} // namespace test
} // namespace onnxruntime

View file

@ -21,6 +21,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, double, ReduceSumTraining);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, int32_t, ReduceSumTraining);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, int64_t, ReduceSumTraining);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SplitTraining);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, ConcatTraining);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SoftmaxCrossEntropy);
@ -111,6 +112,7 @@ Status RegisterCpuTrainingKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, double, ReduceSumTraining)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, int32_t, ReduceSumTraining)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, int64_t, ReduceSumTraining)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SplitTraining)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, ConcatTraining)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SoftmaxCrossEntropy)>,

View file

@ -0,0 +1,153 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include "orttraining/training_ops/cpu/tensor/split.h"
#include "core/providers/common.h"
#include "core/util/math.h"
#include "core/util/math_cpuonly.h"
#include "gsl/gsl"
namespace onnxruntime {
namespace contrib {
ONNX_OPERATOR_KERNEL_EX(
SplitTraining,
kMSDomain,
1,
kCpuExecutionProvider,
KernelDefBuilder()
.TypeConstraint("T", DataTypeImpl::AllTensorTypes()),
SplitTraining);
Status PrepareForTrainingCompute(const TensorShape& input_shape, 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) {
auto& input_dims = input_shape.GetDims();
const auto num_dimensions = gsl::narrow_cast<int64_t>(input_shape.NumDimensions());
int64_t axis_value = axis;
axis = HandleNegativeAxis(axis_value, 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));
std::vector<int64_t> split_sizes_values(split_sizes);
split_sizes.clear();
int64_t split_size_sum = std::accumulate(split_sizes_values.cbegin(), split_sizes_values.cend(), 0LL);
if (split_sizes_values.empty()) {
// equal split based on number of outputs
if (split_dim_size % static_cast<size_t>(num_outputs) != 0) {
return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Input cannot be split evenly on selected axis. Input shape=", input_shape,
" Axis=", axis_value, " NumOutputs=", num_outputs);
}
// populate split_sizes with the same size for each output
split_sizes = std::vector<int64_t>(static_cast<size_t>(num_outputs), split_dim_size / num_outputs);
} else {
if (split_sizes_values.size() != static_cast<size_t>(num_outputs) || split_size_sum != split_dim_size)
return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL,
"Cannot split using values in 'split' input. Axis=", axis_value,
" Input shape=", input_shape,
" NumOutputs=", num_outputs,
" Num entries in 'split' (must equal number of outputs) was ", split_sizes_values.size(),
" Sum of sizes in 'split' (must equal size of selected axis) was ", split_size_sum);
split_sizes = split_sizes_values;
}
return Status::OK();
}
Status SplitTraining::Compute(OpKernelContext* context) const {
const Tensor& input = *context->Input<Tensor>(0);
Status status;
if (input.IsDataType<float>())
status = ComputeImpl<float>(*context, input);
else if (input.IsDataType<int32_t>())
status = ComputeImpl<int32_t>(*context, input);
else if (input.IsDataType<int64_t>())
status = ComputeImpl<int64_t>(*context, input);
else if (input.IsDataTypeString())
status = ComputeImpl<std::string>(*context, input);
else
ORT_THROW("Split operator does not support ", input.DataType(), " yet");
return status;
}
template <typename T>
inline void copy_data(const T* src, T* dst, size_t count) {
memcpy(dst, src, count * sizeof(T));
}
template <>
inline void copy_data<std::string>(const std::string* src, std::string* dst, size_t count) {
const std::string* end = src + count;
std::copy(src, end, dst);
}
template <typename T>
Status SplitTraining::ComputeImpl(OpKernelContext& context, const Tensor& input) const {
auto& input_shape = input.Shape();
auto num_outputs = context.OutputCount();
int64_t axis = axis_;
int before_dims = 0;
int after_dims_including_split_axis = 0;
int after_dims_excluding_split = 0;
//override the attribute value with the input value for split_split
const Tensor* split_tensor = context.Input<Tensor>(1);
ORT_ENFORCE(split_tensor->Shape().NumDimensions() == 1, "An split tensor must be a vector tensor.");
auto nDims = static_cast<size_t>(split_tensor->Shape()[0]);
const auto* data = split_tensor->template Data<int64_t>();
std::vector<int64_t> split_sizes(data, data + nDims);
ORT_RETURN_IF_ERROR(PrepareForTrainingCompute(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;
const T* input_data = input.template Data<T>();
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 = context.Output(i, TensorShape{output_dimensions});
T* output_data = output->template MutableData<T>();
::onnxruntime::math::CopyMatrix<T>(
before_dims, // M
split_size * after_dims_excluding_split, // N
static_cast<const T*>(input_data + input_offset), // A
after_dims_including_split_axis, // lda
static_cast<T*>(output_data), // B
split_size * after_dims_excluding_split, // ldb
[](const T* src, T* dst, size_t count) {
copy_data<T>(src, dst, count);
});
input_offset += split_size * after_dims_excluding_split; // offset by the N data we used in this iteration
}
return Status::OK();
}
} // namespace contrib
} // namespace onnxruntime

View file

@ -0,0 +1,30 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#pragma once
#include <numeric>
#include "core/providers/cpu/tensor/split.h"
#include "core/common/common.h"
#include "core/framework/op_kernel.h"
namespace onnxruntime {
namespace contrib {
class SplitTraining final : public OpKernel, public SplitBase {
public:
SplitTraining(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;
};
Status PrepareForTrainingCompute(const TensorShape& input_shape, 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);
} // namespace contrib
} // namespace onnxruntime

View file

@ -16,6 +16,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, ReduceSumTraining);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, int32_t, ReduceSumTraining);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, ReduceSumTraining);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, SplitTraining);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, ConcatTraining);
// Adam
@ -134,6 +135,7 @@ Status RegisterCudaTrainingKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, ReduceSumTraining)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, int32_t, ReduceSumTraining)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, ReduceSumTraining)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, SplitTraining)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, ConcatTraining)>,
// Adam
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_int64_t_float_float_float_float, AdamOptimizer)>,

View file

@ -0,0 +1,101 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include "orttraining/training_ops/cuda/tensor/split.h"
#include "core/providers/cuda/tensor/split_impl.h"
#include "core/providers/cpu/tensor/utils.h"
#include "core/providers/common.h"
namespace onnxruntime {
namespace cuda {
ONNX_OPERATOR_KERNEL_EX(SplitTraining,
kMSDomain,
1,
kCudaExecutionProvider,
KernelDefBuilder()
.InputMemoryType<OrtMemTypeCPUInput>(1)
.TypeConstraint("T", DataTypeImpl::AllFixedSizeTensorTypes()),
SplitTraining);
Status SplitTraining::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 = HandleNegativeAxis(axis_, input_shape.NumDimensions());
int before_dims = 0;
int block_size_including_axis_dim = 0;
int block_size_inside_axis_dim = 0;
//override the attribute value with the input value for split_split
const Tensor* split_tensor = ctx->Input<Tensor>(1);
ORT_ENFORCE(split_tensor->Shape().NumDimensions() == 1, "An split tensor must be a vector tensor.");
auto nDims = static_cast<size_t>(split_tensor->Shape()[0]);
const auto* data = split_tensor->template Data<int64_t>();
std::vector<int64_t> split_sizes(data, data + nDims);
ORT_RETURN_IF_ERROR(onnxruntime::contrib::PrepareForTrainingCompute(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};
CudaAsyncBuffer<void*> output_ptr(this, num_outputs);
gsl::span<void*> output_ptr_span = output_ptr.CpuSpan();
std::vector<int64_t> axis_dimension_input_output_mapping(input_dims[axis]);
int index = 0;
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;
for (int j = 0; j < split_size; ++j) {
axis_dimension_input_output_mapping.at(index++) = i;
}
}
if (input_tensor->Shape().Size() > 0) {
output_ptr.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];
}
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();
}
} // namespace cuda
} // namespace onnxruntime

View file

@ -0,0 +1,19 @@
// 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"
#include "orttraining/training_ops/cpu/tensor/split.h"
namespace onnxruntime {
namespace cuda {
class SplitTraining final : public CudaKernel, public SplitBase {
public:
SplitTraining(const OpKernelInfo& info) : CudaKernel(info), SplitBase(info) {}
Status ComputeInternal(OpKernelContext* context) const override;
};
} // namespace cuda
} // namespace onnxruntime