diff --git a/orttraining/orttraining/test/gradient/gradient_ops_test.cc b/orttraining/orttraining/test/gradient/gradient_ops_test.cc index 0dc062e095..79d1e045f0 100644 --- a/orttraining/orttraining/test/gradient/gradient_ops_test.cc +++ b/orttraining/orttraining/test/gradient/gradient_ops_test.cc @@ -1385,14 +1385,11 @@ TEST(GradientCheckerTest, UnsqueezeGrad) { // TODO: Reshape missing -#ifdef USE_CUDA -// TODO fix flaky test -// failing random seed: 4133818171 -TEST(GradientCheckerTest, DISABLED_BatchNormalizationGrad) { +TEST(GradientCheckerTest, BatchNormalizationGrad) { float max_error; GradientChecker gradient_checker; - OpDef op_def{"BatchNormalization"}; - float error_tolerance = 1e-2f; + OpDef op_def{"BatchNormInternal", kMSDomain, 1}; + float error_tolerance = 2e-2f; float epsilon = 1e-05f; float momentum = 0.1f; @@ -1499,7 +1496,7 @@ TEST(GradientCheckerTest, DISABLED_BatchNormalizationGrad) { // case for larger multi-dimensional X { int channel_dim = 5; - TensorShape in_out_shape({6, channel_dim, 1, 3, 2, 4}); + TensorShape in_out_shape({6, channel_dim, 3, 2, 4}); TensorShape channel_shape({channel_dim}); // inputs TensorInfo x_info{in_out_shape, true}; @@ -1545,7 +1542,6 @@ TEST(GradientCheckerTest, DISABLED_BatchNormalizationGrad) { } */ } -#endif TEST(GradientCheckerTest, SigmoidGrad) { UnaryOpGradientTest("Sigmoid"); } diff --git a/orttraining/orttraining/test/training_ops/cpu/nn/batchnorm_internal_test.cc b/orttraining/orttraining/test/training_ops/cpu/nn/batchnorm_internal_test.cc new file mode 100644 index 0000000000..e9795a2468 --- /dev/null +++ b/orttraining/orttraining/test/training_ops/cpu/nn/batchnorm_internal_test.cc @@ -0,0 +1,48 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/framework/tensor.h" +#include "core/session/inference_session.h" +#include "test/util/include/default_providers.h" +#include "test/providers/provider_test_utils.h" + +namespace onnxruntime { +namespace contrib { +namespace test { + +using namespace onnxruntime::test; + +TEST(BatchNormInternalTest, ForwardTrainingTest) { + OpTester test("BatchNormInternal", 1, kMSDomain); + float epsilon = 1e-05f; + float momentum = 0.1f; + test.AddAttribute("epsilon", epsilon); + test.AddAttribute("momentum", momentum); + std::vector input_output_dims{2, 2, 2, 2}; + std::vector channel_dims{2}; + test.AddInput("X", input_output_dims, + {-0.2953f, 0.1180f, 1.0973f, -0.1931f, -0.1999f, -0.0237f, 1.5181f, 0.0076f, + -1.0830f, -1.5433f, 0.4327f, -0.9813f, 0.7875f, -0.4080f, -2.3144f, 1.5493f}); + test.AddInput("scale", channel_dims, {1.0f, 1.0f}); + test.AddInput("B", channel_dims, {0.0f, 0.0f}); + test.AddInput("mean", channel_dims, {1.0f, 2.0f}); + test.AddInput("var", channel_dims, {1.0f, 2.0f}); + + test.AddOutput("Y", input_output_dims, + {0.0131f, 0.5210f, 1.7244f, 0.1387f, -0.2708f, -0.1191f, 1.2089f, -0.0922f, + -0.9548f, -1.5203f, 0.9077f, -0.8298f, 0.5796f, -0.4501f, -2.0921f, 1.2358f}); + + test.AddOutput("running_mean", channel_dims, {-0.1754f, 0.303106f}); + test.AddOutput("running_var", channel_dims, {0.696052f, 1.41316f}); + test.AddOutput("saved_mean", channel_dims, {-0.306f, 0.114562f}); + test.AddOutput("saved_inv_std", channel_dims, {1.2288f, 0.861317f}); + + std::vector> execution_providers; + execution_providers.emplace_back(DefaultCpuExecutionProvider()); + + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); +} + +} // namespace test +} // namespace contrib +} // 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 6974c8d31e..dda92fbb41 100644 --- a/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc +++ b/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc @@ -23,6 +23,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, 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, BatchNormInternal); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SoftmaxCrossEntropy); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SoftmaxCrossEntropyGrad); @@ -76,6 +77,7 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, FastG class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, BiasGeluGrad_dX); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, BiasFastGeluGrad_dX); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, GatherNDGrad); +class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, BatchNormalizationGrad); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float_float, Scale); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float_double, Scale); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float_int64_t, Scale); @@ -146,6 +148,7 @@ Status RegisterCpuTrainingKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -172,6 +175,7 @@ Status RegisterCpuTrainingKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/orttraining/orttraining/training_ops/cpu/gist/gistencode_op.cc b/orttraining/orttraining/training_ops/cpu/gist/gistencode_op.cc index cb81fa4ca5..fc8872a825 100644 --- a/orttraining/orttraining/training_ops/cpu/gist/gistencode_op.cc +++ b/orttraining/orttraining/training_ops/cpu/gist/gistencode_op.cc @@ -7,10 +7,10 @@ namespace onnxruntime { namespace contrib { ONNX_OPERATOR_KERNEL_EX( GistBinarizeEncoder, - kMSDomain, + kMSDomain, 1, kCpuExecutionProvider, - KernelDefBuilder().Alias(0,0).TypeConstraint("T", DataTypeImpl::AllTensorTypes()), + KernelDefBuilder().Alias(0, 0).TypeConstraint("T", DataTypeImpl::AllTensorTypes()), GistBinarizeEncoderOp); Status GistBinarizeEncoderOp::Compute(OpKernelContext* context) const { @@ -30,5 +30,5 @@ Status GistBinarizeEncoderOp::Compute(OpKernelContext* context) const { ORT_ENFORCE(target != nullptr); return Status::OK(); } -} -} +} // namespace contrib +} // namespace onnxruntime diff --git a/orttraining/orttraining/training_ops/cpu/nn/batch_norm_grad.cc b/orttraining/orttraining/training_ops/cpu/nn/batch_norm_grad.cc new file mode 100644 index 0000000000..65c9b66a8e --- /dev/null +++ b/orttraining/orttraining/training_ops/cpu/nn/batch_norm_grad.cc @@ -0,0 +1,83 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "orttraining/training_ops/cpu/nn/batch_norm_grad.h" +#include "core/util/math_cpuonly.h" +#include "core/framework/op_kernel_context_internal.h" +#include "core/providers/cpu/nn/batch_norm_helper.h" + +namespace onnxruntime { +namespace contrib { + +template +Status BatchNormalizationGrad::Compute(OpKernelContext* ctx) const { + const Tensor* dY = ctx->Input(0); + const Tensor* X = ctx->Input(1); + const Tensor* scale = ctx->Input(2); + const Tensor* saved_mean = ctx->Input(3); + const Tensor* saved_inv_std = ctx->Input(4); + + const TensorShape X_shape = X->Shape(); + const TensorShape channel_shape = saved_mean->Shape(); + + // no B here, but B has same size as scale, so can validate inputs for gradient with this substitute + ORT_RETURN_IF_ERROR(BatchNormHelper::ValidateInputs(X, scale, scale, saved_mean, saved_inv_std)); + + const auto* dY_data = dY->template Data(); + const auto* X_data = X->template Data(); + const auto* scale_data = scale->template Data(); + const auto* saved_mean_data = saved_mean->template Data(); + const auto* saved_inv_std_data = saved_inv_std->template Data(); + + auto* dX_data = ctx->Output(0, X_shape)->template MutableData(); + auto* dScale_data = ctx->Output(1, channel_shape)->template MutableData(); + auto* dBias_data = ctx->Output(2, channel_shape)->template MutableData(); + + const auto& dims_vec = X_shape.GetDims(); + const size_t N = dims_vec[0]; + const size_t C = dims_vec[1]; // assume NCHW as per the spec + + // calculate sample_size (per individual channel) + size_t sample_size = X_shape.SizeFromDimension(2); + size_t scale_tensor_size = C; + + ConstEigenVectorArrayMap scale_arr(scale_data, scale_tensor_size); + ConstEigenVectorArrayMap mean_arr(saved_mean_data, scale_tensor_size); + ConstEigenVectorArrayMap inv_std_arr(saved_inv_std_data, scale_tensor_size); + + EigenVectorArrayMap dBias_arr(dBias_data, scale_tensor_size); + EigenVectorArrayMap dScale_arr(dScale_data, scale_tensor_size); + + dBias_arr.setZero(); + dScale_arr.setZero(); + + const auto scaled_inv_std = scale_arr * inv_std_arr / (N * sample_size); + + ConstEigenArrayMap X_arr(X_data, sample_size, N * C); + ConstEigenArrayMap dY_arr(dY_data, sample_size, N * C); + EigenArrayMap dX_arr(dX_data, sample_size, N * C); + + for (size_t nc = 0; nc < N * C; ++nc) { + size_t c = nc % C; + dBias_arr(c) += dY_arr.col(nc).sum(); + dScale_arr(c) += ((X_arr.col(nc) - mean_arr(c)) * inv_std_arr(c) * dY_arr.col(nc)).sum(); + } + for (size_t nc = 0; nc < N * C; ++nc) { + size_t c = nc % C; + dX_arr.col(nc) = scaled_inv_std(c) * (dY_arr.col(nc) * N * sample_size - dBias_arr(c) - + (X_arr.col(nc) - mean_arr(c)) * dScale_arr(c) * inv_std_arr(c)); + } + + return Status::OK(); +} + +ONNX_OPERATOR_KERNEL_EX( + BatchNormalizationGrad, kMSDomain, 1, kCpuExecutionProvider, + KernelDefBuilder() + .TypeConstraint("T", DataTypeImpl::GetTensorType()) + .TypeConstraint("T1", DataTypeImpl::GetTensorType()) + .TypeConstraint("T2", DataTypeImpl::GetTensorType()), + BatchNormalizationGrad); + +} // namespace contrib +} // namespace onnxruntime diff --git a/orttraining/orttraining/training_ops/cpu/nn/batch_norm_grad.h b/orttraining/orttraining/training_ops/cpu/nn/batch_norm_grad.h new file mode 100644 index 0000000000..59013838ea --- /dev/null +++ b/orttraining/orttraining/training_ops/cpu/nn/batch_norm_grad.h @@ -0,0 +1,23 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/framework/op_kernel.h" + +namespace onnxruntime { +namespace contrib { + +template +class BatchNormalizationGrad final : public OpKernel { + public: + explicit BatchNormalizationGrad(const OpKernelInfo& info) : OpKernel(info) {} + + Status Compute(OpKernelContext* context) const override; + + private: + ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(BatchNormalizationGrad); +}; + +} // namespace contrib +} // namespace onnxruntime diff --git a/orttraining/orttraining/training_ops/cpu/nn/batch_norm_internal.cc b/orttraining/orttraining/training_ops/cpu/nn/batch_norm_internal.cc new file mode 100644 index 0000000000..755e4853ac --- /dev/null +++ b/orttraining/orttraining/training_ops/cpu/nn/batch_norm_internal.cc @@ -0,0 +1,20 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/providers/cpu/nn/batch_norm.h" + +namespace onnxruntime { +namespace contrib { + +ONNX_OPERATOR_KERNEL_EX( + BatchNormInternal, kMSDomain, 1, kCpuExecutionProvider, + KernelDefBuilder() + .Alias(3, 1) + .Alias(4, 2) + .TypeConstraint("T", DataTypeImpl::GetTensorType()) + .TypeConstraint("T1", DataTypeImpl::GetTensorType()) + .TypeConstraint("T2", DataTypeImpl::GetTensorType()), + BatchNorm); + +} // namespace contrib +} // namespace onnxruntime