diff --git a/orttraining/orttraining/test/training_ops/cpu/reduction/reduction_ops_test.cc b/orttraining/orttraining/test/training_ops/cpu/reduction/reduction_ops_test.cc index 834ff013e1..5cca525229 100644 --- a/orttraining/orttraining/test/training_ops/cpu/reduction/reduction_ops_test.cc +++ b/orttraining/orttraining/test/training_ops/cpu/reduction/reduction_ops_test.cc @@ -84,6 +84,7 @@ TEST(AllOpTest, All_1d_large) { } } } +#endif class ReductionOpTest : public ::testing::TestWithParam { protected: @@ -102,6 +103,7 @@ TEST_P(ReductionOpTest, ReduceAllL2) { test.Run(); } +#ifdef USE_CUDA TEST_P(ReductionOpTest, ReduceAllL2HalfHalf) { OpTester test("ReduceAllL2", 1, onnxruntime::kMSDomain, true); test.SetDeterminism(GetParam()); @@ -163,6 +165,7 @@ TEST_P(ReductionOpTest, ReduceAllL2HalfFloat) { test.AddOutput("reduced", {}, result); test.Run(); } +#endif void TestMultiTensorReduce( const int tensor_count, @@ -226,8 +229,6 @@ TEST_P(ReductionOpTest, ReduceAllL2Many) { // invoke with and without use_determinism flag for session INSTANTIATE_TEST_SUITE_P(ReductionOpTestWrapper, ReductionOpTest, ::testing::Bool()); -#endif - TEST(ReductionOpTest, ReduceSumTraining_int32) { OpTester test("ReduceSumTraining", 1, onnxruntime::kMSDomain); test.AddAttribute("keepdims", (int64_t)1); diff --git a/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc b/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc index 239db66cf4..0fce6f5a6d 100644 --- a/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc +++ b/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc @@ -107,6 +107,8 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, double_int64_t, Scale); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, double_int32_t, Scale); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float_float, ReduceAllL2); + #ifdef USE_MPI // Pipeline communication operators. class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, Send); @@ -217,6 +219,8 @@ Status RegisterCpuTrainingKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + #ifdef USE_MPI BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/orttraining/orttraining/training_ops/cpu/reduction/reduction_all.cc b/orttraining/orttraining/training_ops/cpu/reduction/reduction_all.cc new file mode 100644 index 0000000000..951b66114b --- /dev/null +++ b/orttraining/orttraining/training_ops/cpu/reduction/reduction_all.cc @@ -0,0 +1,50 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "orttraining/training_ops/cpu/reduction/reduction_all.h" +#include "core/providers/cpu/reduction/reduction_ops.h" + +namespace onnxruntime { +namespace contrib { + +#define REGISTER_REDUCEALLL2_KERNEL_TYPED(TIn, TOut) \ + ONNX_OPERATOR_TYPED_KERNEL_EX(ReduceAllL2, kMSDomain, 1, TIn##_##TOut, kCpuExecutionProvider, \ + KernelDefBuilder() \ + .TypeConstraint("TIn", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("TOut", DataTypeImpl::GetTensorType()), \ + ReduceAllL2); + +REGISTER_REDUCEALLL2_KERNEL_TYPED(float, float) + +template +Status ReduceAllL2::Compute(OpKernelContext* ctx) const { + // Get Input tensor count. + const auto total_tensor_count = ctx->InputCount(); + std::vector tensor_pointers(total_tensor_count); + std::vector tensor_sizes(total_tensor_count); + + for (int i = 0; i < total_tensor_count; ++i) { + const Tensor* input = ctx->Input(i); + const auto size = input->Shape().Size(); + ORT_ENFORCE(size <= std::numeric_limits::max(), "Number of reduced elements (", size, + ") exceeds the max allowed value (", std::numeric_limits::max(), ")."); + tensor_pointers[i] = input->template Data(); + tensor_sizes[i] = size; + } + + // Allocate output tensor. + Tensor* output = ctx->Output(0, {}); + TOut* output_data = output->template MutableData(); + *output_data = TOut(0.f); + // perform reduction l2norm = sqrt[sum(tensor[i][j]**2)] for i,j over all tensor elements + for (int i = 0; i < total_tensor_count; ++i) { + *output_data += + ReduceAggregatorSumSquare(tensor_sizes[i], tensor_pointers[i][0]).aggall(tensor_pointers[i]); + } + + *output_data = reduce_sqrt(*output_data); + return Status::OK(); +} + +} // namespace contrib +} // namespace onnxruntime diff --git a/orttraining/orttraining/training_ops/cpu/reduction/reduction_all.h b/orttraining/orttraining/training_ops/cpu/reduction/reduction_all.h new file mode 100644 index 0000000000..cb0d726436 --- /dev/null +++ b/orttraining/orttraining/training_ops/cpu/reduction/reduction_all.h @@ -0,0 +1,20 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/common/common.h" +#include "core/framework/op_kernel.h" + +namespace onnxruntime { +namespace contrib { + +template +class ReduceAllL2 final : public OpKernel { + public: + ReduceAllL2(const OpKernelInfo& info) : OpKernel(info) {} + Status Compute(OpKernelContext* context) const override; +}; + +} // namespace contrib +} // namespace onnxruntime