diff --git a/onnxruntime/core/providers/cuda/tensor/split.cc b/onnxruntime/core/providers/cuda/tensor/split.cc index 779cc62f2c..f447920d69 100644 --- a/onnxruntime/core/providers/cuda/tensor/split.cc +++ b/onnxruntime/core/providers/cuda/tensor/split.cc @@ -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" diff --git a/orttraining/orttraining/core/graph/training_op_defs.cc b/orttraining/orttraining/core/graph/training_op_defs.cc index 16438e4e05..1d7b0abb98 100644 --- a/orttraining/orttraining/core/graph/training_op_defs.cc +++ b/orttraining/orttraining/core/graph/training_op_defs.cc @@ -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(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(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 split = ParseData(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(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(splitDim.dim_value()); + if (split.empty()) { + int chunkSize = + splitDimValue / static_cast(ctx.getNumOutputs()); + int leftOver = splitDimValue - + (chunkSize * static_cast(ctx.getNumOutputs())); + for (int i = 0; i < static_cast(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) diff --git a/orttraining/orttraining/test/training_ops/cpu/tensor/split_op_test.cc b/orttraining/orttraining/test/training_ops/cpu/tensor/split_op_test.cc new file mode 100644 index 0000000000..b727df0950 --- /dev/null +++ b/orttraining/orttraining/test/training_ops/cpu/tensor/split_op_test.cc @@ -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 +using ShapeAndData = std::pair, const std::vector>; + +using ShapeAndFloatData = ShapeAndData; +using ShapeAndStringData = ShapeAndData; +using ExpectResult = OpTester::ExpectResult; + +template +void SplitTrainingOpTester(int64_t axis, const std::vector split_sizes, const ShapeAndData& input, + const std::vector>& 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("input", input.first, input.second); + test.AddInput("split", {static_cast(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(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 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(axis, {}, input, outputs); +} + +TEST(SplitTrainingOpTest, Axis0UnequalSplitFloat) { + const int64_t axis = 0; + std::vector 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 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(axis, splits, input, outputs); +} + + +TEST(SplitTrainingOpTest, Axis0EqualSplitFloat_not_initializer) { + const int64_t axis = 0; + std::vector 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(axis, {}, input, outputs, false); +} + +TEST(SplitTrainingOpTest, Axis0UnequalSplitFloat_not_initializer) { + const int64_t axis = 0; + std::vector 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 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(axis, splits, input, outputs, false); +} + +} // namespace test +} // namespace onnxruntime diff --git a/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc b/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc index bbfdcd41d8..ab37511dc4 100644 --- a/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc +++ b/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc @@ -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, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/orttraining/orttraining/training_ops/cpu/tensor/split.cc b/orttraining/orttraining/training_ops/cpu/tensor/split.cc new file mode 100644 index 0000000000..8332d58c0f --- /dev/null +++ b/orttraining/orttraining/training_ops/cpu/tensor/split.cc @@ -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& split_sizes) { + auto& input_dims = input_shape.GetDims(); + const auto num_dimensions = gsl::narrow_cast(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(input_shape.SizeToDimension(axis)); + after_dims_including_split_axis = gsl::narrow(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(input_shape.SizeFromDimension(axis + 1)); + + std::vector 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(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(static_cast(num_outputs), split_dim_size / num_outputs); + } else { + if (split_sizes_values.size() != static_cast(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(0); + + Status status; + + if (input.IsDataType()) + status = ComputeImpl(*context, input); + else if (input.IsDataType()) + status = ComputeImpl(*context, input); + else if (input.IsDataType()) + status = ComputeImpl(*context, input); + else if (input.IsDataTypeString()) + status = ComputeImpl(*context, input); + else + ORT_THROW("Split operator does not support ", input.DataType(), " yet"); + + return status; +} + +template +inline void copy_data(const T* src, T* dst, size_t count) { + memcpy(dst, src, count * sizeof(T)); +} + +template <> +inline void copy_data(const std::string* src, std::string* dst, size_t count) { + const std::string* end = src + count; + std::copy(src, end, dst); +} + +template +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(1); + ORT_ENFORCE(split_tensor->Shape().NumDimensions() == 1, "An split tensor must be a vector tensor."); + auto nDims = static_cast(split_tensor->Shape()[0]); + const auto* data = split_tensor->template Data(); + std::vector 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 output_dimensions{input_dims}; + + int64_t input_offset = 0; + const T* input_data = input.template Data(); + + for (int i = 0; i < num_outputs; ++i) { + // update size of dimension for axis we're splitting on + auto split_size = gsl::narrow(split_sizes[i]); + output_dimensions[axis] = split_size; + + Tensor* output = context.Output(i, TensorShape{output_dimensions}); + T* output_data = output->template MutableData(); + + ::onnxruntime::math::CopyMatrix( + before_dims, // M + split_size * after_dims_excluding_split, // N + static_cast(input_data + input_offset), // A + after_dims_including_split_axis, // lda + static_cast(output_data), // B + split_size * after_dims_excluding_split, // ldb + [](const T* src, T* dst, size_t count) { + copy_data(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 diff --git a/orttraining/orttraining/training_ops/cpu/tensor/split.h b/orttraining/orttraining/training_ops/cpu/tensor/split.h new file mode 100644 index 0000000000..7abbb5fae3 --- /dev/null +++ b/orttraining/orttraining/training_ops/cpu/tensor/split.h @@ -0,0 +1,30 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include + +#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 + 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& split_sizes); + +} // namespace contrib +} // namespace onnxruntime diff --git a/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc b/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc index 9444b30776..cbbb7dd307 100644 --- a/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc +++ b/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc @@ -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, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, // Adam BuildKernelCreateInfo, diff --git a/orttraining/orttraining/training_ops/cuda/tensor/split.cc b/orttraining/orttraining/training_ops/cuda/tensor/split.cc new file mode 100644 index 0000000000..4a30b785a6 --- /dev/null +++ b/orttraining/orttraining/training_ops/cuda/tensor/split.cc @@ -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(1) + .TypeConstraint("T", DataTypeImpl::AllFixedSizeTensorTypes()), + SplitTraining); + +Status SplitTraining::ComputeInternal(OpKernelContext* ctx) const { + const Tensor* input_tensor = ctx->Input(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(1); + ORT_ENFORCE(split_tensor->Shape().NumDimensions() == 1, "An split tensor must be a vector tensor."); + auto nDims = static_cast(split_tensor->Shape()[0]); + const auto* data = split_tensor->template Data(); + std::vector 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 output_dimensions{input_dims}; + + CudaAsyncBuffer output_ptr(this, num_outputs); + gsl::span output_ptr_span = output_ptr.CpuSpan(); + std::vector 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(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 split_sizes_gpu(this, split_sizes); + split_sizes_gpu.CopyToGpu(); + + std::vector 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 split_sizes_range_gpu(this, split_sizes_range); + split_sizes_range_gpu.CopyToGpu(); + + CudaAsyncBuffer 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 diff --git a/orttraining/orttraining/training_ops/cuda/tensor/split.h b/orttraining/orttraining/training_ops/cuda/tensor/split.h new file mode 100644 index 0000000000..e8e8e746b0 --- /dev/null +++ b/orttraining/orttraining/training_ops/cuda/tensor/split.h @@ -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