mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
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:
parent
aa328c2c20
commit
9888c9e944
9 changed files with 521 additions and 2 deletions
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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)>,
|
||||
|
|
|
|||
153
orttraining/orttraining/training_ops/cpu/tensor/split.cc
Normal file
153
orttraining/orttraining/training_ops/cpu/tensor/split.cc
Normal 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
|
||||
30
orttraining/orttraining/training_ops/cpu/tensor/split.h
Normal file
30
orttraining/orttraining/training_ops/cpu/tensor/split.h
Normal 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
|
||||
|
|
@ -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)>,
|
||||
|
|
|
|||
101
orttraining/orttraining/training_ops/cuda/tensor/split.cc
Normal file
101
orttraining/orttraining/training_ops/cuda/tensor/split.cc
Normal 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
|
||||
19
orttraining/orttraining/training_ops/cuda/tensor/split.h
Normal file
19
orttraining/orttraining/training_ops/cuda/tensor/split.h
Normal 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
|
||||
Loading…
Reference in a new issue