Implement BatchNormInternal for cuda (#8172)

* correct batchnorm replacement output order;

remove bn replacement in grad graph builder

* update op defs and kernel class

* implement batch norm internal and grad.

* change saved_var into saved_inv_std

* cuda test case: bn internal

* remove redundant include

* fix comment; add support and UT for 1d input.

* exclude batch_norm_internal in amd_hipify

* run BNInternal UT for CUDA only

* fix CI error

* fix comment errors

* fix error

* add comment for inconsistency with cudnnBN doc

* additional comments for cudnnBN inconsistency
This commit is contained in:
mindest 2021-07-28 16:04:49 +08:00 committed by GitHub
parent 539d1d44c1
commit a71dab691d
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
11 changed files with 584 additions and 85 deletions

View file

@ -19,13 +19,11 @@ class BatchNormHelper {
const Tensor* var,
bool is_spatial = true) {
const auto& x_dims = X->Shape().GetDims();
if (x_dims.size() < 2) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Invalid input X: The rank of input X must be atleast 2. Got rank: ", x_dims.size());
}
int64_t num_channels = x_dims[1];
int num_feature_dims = static_cast<int>(X->Shape().NumDimensions() - 2); // the first 2 are respectively - N and C
// If x_dims size < 2, num_channels defaults to 1.
int64_t num_channels = x_dims.size() > 1 ? x_dims[1] : 1;
// the first 2 are respectively - N and C.
int num_feature_dims = x_dims.size() > 1 ? static_cast<int>(x_dims.size() - 2) : 0;
// defined as per spec and used for validation
int kNumInputScaleDimensions = (is_spatial ? 1 : num_feature_dims + 1);
@ -109,6 +107,8 @@ class BatchNormHelper {
static void NormalizeDims(const TensorShape& x_shape, std::vector<int64_t>& new_dims) {
new_dims.clear();
auto& orig_dims = x_shape.GetDims();
ORT_ENFORCE(orig_dims.size() < 6,
"Input dim size should be < 6 for BatchNorm, but got ", std::to_string(orig_dims.size()));
if (orig_dims.size() == 4 /*supported size by CUDA*/ ||
orig_dims.size() == 5 /*supported size by CUDA*/) {
new_dims = orig_dims;
@ -118,8 +118,8 @@ class BatchNormHelper {
auto rank = x_shape.NumDimensions();
auto num_samples = rank > 0 ? orig_dims[0] : 1; // NCHW
auto num_channels = rank > 1 ? orig_dims[1] : 1;
auto width = rank > 3 ? orig_dims[3] : 1;
auto height = rank > 2 ? orig_dims[2] : 1;
int64_t width = 1;
new_dims = {num_samples, num_channels, height, width};
}
};

View file

@ -33,7 +33,6 @@ GradientGraphBuilder::GradientGraphBuilder(Graph* graph,
auto rule_based_graph_transformer =
std::make_unique<RuleBasedGraphTransformer>("pre_training_rule_based_graph_transformer");
rule_based_graph_transformer->Register(std::make_unique<InsertMaxPoolOutput>());
rule_based_graph_transformer->Register(std::make_unique<BatchNormReplacement>());
graph_transformation_mgr_.Register(std::move(rule_based_graph_transformer),
TransformerLevel::Level2);

View file

@ -1682,7 +1682,7 @@ Example 4:
})
.SetContextDependentFunctionBodyBuilder(
[](const FunctionBodyBuildContext& ctx, const OpSchema& schema, FunctionProto& functionProto) {
/* DropoutGrad (dy, mask, optional ratio, optional training_mode) => dX
/* DropoutGrad (dy, mask, optional ratio, optional training_mode) => dX
dX = Where (mask, dY / (1-ratio), 0)
where ratio = 0.5 if not specified.
@ -2048,7 +2048,7 @@ Example 4:
.TypeAndShapeInferenceFunction(ONNX_NAMESPACE::propagateShapeAndTypeFromFirstInput)
.SetContextDependentFunctionBodyBuilder(
[](const FunctionBodyBuildContext& ctx, const OpSchema& schema, FunctionProto& functionProto) {
/* Default GeluGrad computation:
/* Default GeluGrad computation:
dX = dY * [0.5f * [erf(sqrt(1/2)*X) + 1.0] + alpha*X*exp(-0.5f * X * X)]
which expands to the following ONNX graph:
*/
@ -2170,22 +2170,30 @@ Example 4:
ONNX_CONTRIB_OPERATOR_SCHEMA(BatchNormalizationGrad)
.SetDomain(kMSDomain)
.SinceVersion(1)
.SetDoc("BatchNormalization")
.SetDoc("BatchNormalizationGrad")
.Attr("epsilon",
"epsilon value",
AttributeProto::FLOAT)
.Input(0, "dY", "Gradient output from previous node", "T")
.Input(1, "X", "Input", "T")
.Input(2, "scale", "Scale tensor", "T")
.Input(3, "mean", "Mean of X", "T")
.Input(4, "variance", "Variance of X", "T")
.Input(2, "scale", "Scale tensor", "T1")
.Input(3, "mean", "Mean of X", "T2")
.Input(4, "variance", "Variance of X", "T2")
.Output(0, "X_grad", "Gradient of the input", "T")
.Output(1, "scale_grad", "Gradient of the scale", "T")
.Output(2, "bias_grad", "Gradient of the bias", "T")
.Output(1, "scale_grad", "Gradient of the scale", "T1")
.Output(2, "bias_grad", "Gradient of the bias", "T1")
.TypeConstraint(
"T",
{"tensor(float16)", "tensor(float)", "tensor(double)", "tensor(bfloat16)"},
"Constrain input and output types to float tensors.");
"Constrain input and output types to float tensors.")
.TypeConstraint(
"T1",
{"tensor(float16)", "tensor(float)", "tensor(double)", "tensor(bfloat16)"},
"Constrain scale and bias types to float tensors.")
.TypeConstraint(
"T2",
{"tensor(float16)", "tensor(float)", "tensor(double)", "tensor(bfloat16)"},
"Constrain mean and variance types to float tensors.");
ONNX_CONTRIB_OPERATOR_SCHEMA(Group)
.SetDomain(kMSDomain)
@ -2362,30 +2370,40 @@ Return true if all elements are true and false otherwise.
.Attr("momentum", "momentum value", AttributeProto::FLOAT, 0.9f)
.Attr("training_mode", "true if training", AttributeProto::INT, static_cast<int64_t>(1))
.Input(0, "X", "Input tensor.", "T", OpSchema::Single, true, 1, OpSchema::Differentiable)
.Input(1, "scale", "Scale tensor of shape (C).", "T", OpSchema::Single, true, 1, OpSchema::Differentiable)
.Input(2, "B", "Bias tensor of shape (C).", "T", OpSchema::Single, true, 1, OpSchema::Differentiable)
.Input(3, "input_mean", "running mean tensor of shape (C).", "U", OpSchema::Single, true, 1, OpSchema::Differentiable)
.Input(4, "input_var", "running variance tensor of shape (C).", "U", OpSchema::Single, true, 1, OpSchema::Differentiable)
.Input(1, "scale", "Scale tensor of shape (C).", "T1", OpSchema::Single, true, 1, OpSchema::Differentiable)
.Input(2, "B", "Bias tensor of shape (C).", "T1", OpSchema::Single, true, 1, OpSchema::Differentiable)
.Input(3, "input_mean", "running mean tensor of shape (C).", "T2", OpSchema::Single, true, 1, OpSchema::Differentiable)
.Input(4, "input_var", "running variance tensor of shape (C).", "T2", OpSchema::Single, true, 1, OpSchema::Differentiable)
.Output(0, "Y", "The output tensor of the same shape as X", "T", OpSchema::Single, true, 1, OpSchema::Differentiable)
.Output(1, "running_mean", "The running mean after BN.", "U", OpSchema::Optional, true, 1, OpSchema::NonDifferentiable)
.Output(2, "running_var", "Running var after BN", "U", OpSchema::Optional, true, 1, OpSchema::NonDifferentiable)
.Output(3, "saved_mean", "Mean of the batch", "U", OpSchema::Optional, true, 1, OpSchema::NonDifferentiable)
.Output(4, "saved_inv_std", "Inverse standard deviation for the batch", "U", OpSchema::Optional, true, 1, OpSchema::NonDifferentiable)
.Output(1, "running_mean", "The running mean after BN.", "T2", OpSchema::Optional, true, 1, OpSchema::NonDifferentiable)
.Output(2, "running_var", "Running var after BN", "T2", OpSchema::Optional, true, 1, OpSchema::NonDifferentiable)
.Output(3, "saved_mean", "Mean of the batch", "T2", OpSchema::Optional, true, 1, OpSchema::NonDifferentiable)
.Output(4, "saved_inv_std", "Inverse standard deviation for the batch", "T2", OpSchema::Optional, true, 1, OpSchema::NonDifferentiable)
.TypeConstraint(
"T",
{"tensor(float16)", "tensor(float)", "tensor(double)", "tensor(bfloat16)"},
"Constrain input and output types to float tensors.")
.TypeConstraint(
"U",
"T1",
{"tensor(float16)", "tensor(float)", "tensor(double)", "tensor(bfloat16)"},
"Constrain mean and variance types to float tensors. It allows all float type for U.")
"Constrain scale and bias types to float tensors.")
.TypeConstraint(
"T2",
{"tensor(float16)", "tensor(float)", "tensor(double)", "tensor(bfloat16)"},
"Constrain mean and variance types to float tensors.")
.TypeAndShapeInferenceFunction([](ONNX_NAMESPACE::InferenceContext& ctx) {
propagateShapeAndTypeFromFirstInput(ctx);
propagateShapeFromInputToOutput(ctx, 0, 0);
Dim num_channels;
unifyInputDim(ctx, 0, 1, num_channels);
// Add support for 1D input X, in which case num_channels should default to 1.
auto& input_shape = getInputShape(ctx, 0);
if (input_shape.dim_size() <= 1) {
num_channels.set_dim_value(1);
} else {
unifyInputDim(ctx, 0, 1, num_channels);
}
unifyInputDim(ctx, 1, 0, num_channels);
unifyInputDim(ctx, 2, 0, num_channels);
unifyInputDim(ctx, 3, 0, num_channels);

View file

@ -27,20 +27,20 @@ Status BatchNormReplacement::Apply(Graph& graph, Node& bn_node, RewriteRuleEffec
if (bn_outputs.size() == 3) {
NodeArg& saved_mean_def = graph.GetOrCreateNodeArg(graph.GenerateNodeArgName("saved_mean_def"), scale_input_def_type_proto);
NodeArg& saved_inv_std = graph.GetOrCreateNodeArg(graph.GenerateNodeArgName("saved_inv_std"), scale_input_def_type_proto);
bn_outputs.push_back(&saved_inv_std);
bn_outputs.push_back(&saved_mean_def);
bn_outputs.push_back(&saved_inv_std);
}
// check Batch Normalization node has 5 output node args for training mode
ORT_ENFORCE(bn_node.OutputDefs().size() == 5);
Node& batchnorm_internal_node = graph.AddNode(graph.GenerateNodeName(bn_node.Name() + "_BatchNormInternal"),
"BatchNormInternal",
"BatchNormalization with saved mean/inv_std_dev",
bn_inputs,
bn_outputs,
&bn_node.GetAttributes(),
kMSDomain);
"BatchNormInternal",
"BatchNormalization with saved mean/inv_std",
bn_inputs,
bn_outputs,
&bn_node.GetAttributes(),
kMSDomain);
batchnorm_internal_node.AddAttribute("training_mode", static_cast<int64_t>(1));
// Assign provider to this new node. Provider should be same as the provider for old node.
batchnorm_internal_node.SetExecutionProviderType(bn_node.GetExecutionProviderType());

View file

@ -0,0 +1,202 @@
// 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/providers/provider_test_utils.h"
#include "gtest/gtest.h"
#include "gmock/gmock.h"
using namespace std;
namespace onnxruntime {
namespace contrib {
namespace test {
using namespace onnxruntime::test;
#ifdef USE_CUDA
static void TestBatchNormInternal(bool test_double = false, bool T_is_half = false,
bool T1_is_half = false, bool T2_is_half = false,
const std::vector<int64_t>& input_output_dims = {2, 2, 2, 2}) {
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> channel_dims{2};
std::vector<float> X = {-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};
std::vector<float> scale = {1.0f, 1.0f};
std::vector<float> B = {0.0f, 0.0f};
std::vector<float> mean = {1.0f, 2.0f};
std::vector<float> var = {1.0f, 2.0f};
// cudnnBatchNorm uses biased `batch_var` to calculate `Y` and `saved_inv_std`, while
// uses unbiased `batch_var` to update `running_var`:
// running_var = (1 - momentum) * unbiased_batch_var + momentum * running_var.
// When using biased `batch_var`, the new `running_var` should be {0.696052f, 1.41316f}.
std::vector<float> Y = {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};
std::vector<float> running_mean = {-0.1754f, 0.303106f};
std::vector<float> running_var = {0.7812f, 1.5865f};
std::vector<float> saved_mean = {-0.306f, 0.114562f};
std::vector<float> saved_inv_std = {1.2288f, 0.861317f};
if (test_double) {
std::vector<double> X_double (X.begin(), X.end());
std::vector<double> scale_double (scale.begin(), scale.end());
std::vector<double> B_double (B.begin(), B.end());
std::vector<double> mean_double (mean.begin(), mean.end());
std::vector<double> var_double (var.begin(), var.end());
std::vector<double> Y_double (Y.begin(), Y.end());
std::vector<double> running_mean_double (running_mean.begin(), running_mean.end());
std::vector<double> running_var_double (running_var.begin(), running_var.end());
std::vector<double> saved_mean_double (saved_mean.begin(), saved_mean.end());
std::vector<double> saved_inv_std_double (saved_inv_std.begin(), saved_inv_std.end());
test.AddInput<double>("X", input_output_dims, X_double);
test.AddInput<double>("scale", channel_dims, scale_double);
test.AddInput<double>("B", channel_dims, B_double);
test.AddInput<double>("mean", channel_dims, mean_double);
test.AddInput<double>("var", channel_dims, var_double);
test.AddOutput<double>("Y", input_output_dims, Y_double);
test.AddOutput<double>("running_mean", channel_dims, running_mean_double);
test.AddOutput<double>("running_var", channel_dims, running_var_double);
test.AddOutput<double>("saved_mean", channel_dims, saved_mean_double);
test.AddOutput<double>("saved_inv_std", channel_dims, saved_inv_std_double);
} else {
if (T_is_half) {
std::vector<MLFloat16> X_half(X.size());
ConvertFloatToMLFloat16(X.data(), X_half.data(), int(X.size()));
test.AddInput<MLFloat16>("X", input_output_dims, X_half);
std::vector<MLFloat16> Y_half(Y.size());
ConvertFloatToMLFloat16(Y.data(), Y_half.data(), int(Y.size()));
test.AddOutput<MLFloat16>("Y", input_output_dims, Y_half);
} else {
test.AddInput<float>("X", input_output_dims, X);
test.AddOutput<float>("Y", input_output_dims, Y);
}
if (T1_is_half) {
std::vector<MLFloat16> scale_half(scale.size());
ConvertFloatToMLFloat16(scale.data(), scale_half.data(), int(scale.size()));
test.AddInput<MLFloat16>("scale", channel_dims, scale_half);
std::vector<MLFloat16> B_half(B.size());
ConvertFloatToMLFloat16(B.data(), B_half.data(), int(B.size()));
test.AddInput<MLFloat16>("B", channel_dims, B_half);
} else {
test.AddInput<float>("scale", channel_dims, scale);
test.AddInput<float>("B", channel_dims, B);
}
if (T2_is_half) {
std::vector<MLFloat16> mean_half(mean.size());
ConvertFloatToMLFloat16(mean.data(), mean_half.data(), int(mean.size()));
test.AddInput<MLFloat16>("mean", channel_dims, mean_half);
std::vector<MLFloat16> var_half(var.size());
ConvertFloatToMLFloat16(var.data(), var_half.data(), int(var.size()));
test.AddInput<MLFloat16>("var", channel_dims, var_half);
std::vector<MLFloat16> running_mean_half(running_mean.size());
ConvertFloatToMLFloat16(running_mean.data(), running_mean_half.data(), int(running_mean.size()));
test.AddOutput<MLFloat16>("running_mean", channel_dims, running_mean_half);
std::vector<MLFloat16> running_var_half(running_var.size());
ConvertFloatToMLFloat16(running_var.data(), running_var_half.data(), int(running_var.size()));
test.AddOutput<MLFloat16>("running_var", channel_dims, running_var_half);
std::vector<MLFloat16> saved_mean_half(saved_mean.size());
ConvertFloatToMLFloat16(saved_mean.data(), saved_mean_half.data(), int(saved_mean.size()));
test.AddOutput<MLFloat16>("saved_mean", channel_dims, saved_mean_half);
std::vector<MLFloat16> saved_inv_std_half(saved_inv_std.size());
ConvertFloatToMLFloat16(saved_inv_std.data(), saved_inv_std_half.data(), int(saved_inv_std.size()));
test.AddOutput<MLFloat16>("saved_inv_std", channel_dims, saved_inv_std_half);
} else {
test.AddInput<float>("mean", channel_dims, mean);
test.AddInput<float>("var", channel_dims, var);
test.AddOutput<float>("running_mean", channel_dims, running_mean);
test.AddOutput<float>("running_var", channel_dims, running_var);
test.AddOutput<float>("saved_mean", channel_dims, saved_mean);
test.AddOutput<float>("saved_inv_std", channel_dims, saved_inv_std);
}
}
test.Run(OpTester::ExpectResult::kExpectSuccess, "",
{kCpuExecutionProvider, kTensorrtExecutionProvider, kOpenVINOExecutionProvider});
}
TEST(CudaKernelTest, BNInternalBasic) { // float case
TestBatchNormInternal();
}
TEST(CudaKernelTest, BNInternalDouble) { // double case
TestBatchNormInternal(true);
}
TEST(CudaKernelTest, BNInternalHalf) { // half case
TestBatchNormInternal(false, true, true, true);
}
TEST(CudaKernelTest, BNInternalHalfHalfFloat) { // half X/Y & scale/B, float mean/var
TestBatchNormInternal(false, true, true);
}
TEST(CudaKernelTest, BNInternalHalfFloatFloat) { // half X/Y, float scale/B & mean/var
TestBatchNormInternal(false, true);
}
TEST(CudaKernelTest, BNInternal3DInput) { // float case, 3d input
TestBatchNormInternal(false, false, false, false, {2, 2, 4});
}
TEST(CudaKernelTest, BNInternal5DInput) { // float case, 5d input
TestBatchNormInternal(false, false, false, false, {2, 2, 2, 1, 2});
}
TEST(CudaKernelTest, BNInternal1DInput) { // float case, 1d input
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{16};
std::vector<int64_t> channel_dims{1};
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});
test.AddInput<float>("B", channel_dims, {0.0f});
test.AddInput<float>("mean", channel_dims, {1.0f});
test.AddInput<float>("var", channel_dims, {1.0f});
// cudnnBatchNorm uses biased `batch_var` to calculate `Y` and `saved_inv_std`, while
// uses unbiased `batch_var` to update `running_var`:
// running_var = (1 - momentum) * unbiased_batch_var + momentum * running_var.
// When using biased `batch_var`, the new `running_var` should be {1.0444f}.
test.AddOutput<float>("Y", input_output_dims,
{-0.1948f, 0.2086f, 1.1646f, -0.0951f, -0.1017f, 0.0703f, 1.5754f, 0.1009f,
-0.9638f, -1.4131f, 0.5158f, -0.8645f, 0.8622f, -0.3049f, -2.1659f, 1.6059f});
test.AddOutput<float>("running_mean", channel_dims, {0.0139f});
test.AddOutput<float>("running_var", channel_dims, {1.1074f});
test.AddOutput<float>("saved_mean", channel_dims, {-0.0957f});
test.AddOutput<float>("saved_inv_std", channel_dims, {0.9762f});
test.Run(OpTester::ExpectResult::kExpectSuccess, "",
{kCpuExecutionProvider, kTensorrtExecutionProvider, kOpenVINOExecutionProvider});
}
#endif // USE_CUDA
} // namespace test
} // namespace contrib
} // namespace onnxruntime

View file

@ -69,8 +69,11 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, LogSoftmaxGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, LogSoftmaxGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, LogSoftmaxGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, BatchNormalizationGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, BatchNormalizationGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float_float, BatchNormalizationGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double_double, BatchNormalizationGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_MLFloat16, BatchNormalizationGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_float, BatchNormalizationGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float_float, BatchNormalizationGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, ConvGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, ConvGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, ConvGrad);
@ -144,6 +147,11 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, GistPack16Decoder);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, GistPackMsfp15Encoder);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, GistPackMsfp15Decoder);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float_float, BatchNormInternal);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double_double, BatchNormInternal);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_MLFloat16, BatchNormInternal);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_float, BatchNormInternal);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float_float, BatchNormInternal);
#if defined(CUDA_VERSION) && CUDA_VERSION >= 11000
// Adam
@ -275,8 +283,11 @@ Status RegisterCudaTrainingKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, int64_t, SoftmaxCrossEntropyLossInternal)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, int64_t, SoftmaxCrossEntropyLossInternalGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, int64_t, SoftmaxCrossEntropyLossInternalGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, BatchNormalizationGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, BatchNormalizationGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float_float, BatchNormalizationGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double_double, BatchNormalizationGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_MLFloat16, BatchNormalizationGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_float, BatchNormalizationGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float_float, BatchNormalizationGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, ConvGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, ConvGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, ConvGrad)>,
@ -340,6 +351,11 @@ Status RegisterCudaTrainingKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, GistPack16Decoder)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, GistPackMsfp15Encoder)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, GistPackMsfp15Decoder)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float_float, BatchNormInternal)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double_double, BatchNormInternal)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_MLFloat16, BatchNormInternal)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_float, BatchNormInternal)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float_float, BatchNormInternal)>,
#if defined(CUDA_VERSION) && CUDA_VERSION >= 11000
// Adam

View file

@ -5,45 +5,52 @@
#include "core/providers/common.h"
#include "core/providers/cuda/cudnn_common.h"
#include "core/providers/cpu/nn/batch_norm_helper.h"
#include "core/providers/cuda/math/unary_elementwise_ops_impl.h"
using namespace std;
namespace onnxruntime {
namespace cuda {
#define REGISTER_GRADIENT_KERNEL_TYPED(T) \
ONNX_OPERATOR_TYPED_KERNEL_EX( \
BatchNormalizationGrad, \
kMSDomain, \
1, \
T, \
kCudaExecutionProvider, \
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
BatchNormalizationGrad<T>);
#define REGISTER_GRADIENT_KERNEL_TYPED(T, T1, T2) \
ONNX_OPERATOR_TYPED_KERNEL_EX( \
BatchNormalizationGrad, \
kMSDomain, \
1, \
T##_##T1##_##T2, \
kCudaExecutionProvider, \
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()) \
.TypeConstraint("T1", DataTypeImpl::GetTensorType<T1>()) \
.TypeConstraint("T2", DataTypeImpl::GetTensorType<T2>()), \
BatchNormalizationGrad<T, T1, T2>);
template <typename T>
Status BatchNormalizationGrad<T>::ComputeInternal(OpKernelContext* ctx) const {
template <typename T, typename T1, typename T2>
Status BatchNormalizationGrad<T, T1, T2>::ComputeInternal(OpKernelContext* ctx) const {
typedef typename ToCudaType<T>::MappedType CudaT;
typedef typename ToCudaType<T1>::MappedType CudaT1;
typedef typename ToCudaType<T2>::MappedType CudaT2;
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_variance = ctx->Input<Tensor>(4);
// cudnnBatchNormalizationBackward() claims to use `savedInvVariance`, but the value
// is actually equal to the batch inv_std, so we use name `saved_inv_std` here.
const Tensor* saved_inv_std = ctx->Input<Tensor>(4);
const TensorShape input_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_variance));
ORT_RETURN_IF_ERROR(BatchNormHelper::ValidateInputs(X, Scale, Scale, saved_mean, saved_inv_std));
auto dY_data = reinterpret_cast<const CudaT*>(dY->template Data<T>());
auto X_data = reinterpret_cast<const CudaT*>(X->template Data<T>());
auto Scale_data = reinterpret_cast<const CudaT*>(Scale->template Data<T>());
auto saved_mean_data = reinterpret_cast<const CudaT*>(saved_mean->template Data<T>());
auto saved_variance_data = reinterpret_cast<const CudaT*>(saved_variance->template Data<T>());
auto Scale_data = reinterpret_cast<const CudaT1*>(Scale->template Data<T1>());
auto saved_mean_data = reinterpret_cast<const CudaT2*>(saved_mean->template Data<T2>());
auto saved_inv_std_data = reinterpret_cast<const CudaT2*>(saved_inv_std->template Data<T2>());
auto dX_data = reinterpret_cast<CudaT*>(ctx->Output(0, input_shape)->template MutableData<T>());
auto dScale_data = reinterpret_cast<CudaT*>(ctx->Output(1, channel_shape)->template MutableData<T>());
auto dBias_data = reinterpret_cast<CudaT*>(ctx->Output(2, channel_shape)->template MutableData<T>());
auto dScale_data = reinterpret_cast<CudaT1*>(ctx->Output(1, channel_shape)->template MutableData<T1>());
auto dBias_data = reinterpret_cast<CudaT1*>(ctx->Output(2, channel_shape)->template MutableData<T1>());
const auto alpha = Consts<CudaT>::One;
const auto beta = Consts<CudaT>::Zero;
@ -52,39 +59,79 @@ Status BatchNormalizationGrad<T>::ComputeInternal(OpKernelContext* ctx) const {
vector<int64_t> new_dims;
BatchNormHelper::NormalizeDims(input_shape, new_dims);
ORT_RETURN_IF_ERROR(input_tensor.Set(new_dims, CudnnTensor::GetDataType<CudaT>()));
// for fp16 input, `scale_bias_tensor` will have a float type; otherwise it will be the same as input type.
ORT_RETURN_IF_ERROR(scale_bias_tensor.Set(input_tensor, cudnn_batch_norm_mode_));
// note this is only valid for cudnnBatchNormalizationForwardTraining, not ForwardInference
CUDNN_RETURN_IF_ERROR(
cudnnBatchNormalizationBackward(
CudnnHandle(),
cudnn_batch_norm_mode_,
&alpha,
&beta,
&alpha,
&beta,
input_tensor,
X_data,
input_tensor,
dY_data,
input_tensor,
dX_data,
scale_bias_tensor,
Scale_data,
dScale_data,
dBias_data,
epsilon_,
saved_mean_data,
saved_variance_data));
const int64_t C = new_dims[1];
auto p_scale = reinterpret_cast<const void*>(Scale_data);
auto p_saved_mean = reinterpret_cast<const void*>(saved_mean_data);
auto p_saved_inv_std = reinterpret_cast<const void*>(saved_inv_std_data);
auto p_dScale = reinterpret_cast<void*>(dScale_data);
auto p_dBias = reinterpret_cast<void*>(dBias_data);
IAllocatorUniquePtr<float> p_f_scale, p_f_dScale, p_f_dBias, p_f_saved_mean, p_f_saved_inv_std;
if (std::is_same<T1, MLFloat16>::value) {
p_f_scale = GetScratchBuffer<float>(C);
p_f_dScale = GetScratchBuffer<float>(C);
p_f_dBias = GetScratchBuffer<float>(C);
Impl_Cast<CudaT1, float>(Stream(), Scale_data, p_f_scale.get(), C);
p_scale = p_f_scale.get();
p_dScale = p_f_dScale.get();
p_dBias = p_f_dBias.get();
}
if (std::is_same<T2, MLFloat16>::value) {
p_f_saved_mean = GetScratchBuffer<float>(C);
p_f_saved_inv_std = GetScratchBuffer<float>(C);
Impl_Cast<CudaT2, float>(Stream(), saved_mean_data, p_f_saved_mean.get(), C);
Impl_Cast<CudaT2, float>(Stream(), saved_inv_std_data, p_f_saved_inv_std.get(), C);
p_saved_mean = p_f_saved_mean.get();
p_saved_inv_std = p_f_saved_inv_std.get();
}
CUDNN_RETURN_IF_ERROR(cudnnBatchNormalizationBackward(
CudnnHandle(),
cudnn_batch_norm_mode_,
&alpha,
&beta,
&alpha,
&beta,
input_tensor,
X_data,
input_tensor,
dY_data,
input_tensor,
dX_data,
scale_bias_tensor,
p_scale,
p_dScale,
p_dBias,
epsilon_,
p_saved_mean,
p_saved_inv_std));
if (std::is_same<T1, MLFloat16>::value) {
Impl_Cast<float, CudaT1>(Stream(), reinterpret_cast<float*>(p_dScale), dScale_data, C);
Impl_Cast<float, CudaT1>(Stream(), reinterpret_cast<float*>(p_dBias), dBias_data, C);
}
return Status::OK();
}
#define SPECIALIZED_GRADIENT(T) \
REGISTER_GRADIENT_KERNEL_TYPED(T) \
template Status BatchNormalizationGrad<T>::ComputeInternal(OpKernelContext* ctx) const;
#define SPECIALIZED_GRADIENT(T, T1, T2) \
REGISTER_GRADIENT_KERNEL_TYPED(T, T1, T2) \
template Status BatchNormalizationGrad<T, T1, T2>::ComputeInternal(OpKernelContext* ctx) const;
SPECIALIZED_GRADIENT(float)
SPECIALIZED_GRADIENT(double)
SPECIALIZED_GRADIENT(float, float, float)
SPECIALIZED_GRADIENT(double, double, double)
SPECIALIZED_GRADIENT(MLFloat16, MLFloat16, MLFloat16)
SPECIALIZED_GRADIENT(MLFloat16, MLFloat16, float)
SPECIALIZED_GRADIENT(MLFloat16, float, float)
} // namespace cuda
} // namespace onnxruntime

View file

@ -11,7 +11,7 @@
namespace onnxruntime {
namespace cuda {
template <typename T>
template <typename T, typename T1, typename T2>
class BatchNormalizationGrad final : public CudaKernel {
public:
BatchNormalizationGrad(const OpKernelInfo& info)

View file

@ -0,0 +1,165 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include "orttraining/training_ops/cuda/nn/batch_norm_internal.h"
#include "core/providers/common.h"
#include "core/providers/cuda/cudnn_common.h"
#include "core/providers/cpu/nn/batch_norm_helper.h"
#include "core/providers/cuda/math/unary_elementwise_ops_impl.h"
using namespace std;
namespace onnxruntime {
namespace cuda {
#define REGISTER_KERNEL_TYPED(T, T1, T2) \
ONNX_OPERATOR_TYPED_KERNEL_EX( \
BatchNormInternal, \
kMSDomain, \
1, \
T##_##T1##_##T2, \
kCudaExecutionProvider, \
(*KernelDefBuilder::Create()) \
.Alias(3, 1) \
.Alias(4, 2) \
.TypeConstraint("T", DataTypeImpl::GetTensorType<T>()) \
.TypeConstraint("T1", DataTypeImpl::GetTensorType<T1>()) \
.TypeConstraint("T2", DataTypeImpl::GetTensorType<T2>()), \
BatchNormInternal<T, T1, T2>);
template <typename T, typename T1, typename T2>
Status BatchNormInternal<T, T1, T2>::ComputeInternal(OpKernelContext* p_op_kernel_context) const {
typedef typename ToCudaType<T>::MappedType CudaT;
typedef typename ToCudaType<T1>::MappedType CudaT1;
typedef typename ToCudaType<T2>::MappedType CudaT2;
const Tensor* X = p_op_kernel_context->Input<Tensor>(0);
const Tensor* scale = p_op_kernel_context->Input<Tensor>(1);
const Tensor* B = p_op_kernel_context->Input<Tensor>(2);
const Tensor* mean = p_op_kernel_context->Input<Tensor>(3);
const Tensor* var = p_op_kernel_context->Input<Tensor>(4);
ORT_RETURN_IF_ERROR(BatchNormHelper::ValidateInputs(X, scale, B, mean, var, spatial_ == 1));
const TensorShape& x_shape = X->Shape();
const TensorShape& channel_shape = mean->Shape();
Tensor* Y = p_op_kernel_context->Output(0, x_shape);
Tensor* running_mean = p_op_kernel_context->Output(1, channel_shape);
Tensor* running_var = p_op_kernel_context->Output(2, channel_shape);
Tensor* saved_mean = p_op_kernel_context->Output(3, channel_shape);
// cudnnBatchNormalizationForwardTraining() claims to output `resultSaveInvVariance`, but the value
// is actually equal to the batch inv_std, so we use name `saved_inv_std` here.
Tensor* saved_inv_std = p_op_kernel_context->Output(4, channel_shape);
auto x_data = reinterpret_cast<const CudaT*>(X->template Data<T>());
auto scale_data = reinterpret_cast<const CudaT1*>(scale->template Data<T1>());
auto b_data = reinterpret_cast<const CudaT1*>(B->template Data<T1>());
auto mean_data = reinterpret_cast<const CudaT2*>(mean->template Data<T2>());
auto var_data = reinterpret_cast<const CudaT2*>(var->template Data<T2>());
auto y_data = reinterpret_cast<CudaT*>(Y->template MutableData<T>());
const auto alpha = Consts<CudaT>::One;
const auto beta = Consts<CudaT>::Zero;
CudnnTensor data_desc, bn_tensor_desc;
vector<int64_t> new_dims;
BatchNormHelper::NormalizeDims(x_shape, new_dims);
ORT_RETURN_IF_ERROR(data_desc.Set(new_dims, CudnnTensor::GetDataType<CudaT>()));
// for fp16 input, `bn_tensor_desc` will have a float type; otherwise it will be the same as input type.
ORT_RETURN_IF_ERROR(bn_tensor_desc.Set(data_desc, cudnn_batch_norm_mode_));
auto running_mean_data = reinterpret_cast<CudaT2*>(running_mean->template MutableData<T2>());
auto running_var_data = reinterpret_cast<CudaT2*>(running_var->template MutableData<T2>());
auto saved_mean_data = reinterpret_cast<CudaT2*>(saved_mean->template MutableData<T2>());
auto saved_inv_std_data = reinterpret_cast<CudaT2*>(saved_inv_std->template MutableData<T2>());
auto p_scale = reinterpret_cast<const void*>(scale_data);
auto p_B = reinterpret_cast<const void*>(b_data);
auto p_running_mean = reinterpret_cast<void*>(running_mean_data);
auto p_running_var = reinterpret_cast<void*>(running_var_data);
auto p_saved_mean = reinterpret_cast<void*>(saved_mean_data);
auto p_saved_inv_std = reinterpret_cast<void*>(saved_inv_std_data);
const int64_t C = new_dims[1];
IAllocatorUniquePtr<float> p_f_scale, p_f_B, p_f_running_mean, p_f_running_var, p_f_saved_mean, p_f_saved_inv_std;
if (std::is_same<T1, MLFloat16>::value) {
// Convert scale/B to float
p_f_scale = GetScratchBuffer<float>(C);
p_f_B = GetScratchBuffer<float>(C);
Impl_Cast<CudaT1, float>(Stream(), scale_data, p_f_scale.get(), C);
Impl_Cast<CudaT1, float>(Stream(), b_data, p_f_B.get(), C);
p_scale = p_f_scale.get();
p_B = p_f_B.get();
}
if (std::is_same<T2, MLFloat16>::value) {
// Convert mean/var to float
p_f_running_mean = GetScratchBuffer<float>(C);
p_f_running_var = GetScratchBuffer<float>(C);
p_f_saved_mean = GetScratchBuffer<float>(C);
p_f_saved_inv_std = GetScratchBuffer<float>(C);
Impl_Cast<CudaT2, float>(Stream(), mean_data, p_f_running_mean.get(), C);
Impl_Cast<CudaT2, float>(Stream(), var_data, p_f_running_var.get(), C);
p_running_mean = p_f_running_mean.get();
p_running_var = p_f_running_var.get();
p_saved_mean = p_f_saved_mean.get();
p_saved_inv_std = p_f_saved_inv_std.get();
} else if (mean_data != running_mean_data) {
CUDA_RETURN_IF_ERROR(
cudaMemcpyAsync(running_mean_data, mean_data, C * sizeof(T2), cudaMemcpyDeviceToDevice, Stream()));
CUDA_RETURN_IF_ERROR(
cudaMemcpyAsync(running_var_data, var_data, C * sizeof(T2), cudaMemcpyDeviceToDevice, Stream()));
}
// NOTE: in cudnnBatchNorm, biased std/var is used when calculating `save_inv_std` and `y`, while
// `running_var` is updated using unbiased `batch_var`:
// running_var = (1 - momentum_) * unbiased_batch_var + momentum_ * running_var
// This is inconsistent with BatchNormalization Onnx spec, which uses population variance (biased).
CUDNN_RETURN_IF_ERROR(cudnnBatchNormalizationForwardTraining(
CudnnHandle(),
cudnn_batch_norm_mode_,
&alpha,
&beta,
data_desc,
x_data,
data_desc,
y_data,
bn_tensor_desc,
p_scale,
p_B,
1.0 - momentum_,
p_running_mean,
p_running_var,
epsilon_,
p_saved_mean,
p_saved_inv_std));
if (std::is_same<T2, MLFloat16>::value) {
Impl_Cast<float, CudaT2>(Stream(), reinterpret_cast<float*>(p_running_mean), running_mean_data, C);
Impl_Cast<float, CudaT2>(Stream(), reinterpret_cast<float*>(p_running_var), running_var_data, C);
Impl_Cast<float, CudaT2>(Stream(), reinterpret_cast<float*>(p_saved_mean), saved_mean_data, C);
Impl_Cast<float, CudaT2>(Stream(), reinterpret_cast<float*>(p_saved_inv_std), saved_inv_std_data, C);
}
return Status::OK();
}
#define SPECIALIZED_COMPUTE(T, T1, T2) \
REGISTER_KERNEL_TYPED(T, T1, T2) \
template Status BatchNormInternal<T, T1, T2>::ComputeInternal(OpKernelContext* ctx) const;
SPECIALIZED_COMPUTE(float, float, float)
SPECIALIZED_COMPUTE(double, double, double)
SPECIALIZED_COMPUTE(MLFloat16, MLFloat16, MLFloat16)
SPECIALIZED_COMPUTE(MLFloat16, MLFloat16, float)
SPECIALIZED_COMPUTE(MLFloat16, float, float)
} // namespace cuda
} // namespace onnxruntime

View file

@ -0,0 +1,50 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#pragma once
#include "gsl/gsl"
#include "core/providers/cuda/cuda_kernel.h"
#include "core/providers/cuda/cudnn_common.h"
namespace onnxruntime {
namespace cuda {
template <typename T, typename T1, typename T2>
class BatchNormInternal final : public CudaKernel {
public:
BatchNormInternal(const OpKernelInfo& op_kernel_info)
: CudaKernel{op_kernel_info},
cudnn_batch_norm_mode_(CUDNN_BATCHNORM_SPATIAL),
momentum_(0.9) {
float tmp_epsilon;
ORT_ENFORCE(op_kernel_info.GetAttr<float>("epsilon", &tmp_epsilon).IsOK());
epsilon_ = ClampCudnnBatchNormEpsilon(static_cast<double>(tmp_epsilon));
// spatial or not
int64_t tmp_spatial;
if (op_kernel_info.GetAttr<int64_t>("spatial", &tmp_spatial).IsOK()) {
spatial_ = tmp_spatial;
}
if (spatial_ == 0) {
cudnn_batch_norm_mode_ = CUDNN_BATCHNORM_PER_ACTIVATION;
}
float tmp_momentum;
if (op_kernel_info.GetAttr<float>("momentum", &tmp_momentum).IsOK()) {
momentum_ = static_cast<double>(tmp_momentum);
}
}
Status ComputeInternal(OpKernelContext* context) const override;
private:
double epsilon_;
int64_t spatial_ = 1; // default as per spec
cudnnBatchNormMode_t cudnn_batch_norm_mode_;
double momentum_;
};
} // namespace cuda
} // namespace onnxruntime

View file

@ -199,6 +199,8 @@ training_ops_excluded_files = [
'math/softmax_grad.cc',
'nn/batch_norm_grad.cc',
'nn/batch_norm_grad.h',
'nn/batch_norm_internal.cc',
'nn/batch_norm_internal.h',
'nn/conv_grad.cc',
'nn/conv_grad.h',
'reduction/reduction_all.cc',