mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
Implement BatchNormGradient kernel for CPU EP (#7622)
**Description**: Register an implementation for BatchNormInternal and add a CPU kernel for BatchNormGradient. This is the third in a series of PRs to implement BN training on CPU (first was #6946, second was #7539). **Motivation and Context** Support training networks with BatchNorm (e.g. convnets). Also note that there exists a CUDA kernel for BN (forward training & backwards) but it's currently disabled due to flaky failures; someone more familiar with those parts can register the implementation for BNInternal on CUDA (gradient kernel doesn't have to change). --------- Co-authored-by: Simon Zirui Guo <simonguozirui@berkeley.edu> Co-authored-by: mindest <linminuser@gmail.com> Co-authored-by: mindest <30493312+mindest@users.noreply.github.com>
This commit is contained in:
parent
5e2f46df2b
commit
3c5d02a9ce
7 changed files with 186 additions and 12 deletions
|
|
@ -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<float, float, float> 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"); }
|
||||
|
||||
|
|
|
|||
|
|
@ -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<int64_t> input_output_dims{2, 2, 2, 2};
|
||||
std::vector<int64_t> channel_dims{2};
|
||||
test.AddInput<float>("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<float>("scale", channel_dims, {1.0f, 1.0f});
|
||||
test.AddInput<float>("B", channel_dims, {0.0f, 0.0f});
|
||||
test.AddInput<float>("mean", channel_dims, {1.0f, 2.0f});
|
||||
test.AddInput<float>("var", channel_dims, {1.0f, 2.0f});
|
||||
|
||||
test.AddOutput<float>("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<float>("running_mean", channel_dims, {-0.1754f, 0.303106f});
|
||||
test.AddOutput<float>("running_var", channel_dims, {0.696052f, 1.41316f});
|
||||
test.AddOutput<float>("saved_mean", channel_dims, {-0.306f, 0.114562f});
|
||||
test.AddOutput<float>("saved_inv_std", channel_dims, {1.2288f, 0.861317f});
|
||||
|
||||
std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
|
||||
execution_providers.emplace_back(DefaultCpuExecutionProvider());
|
||||
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
|
||||
}
|
||||
|
||||
} // namespace test
|
||||
} // namespace contrib
|
||||
} // namespace onnxruntime
|
||||
|
|
@ -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<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, BatchNormInternal)>,
|
||||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SoftmaxCrossEntropy)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SoftmaxCrossEntropyGrad)>,
|
||||
|
|
@ -172,6 +175,7 @@ Status RegisterCpuTrainingKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, GatherGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, GatherElementsGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, GeluGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, BatchNormalizationGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SigmoidGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, TanhGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, QuickGeluGrad)>,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 <typename T>
|
||||
Status BatchNormalizationGrad<T>::Compute(OpKernelContext* ctx) const {
|
||||
const Tensor* dY = ctx->Input<Tensor>(0);
|
||||
const Tensor* X = ctx->Input<Tensor>(1);
|
||||
const Tensor* scale = ctx->Input<Tensor>(2);
|
||||
const Tensor* saved_mean = ctx->Input<Tensor>(3);
|
||||
const Tensor* saved_inv_std = ctx->Input<Tensor>(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<T>();
|
||||
const auto* X_data = X->template Data<T>();
|
||||
const auto* scale_data = scale->template Data<T>();
|
||||
const auto* saved_mean_data = saved_mean->template Data<T>();
|
||||
const auto* saved_inv_std_data = saved_inv_std->template Data<T>();
|
||||
|
||||
auto* dX_data = ctx->Output(0, X_shape)->template MutableData<T>();
|
||||
auto* dScale_data = ctx->Output(1, channel_shape)->template MutableData<T>();
|
||||
auto* dBias_data = ctx->Output(2, channel_shape)->template MutableData<T>();
|
||||
|
||||
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<T> scale_arr(scale_data, scale_tensor_size);
|
||||
ConstEigenVectorArrayMap<T> mean_arr(saved_mean_data, scale_tensor_size);
|
||||
ConstEigenVectorArrayMap<T> inv_std_arr(saved_inv_std_data, scale_tensor_size);
|
||||
|
||||
EigenVectorArrayMap<T> dBias_arr(dBias_data, scale_tensor_size);
|
||||
EigenVectorArrayMap<T> 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<T> X_arr(X_data, sample_size, N * C);
|
||||
ConstEigenArrayMap<T> dY_arr(dY_data, sample_size, N * C);
|
||||
EigenArrayMap<T> 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<float>())
|
||||
.TypeConstraint("T1", DataTypeImpl::GetTensorType<float>())
|
||||
.TypeConstraint("T2", DataTypeImpl::GetTensorType<float>()),
|
||||
BatchNormalizationGrad<float>);
|
||||
|
||||
} // namespace contrib
|
||||
} // namespace onnxruntime
|
||||
|
|
@ -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 <typename T>
|
||||
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
|
||||
|
|
@ -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<float>())
|
||||
.TypeConstraint("T1", DataTypeImpl::GetTensorType<float>())
|
||||
.TypeConstraint("T2", DataTypeImpl::GetTensorType<float>()),
|
||||
BatchNorm<float>);
|
||||
|
||||
} // namespace contrib
|
||||
} // namespace onnxruntime
|
||||
Loading…
Reference in a new issue