mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-25 19:48:11 +00:00
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:
parent
539d1d44c1
commit
a71dab691d
11 changed files with 584 additions and 85 deletions
|
|
@ -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};
|
||||
}
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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',
|
||||
|
|
|
|||
Loading…
Reference in a new issue