From a71dab691d1d67469654e8e35043aa8bceea8b9e Mon Sep 17 00:00:00 2001 From: mindest <30493312+mindest@users.noreply.github.com> Date: Wed, 28 Jul 2021 16:04:49 +0800 Subject: [PATCH] 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 --- .../core/providers/cpu/nn/batch_norm_helper.h | 14 +- .../core/framework/gradient_graph_builder.cc | 1 - .../core/graph/training_op_defs.cc | 58 +++-- .../core/optimizer/batchnorm_replacement.cc | 14 +- .../cuda/batch_norm_internal_test.cc | 202 ++++++++++++++++++ .../cuda/cuda_training_kernels.cc | 24 ++- .../training_ops/cuda/nn/batch_norm_grad.cc | 137 ++++++++---- .../training_ops/cuda/nn/batch_norm_grad.h | 2 +- .../cuda/nn/batch_norm_internal.cc | 165 ++++++++++++++ .../cuda/nn/batch_norm_internal.h | 50 +++++ tools/ci_build/amd_hipify.py | 2 + 11 files changed, 584 insertions(+), 85 deletions(-) create mode 100644 orttraining/orttraining/test/training_ops/cuda/batch_norm_internal_test.cc create mode 100644 orttraining/orttraining/training_ops/cuda/nn/batch_norm_internal.cc create mode 100644 orttraining/orttraining/training_ops/cuda/nn/batch_norm_internal.h diff --git a/onnxruntime/core/providers/cpu/nn/batch_norm_helper.h b/onnxruntime/core/providers/cpu/nn/batch_norm_helper.h index 92efe04472..c813773b7c 100644 --- a/onnxruntime/core/providers/cpu/nn/batch_norm_helper.h +++ b/onnxruntime/core/providers/cpu/nn/batch_norm_helper.h @@ -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(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(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& 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}; } }; diff --git a/orttraining/orttraining/core/framework/gradient_graph_builder.cc b/orttraining/orttraining/core/framework/gradient_graph_builder.cc index d54c35f807..9ac1b36f95 100644 --- a/orttraining/orttraining/core/framework/gradient_graph_builder.cc +++ b/orttraining/orttraining/core/framework/gradient_graph_builder.cc @@ -33,7 +33,6 @@ GradientGraphBuilder::GradientGraphBuilder(Graph* graph, auto rule_based_graph_transformer = std::make_unique("pre_training_rule_based_graph_transformer"); rule_based_graph_transformer->Register(std::make_unique()); - rule_based_graph_transformer->Register(std::make_unique()); graph_transformation_mgr_.Register(std::move(rule_based_graph_transformer), TransformerLevel::Level2); diff --git a/orttraining/orttraining/core/graph/training_op_defs.cc b/orttraining/orttraining/core/graph/training_op_defs.cc index df8e42190c..5c32e1b992 100644 --- a/orttraining/orttraining/core/graph/training_op_defs.cc +++ b/orttraining/orttraining/core/graph/training_op_defs.cc @@ -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(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); diff --git a/orttraining/orttraining/core/optimizer/batchnorm_replacement.cc b/orttraining/orttraining/core/optimizer/batchnorm_replacement.cc index 28e6cb93cd..ec68e83980 100644 --- a/orttraining/orttraining/core/optimizer/batchnorm_replacement.cc +++ b/orttraining/orttraining/core/optimizer/batchnorm_replacement.cc @@ -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(1)); // Assign provider to this new node. Provider should be same as the provider for old node. batchnorm_internal_node.SetExecutionProviderType(bn_node.GetExecutionProviderType()); diff --git a/orttraining/orttraining/test/training_ops/cuda/batch_norm_internal_test.cc b/orttraining/orttraining/test/training_ops/cuda/batch_norm_internal_test.cc new file mode 100644 index 0000000000..0e329537d7 --- /dev/null +++ b/orttraining/orttraining/test/training_ops/cuda/batch_norm_internal_test.cc @@ -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& 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 channel_dims{2}; + + std::vector 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 scale = {1.0f, 1.0f}; + std::vector B = {0.0f, 0.0f}; + std::vector mean = {1.0f, 2.0f}; + std::vector 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 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 running_mean = {-0.1754f, 0.303106f}; + std::vector running_var = {0.7812f, 1.5865f}; + std::vector saved_mean = {-0.306f, 0.114562f}; + std::vector saved_inv_std = {1.2288f, 0.861317f}; + + if (test_double) { + std::vector X_double (X.begin(), X.end()); + std::vector scale_double (scale.begin(), scale.end()); + std::vector B_double (B.begin(), B.end()); + std::vector mean_double (mean.begin(), mean.end()); + std::vector var_double (var.begin(), var.end()); + + std::vector Y_double (Y.begin(), Y.end()); + std::vector running_mean_double (running_mean.begin(), running_mean.end()); + std::vector running_var_double (running_var.begin(), running_var.end()); + std::vector saved_mean_double (saved_mean.begin(), saved_mean.end()); + std::vector saved_inv_std_double (saved_inv_std.begin(), saved_inv_std.end()); + + test.AddInput("X", input_output_dims, X_double); + test.AddInput("scale", channel_dims, scale_double); + test.AddInput("B", channel_dims, B_double); + test.AddInput("mean", channel_dims, mean_double); + test.AddInput("var", channel_dims, var_double); + + test.AddOutput("Y", input_output_dims, Y_double); + test.AddOutput("running_mean", channel_dims, running_mean_double); + test.AddOutput("running_var", channel_dims, running_var_double); + test.AddOutput("saved_mean", channel_dims, saved_mean_double); + test.AddOutput("saved_inv_std", channel_dims, saved_inv_std_double); + } else { + if (T_is_half) { + std::vector X_half(X.size()); + ConvertFloatToMLFloat16(X.data(), X_half.data(), int(X.size())); + test.AddInput("X", input_output_dims, X_half); + + std::vector Y_half(Y.size()); + ConvertFloatToMLFloat16(Y.data(), Y_half.data(), int(Y.size())); + test.AddOutput("Y", input_output_dims, Y_half); + } else { + test.AddInput("X", input_output_dims, X); + test.AddOutput("Y", input_output_dims, Y); + } + + if (T1_is_half) { + std::vector scale_half(scale.size()); + ConvertFloatToMLFloat16(scale.data(), scale_half.data(), int(scale.size())); + test.AddInput("scale", channel_dims, scale_half); + + std::vector B_half(B.size()); + ConvertFloatToMLFloat16(B.data(), B_half.data(), int(B.size())); + test.AddInput("B", channel_dims, B_half); + } else { + test.AddInput("scale", channel_dims, scale); + test.AddInput("B", channel_dims, B); + } + + if (T2_is_half) { + std::vector mean_half(mean.size()); + ConvertFloatToMLFloat16(mean.data(), mean_half.data(), int(mean.size())); + test.AddInput("mean", channel_dims, mean_half); + + std::vector var_half(var.size()); + ConvertFloatToMLFloat16(var.data(), var_half.data(), int(var.size())); + test.AddInput("var", channel_dims, var_half); + + std::vector running_mean_half(running_mean.size()); + ConvertFloatToMLFloat16(running_mean.data(), running_mean_half.data(), int(running_mean.size())); + test.AddOutput("running_mean", channel_dims, running_mean_half); + + std::vector running_var_half(running_var.size()); + ConvertFloatToMLFloat16(running_var.data(), running_var_half.data(), int(running_var.size())); + test.AddOutput("running_var", channel_dims, running_var_half); + + std::vector saved_mean_half(saved_mean.size()); + ConvertFloatToMLFloat16(saved_mean.data(), saved_mean_half.data(), int(saved_mean.size())); + test.AddOutput("saved_mean", channel_dims, saved_mean_half); + + std::vector 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("saved_inv_std", channel_dims, saved_inv_std_half); + } else { + test.AddInput("mean", channel_dims, mean); + test.AddInput("var", channel_dims, var); + test.AddOutput("running_mean", channel_dims, running_mean); + test.AddOutput("running_var", channel_dims, running_var); + test.AddOutput("saved_mean", channel_dims, saved_mean); + test.AddOutput("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 input_output_dims{16}; + std::vector channel_dims{1}; + + 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}); + test.AddInput("B", channel_dims, {0.0f}); + test.AddInput("mean", channel_dims, {1.0f}); + test.AddInput("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("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("running_mean", channel_dims, {0.0139f}); + test.AddOutput("running_var", channel_dims, {1.1074f}); + test.AddOutput("saved_mean", channel_dims, {-0.0957f}); + test.AddOutput("saved_inv_std", channel_dims, {0.9762f}); + + test.Run(OpTester::ExpectResult::kExpectSuccess, "", + {kCpuExecutionProvider, kTensorrtExecutionProvider, kOpenVINOExecutionProvider}); +} +#endif // USE_CUDA + +} // namespace test +} // 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 dd5c075686..6ad9d65ee5 100644 --- a/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc +++ b/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc @@ -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, BuildKernelCreateInfo, BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -340,6 +351,11 @@ Status RegisterCudaTrainingKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, #if defined(CUDA_VERSION) && CUDA_VERSION >= 11000 // Adam diff --git a/orttraining/orttraining/training_ops/cuda/nn/batch_norm_grad.cc b/orttraining/orttraining/training_ops/cuda/nn/batch_norm_grad.cc index 4f59eba9ee..422c1ee1f6 100644 --- a/orttraining/orttraining/training_ops/cuda/nn/batch_norm_grad.cc +++ b/orttraining/orttraining/training_ops/cuda/nn/batch_norm_grad.cc @@ -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()), \ - BatchNormalizationGrad); +#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()) \ + .TypeConstraint("T1", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("T2", DataTypeImpl::GetTensorType()), \ + BatchNormalizationGrad); -template -Status BatchNormalizationGrad::ComputeInternal(OpKernelContext* ctx) const { +template +Status BatchNormalizationGrad::ComputeInternal(OpKernelContext* ctx) const { typedef typename ToCudaType::MappedType CudaT; + typedef typename ToCudaType::MappedType CudaT1; + typedef typename ToCudaType::MappedType CudaT2; 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_variance = ctx->Input(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(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(dY->template Data()); auto X_data = reinterpret_cast(X->template Data()); - auto Scale_data = reinterpret_cast(Scale->template Data()); - auto saved_mean_data = reinterpret_cast(saved_mean->template Data()); - auto saved_variance_data = reinterpret_cast(saved_variance->template Data()); + auto Scale_data = reinterpret_cast(Scale->template Data()); + auto saved_mean_data = reinterpret_cast(saved_mean->template Data()); + auto saved_inv_std_data = reinterpret_cast(saved_inv_std->template Data()); auto dX_data = reinterpret_cast(ctx->Output(0, input_shape)->template MutableData()); - auto dScale_data = reinterpret_cast(ctx->Output(1, channel_shape)->template MutableData()); - auto dBias_data = reinterpret_cast(ctx->Output(2, channel_shape)->template MutableData()); + auto dScale_data = reinterpret_cast(ctx->Output(1, channel_shape)->template MutableData()); + auto dBias_data = reinterpret_cast(ctx->Output(2, channel_shape)->template MutableData()); const auto alpha = Consts::One; const auto beta = Consts::Zero; @@ -52,39 +59,79 @@ Status BatchNormalizationGrad::ComputeInternal(OpKernelContext* ctx) const { vector new_dims; BatchNormHelper::NormalizeDims(input_shape, new_dims); ORT_RETURN_IF_ERROR(input_tensor.Set(new_dims, CudnnTensor::GetDataType())); + // 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(Scale_data); + auto p_saved_mean = reinterpret_cast(saved_mean_data); + auto p_saved_inv_std = reinterpret_cast(saved_inv_std_data); + auto p_dScale = reinterpret_cast(dScale_data); + auto p_dBias = reinterpret_cast(dBias_data); + + IAllocatorUniquePtr p_f_scale, p_f_dScale, p_f_dBias, p_f_saved_mean, p_f_saved_inv_std; + + if (std::is_same::value) { + p_f_scale = GetScratchBuffer(C); + p_f_dScale = GetScratchBuffer(C); + p_f_dBias = GetScratchBuffer(C); + + Impl_Cast(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::value) { + p_f_saved_mean = GetScratchBuffer(C); + p_f_saved_inv_std = GetScratchBuffer(C); + + Impl_Cast(Stream(), saved_mean_data, p_f_saved_mean.get(), C); + Impl_Cast(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::value) { + Impl_Cast(Stream(), reinterpret_cast(p_dScale), dScale_data, C); + Impl_Cast(Stream(), reinterpret_cast(p_dBias), dBias_data, C); + } + return Status::OK(); } -#define SPECIALIZED_GRADIENT(T) \ - REGISTER_GRADIENT_KERNEL_TYPED(T) \ - template Status BatchNormalizationGrad::ComputeInternal(OpKernelContext* ctx) const; +#define SPECIALIZED_GRADIENT(T, T1, T2) \ + REGISTER_GRADIENT_KERNEL_TYPED(T, T1, T2) \ + template Status BatchNormalizationGrad::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 diff --git a/orttraining/orttraining/training_ops/cuda/nn/batch_norm_grad.h b/orttraining/orttraining/training_ops/cuda/nn/batch_norm_grad.h index ff29e18dbd..c462829f2b 100644 --- a/orttraining/orttraining/training_ops/cuda/nn/batch_norm_grad.h +++ b/orttraining/orttraining/training_ops/cuda/nn/batch_norm_grad.h @@ -11,7 +11,7 @@ namespace onnxruntime { namespace cuda { -template +template class BatchNormalizationGrad final : public CudaKernel { public: BatchNormalizationGrad(const OpKernelInfo& info) diff --git a/orttraining/orttraining/training_ops/cuda/nn/batch_norm_internal.cc b/orttraining/orttraining/training_ops/cuda/nn/batch_norm_internal.cc new file mode 100644 index 0000000000..14015f54d1 --- /dev/null +++ b/orttraining/orttraining/training_ops/cuda/nn/batch_norm_internal.cc @@ -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()) \ + .TypeConstraint("T1", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("T2", DataTypeImpl::GetTensorType()), \ + BatchNormInternal); + +template +Status BatchNormInternal::ComputeInternal(OpKernelContext* p_op_kernel_context) const { + typedef typename ToCudaType::MappedType CudaT; + typedef typename ToCudaType::MappedType CudaT1; + typedef typename ToCudaType::MappedType CudaT2; + + const Tensor* X = p_op_kernel_context->Input(0); + const Tensor* scale = p_op_kernel_context->Input(1); + const Tensor* B = p_op_kernel_context->Input(2); + const Tensor* mean = p_op_kernel_context->Input(3); + const Tensor* var = p_op_kernel_context->Input(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(X->template Data()); + auto scale_data = reinterpret_cast(scale->template Data()); + auto b_data = reinterpret_cast(B->template Data()); + auto mean_data = reinterpret_cast(mean->template Data()); + auto var_data = reinterpret_cast(var->template Data()); + + auto y_data = reinterpret_cast(Y->template MutableData()); + + const auto alpha = Consts::One; + const auto beta = Consts::Zero; + + CudnnTensor data_desc, bn_tensor_desc; + vector new_dims; + BatchNormHelper::NormalizeDims(x_shape, new_dims); + ORT_RETURN_IF_ERROR(data_desc.Set(new_dims, CudnnTensor::GetDataType())); + // 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(running_mean->template MutableData()); + auto running_var_data = reinterpret_cast(running_var->template MutableData()); + auto saved_mean_data = reinterpret_cast(saved_mean->template MutableData()); + auto saved_inv_std_data = reinterpret_cast(saved_inv_std->template MutableData()); + + auto p_scale = reinterpret_cast(scale_data); + auto p_B = reinterpret_cast(b_data); + auto p_running_mean = reinterpret_cast(running_mean_data); + auto p_running_var = reinterpret_cast(running_var_data); + auto p_saved_mean = reinterpret_cast(saved_mean_data); + auto p_saved_inv_std = reinterpret_cast(saved_inv_std_data); + + + const int64_t C = new_dims[1]; + IAllocatorUniquePtr 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::value) { + // Convert scale/B to float + p_f_scale = GetScratchBuffer(C); + p_f_B = GetScratchBuffer(C); + + Impl_Cast(Stream(), scale_data, p_f_scale.get(), C); + Impl_Cast(Stream(), b_data, p_f_B.get(), C); + + p_scale = p_f_scale.get(); + p_B = p_f_B.get(); + } + + if (std::is_same::value) { + // Convert mean/var to float + p_f_running_mean = GetScratchBuffer(C); + p_f_running_var = GetScratchBuffer(C); + p_f_saved_mean = GetScratchBuffer(C); + p_f_saved_inv_std = GetScratchBuffer(C); + + Impl_Cast(Stream(), mean_data, p_f_running_mean.get(), C); + Impl_Cast(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::value) { + Impl_Cast(Stream(), reinterpret_cast(p_running_mean), running_mean_data, C); + Impl_Cast(Stream(), reinterpret_cast(p_running_var), running_var_data, C); + Impl_Cast(Stream(), reinterpret_cast(p_saved_mean), saved_mean_data, C); + Impl_Cast(Stream(), reinterpret_cast(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::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 diff --git a/orttraining/orttraining/training_ops/cuda/nn/batch_norm_internal.h b/orttraining/orttraining/training_ops/cuda/nn/batch_norm_internal.h new file mode 100644 index 0000000000..3f46c91f22 --- /dev/null +++ b/orttraining/orttraining/training_ops/cuda/nn/batch_norm_internal.h @@ -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 +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("epsilon", &tmp_epsilon).IsOK()); + epsilon_ = ClampCudnnBatchNormEpsilon(static_cast(tmp_epsilon)); + + // spatial or not + int64_t tmp_spatial; + if (op_kernel_info.GetAttr("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("momentum", &tmp_momentum).IsOK()) { + momentum_ = static_cast(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 diff --git a/tools/ci_build/amd_hipify.py b/tools/ci_build/amd_hipify.py index 865a7092b2..cc105a5143 100644 --- a/tools/ci_build/amd_hipify.py +++ b/tools/ci_build/amd_hipify.py @@ -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',