BiasDropoutFusion (#4167)

* Implement BiasDropout Fusion and Kernel

Dropout kernel for residual input

BiasDropout Fusion to take residual input

Fix BiasDropout Kernel

Optimize DropoutGrad with 4 elements per thread

* Add graph transformer UT

* MLTypeCallDispatcher for RatioData

* Use MLTypeDispatcher for ratio tensor

* Handle traing_mode input for BiasDropout fusion

* Add test case for missing ratio input

* Replace using FinalizeNodeFusion

* Make BiasDropout kernel template-less

* Make DropoutGrad template-less

* Make Dropout and TrainableDropout template-less

* Regenerate onnx file for UT

* Minior fix on divmod in BiasDropoutKernel

* Adjust pt frontend test due to dropout randomnesss

* Make dropout kernel opeartion in fp32

Co-authored-by: Sherlock Huang <bahuang@OrtTrainingDev3.af05slrtruoetgaxwwjv5nsq5e.px.internal.cloudapp.net>
This commit is contained in:
Sherlock 2020-06-30 15:43:14 -07:00 committed by GitHub
parent 0404763f23
commit 6365760906
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
23 changed files with 1026 additions and 316 deletions

View file

@ -813,15 +813,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, int64_t, GatherND);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, MLFloat16_MLFloat16, Dropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, MLFloat16_float, Dropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, MLFloat16_double, Dropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, float_MLFloat16, Dropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, float_float, Dropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, float_double, Dropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, double_MLFloat16, Dropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, double_float, Dropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, double_double, Dropout);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, Dropout);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, Einsum);
static Status RegisterCudaKernels(KernelRegistry& kernel_registry) {
@ -1341,15 +1333,7 @@ static Status RegisterCudaKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, int64_t, GatherND)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, MLFloat16_MLFloat16, Dropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, MLFloat16_float, Dropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, MLFloat16_double, Dropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, float_MLFloat16, Dropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, float_float, Dropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, float_double, Dropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, double_MLFloat16, Dropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, double_float, Dropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, double_double, Dropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, Dropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, Einsum)>,
};

View file

@ -6,30 +6,18 @@
namespace onnxruntime {
namespace cuda {
#define REGISTER_KERNEL_TYPED(T1, T2) \
ONNX_OPERATOR_TYPED_KERNEL_EX( \
Dropout, \
kOnnxDomain, \
12, \
T1##_##T2, \
kCudaExecutionProvider, \
KernelDefBuilder() \
.TypeConstraint("T", DataTypeImpl::GetTensorType<T1>()) \
.TypeConstraint("T1", DataTypeImpl::GetTensorType<T2>()) \
.TypeConstraint("T2", DataTypeImpl::GetTensorType<bool>()) \
.InputMemoryType<OrtMemTypeCPUInput>(1) \
.InputMemoryType<OrtMemTypeCPUInput>(2), \
Dropout<T1, T2, false>);
REGISTER_KERNEL_TYPED(MLFloat16, MLFloat16)
REGISTER_KERNEL_TYPED(MLFloat16, float)
REGISTER_KERNEL_TYPED(MLFloat16, double)
REGISTER_KERNEL_TYPED(float, MLFloat16)
REGISTER_KERNEL_TYPED(float, float)
REGISTER_KERNEL_TYPED(float, double)
REGISTER_KERNEL_TYPED(double, MLFloat16)
REGISTER_KERNEL_TYPED(double, float)
REGISTER_KERNEL_TYPED(double, double)
ONNX_OPERATOR_KERNEL_EX(
Dropout,
kOnnxDomain,
12,
kCudaExecutionProvider,
KernelDefBuilder()
.TypeConstraint("T", DataTypeImpl::AllIEEEFloatTensorTypes())
.TypeConstraint("T1", DataTypeImpl::AllIEEEFloatTensorTypes())
.TypeConstraint("T2", DataTypeImpl::GetTensorType<bool>())
.InputMemoryType<OrtMemTypeCPUInput>(1)
.InputMemoryType<OrtMemTypeCPUInput>(2),
Dropout<false>);
} // namespace cuda
} // namespace onnxruntime

View file

@ -12,10 +12,36 @@
namespace onnxruntime {
namespace cuda {
template <typename T1, typename T2, bool trainable_dropout>
template <typename T>
struct GetRatioDataImpl {
void operator()(const Tensor* ratio, float& ratio_data) const {
ratio_data = static_cast<float>(*(ratio->template Data<T>()));
ORT_ENFORCE(ratio_data >= 0.0f && ratio_data < 1.0f, "ratio_data is outside range [0, 1)");
}
};
template <typename T>
struct DropoutComputeImpl {
void operator()(const cudaDeviceProp& prop,
const int64_t N,
const float ratio_data,
PhiloxGenerator& generator,
const Tensor& X,
Tensor& Y,
bool* mask_data) const {
typedef typename ToCudaType<T>::MappedType CudaT;
const CudaT* X_data = reinterpret_cast<const CudaT*>(X.template Data<T>());
CudaT* Y_data = reinterpret_cast<CudaT*>(Y.template MutableData<T>());
DropoutKernelImpl<CudaT>(prop, N, ratio_data, generator, X_data, Y_data, mask_data);
}
};
template <bool trainable_dropout>
class Dropout final : public CudaKernel {
public:
Dropout(const OpKernelInfo& info) : CudaKernel(info), default_ratio_(0.5) {
Dropout(const OpKernelInfo& info) : CudaKernel(info) {
int64_t seed = 0;
if (info.GetAttr<int64_t>("seed", &seed).IsOK()) {
generator_ = onnxruntime::make_unique<PhiloxGenerator>(static_cast<uint64_t>(seed));
@ -26,48 +52,40 @@ class Dropout final : public CudaKernel {
private:
mutable std::unique_ptr<PhiloxGenerator> generator_;
const float default_ratio_;
static constexpr float default_ratio_ = 0.5f;
};
template <typename T1, typename T2, bool trainable_dropout>
Status Dropout<T1, T2, trainable_dropout>::ComputeInternal(OpKernelContext* context) const {
typedef typename ToCudaType<T1>::MappedType CudaT;
template <bool trainable_dropout>
Status Dropout<trainable_dropout>::ComputeInternal(OpKernelContext* context) const {
//Get X_data
const Tensor* X = context->Input<Tensor>(0);
if (X == nullptr) return Status(common::ONNXRUNTIME, common::FAIL, "X Input is not available.");
const TensorShape& shape = X->Shape();
auto X_data = reinterpret_cast<const CudaT*>(X->template Data<T1>());
const int64_t N = shape.Size();
//Get Y_data
auto Y = context->Output(0, shape);
auto Y_data = reinterpret_cast<CudaT*>(Y->template MutableData<T1>());
//Get mask_data
auto mask = context->Output(1, shape);
ORT_ENFORCE(!mask || mask->Shape().Size() == N);
//Get the ratio_data
float ratio_data;
float ratio_data = default_ratio_;
auto ratio = context->Input<Tensor>(1);
static_assert(std::is_same<T2, MLFloat16>::value || std::is_same<T2, float>::value || std::is_same<T2, double>::value,
"T2 must be float16 or float or double");
if (ratio) {
ratio_data = static_cast<float>(*(ratio->template Data<T2>()));
} else {
ratio_data = default_ratio_;
utils::MLTypeCallDispatcher<GetRatioDataImpl, float, MLFloat16, double> t_disp(ratio->GetElementType());
t_disp.Invoke(ratio, ratio_data);
}
ORT_ENFORCE(ratio_data >= 0.0f && ratio_data < 1.0f);
const Tensor* training_mode = context->Input<Tensor>(2);
//Check for inference mode.
if ((0 == ratio_data /*Backward compat with TrainableDropout*/) ||
(!trainable_dropout && (training_mode == nullptr || *(training_mode->Data<bool>()) == false))) {
const void* X_data = X->DataRaw();
void* Y_data = Y->MutableDataRaw();
if (Y_data != X_data) {
CUDA_CALL_THROW(cudaMemcpyAsync(Y_data, X_data, N * sizeof(T1), cudaMemcpyDeviceToDevice));
CUDA_CALL_THROW(cudaMemcpyAsync(Y_data, X_data, X->SizeInBytes(), cudaMemcpyDeviceToDevice));
}
// If mask is requested, return all 1s.
@ -85,8 +103,10 @@ Status Dropout<T1, T2, trainable_dropout>::ComputeInternal(OpKernelContext* cont
return temp_mask_buffer.get();
}();
PhiloxGenerator& generator = generator_ != nullptr ? *generator_.get() : PhiloxGenerator::Default();
DropoutKernelImpl(GetDeviceProp(), N, ratio_data, generator, X_data, Y_data, mask_data);
PhiloxGenerator& generator = generator_ ? *generator_ : PhiloxGenerator::Default();
utils::MLTypeCallDispatcher<DropoutComputeImpl, float, MLFloat16, double> t_disp(X->GetElementType());
t_disp.Invoke(GetDeviceProp(), N, ratio_data, generator, *X, *Y, mask_data);
return Status::OK();
}

View file

@ -35,7 +35,7 @@ __global__ void DropoutKernel(
T* Y_data,
bool* mask_data) {
const float p = 1.0f - ratio;
const T scale = T(1.0f / p);
const float scale = 1.0f / p;
CUDA_LONG idx = blockDim.x * blockIdx.x + threadIdx.x;
CUDA_LONG step_size = gridDim.x * blockDim.x * UNROLL;
@ -52,12 +52,13 @@ __global__ void DropoutKernel(
// use of Philox_4x32_10 is to generate a multiple of 4 times number of threads.
for (CUDA_LONG id = idx; id < rounded_size; id += step_size) {
float4 rand = curand_uniform4(&state);
#pragma unroll
for (CUDA_LONG i = 0; i < UNROLL; i++) {
CUDA_LONG li = id + gridDim.x * blockDim.x * i;
if (li < N) {
mask_data[li] = (&rand.x)[i] < p;
Y_data[li] = X_data[li] * T(mask_data[li]) * scale;
Y_data[li] = T(float(X_data[li]) * mask_data[li] * scale);
}
}
@ -76,7 +77,7 @@ void DropoutKernelImpl(
bool* mask_data) {
const int block_size = 256;
const int blocks_per_sm = prop.maxThreadsPerMultiProcessor / block_size;
const int grid_size = std::min(prop.multiProcessorCount * blocks_per_sm, static_cast<int>(CeilDiv(N, block_size)));
const int grid_size = std::min(prop.multiProcessorCount * blocks_per_sm, static_cast<int>(CeilDiv(N, block_size * UNROLL)));
// Compute the number of random numbers generated by each thread, and increment philox generator offset by that amount.
const uint64_t counter_offset = static_cast<uint64_t>(((N - 1) / (block_size * grid_size * UNROLL) + 1) * UNROLL);

View file

@ -655,9 +655,7 @@ class TestOrtTrainer(unittest.TestCase):
assert np.array_equal(state_dict[key], loaded_state_dict[key])
def testBertTrainingBasic(self):
expected_losses = [
11.02906322479248, 11.094074249267578, 11.00899887084961, 11.06129264831543,
11.029067039489746, 11.040265083312988, 11.046793937683105, 10.993699073791504]
expected_losses = [11.034271, 11.125311, 11.006095, 11.046938, 11.027476, 11.015745, 11.060884, 10.971851]
expected_eval_loss = [10.95898914]
actual_losses, actual_eval_loss = runBertTrainingTest(
gradient_accumulation_steps=1, use_mixed_precision=False, allreduce_post_accumulation=False)
@ -669,14 +667,12 @@ class TestOrtTrainer(unittest.TestCase):
# print('eval_loss actual: ', actual_eval_loss)
# import pdb; pdb.set_trace()
rtol = 1e-04
rtol = 1e-03
assert_allclose(expected_losses, actual_losses, rtol=rtol, err_msg="loss mismatch")
assert_allclose(expected_eval_loss, actual_eval_loss, rtol=rtol, err_msg="evaluation loss mismatch")
def testBertTrainingGradientAccumulation(self):
expected_losses = [
11.02906322479248, 11.094074249267578, 11.008995056152344, 11.061283111572266,
11.029059410095215, 11.04024887084961, 11.04680347442627, 10.993708610534668]
expected_losses = [11.034271, 11.125311, 11.006093, 11.046929, 11.027471, 11.015731, 11.060894, 10.971855]
expected_eval_loss = [10.959011]
actual_losses, actual_eval_loss = runBertTrainingTest(
@ -689,7 +685,7 @@ class TestOrtTrainer(unittest.TestCase):
# print('eval_loss actual: ', actual_eval_loss)
# import pdb; pdb.set_trace()
rtol = 1e-04
rtol = 1e-03
assert_allclose(expected_losses, actual_losses, rtol=rtol, err_msg="loss mismatch")
assert_allclose(expected_eval_loss, actual_eval_loss, rtol=rtol, err_msg="evaluation loss mismatch")

View file

@ -0,0 +1,114 @@
import onnx
from onnx import helper
from onnx import TensorProto, OperatorSetIdProto
# inputs/outputs
A = helper.make_tensor_value_info('A', TensorProto.FLOAT, ['unk_1', 'unk_2', 3072])
B = helper.make_tensor_value_info('B', TensorProto.FLOAT, [3072])
R = helper.make_tensor_value_info('R', TensorProto.FLOAT, ['unk_1', 'unk_2', 3072])
C = helper.make_tensor_value_info('C', TensorProto.FLOAT, ['unk_1', 'unk_2', 3072])
mask = helper.make_tensor_value_info('mask', TensorProto.BOOL, ['unk_1', 'unk_2', 3072])
# initializers
ratio = helper.make_tensor('ratio_const', TensorProto.FLOAT, [], [0.8])
training_mode = helper.make_tensor('training_mode', TensorProto.BOOL, [], [1])
opsets = []
onnxdomain = OperatorSetIdProto()
onnxdomain.version = 12
onnxdomain.domain = "" # The empty string ("") or absence of this field implies the operator set that is defined as part of the ONNX specification.
opsets.append(onnxdomain)
kwargs={}
kwargs['opset_imports'] = opsets
# Create the model (ModelProto)
bias = helper.make_node("Add", ["A", "B"], ["add0_out"], "add0")
dropout_12 = helper.make_node("Dropout", ["add0_out", "ratio_const", "training_mode"], ["C", "mask"], "dropout0")
graph = helper.make_graph(
[bias, dropout_12],
"Bias_Dropout_Fusion", #name
[A, B],
[C],
[ratio, training_mode])
model = helper.make_model(graph, producer_name='onnx-example', **kwargs)
onnx.save(model, 'bias_dropout_fusion1.onnx')
# Create the model (ModelProto)
bias = helper.make_node("Add", ["B", "A"], ["add0_out"], "add0")
dropout_12 = helper.make_node("Dropout", ["add0_out", "ratio_const", "training_mode"], ["C", "mask"], "dropout0")
graph = helper.make_graph(
[bias, dropout_12],
"Bias_Dropout_Fusion", #name
[A, B],
[C],
[ratio, training_mode])
model = helper.make_model(graph, producer_name='onnx-example', **kwargs)
onnx.save(model, 'bias_dropout_fusion2.onnx')
# Create the model (ModelProto)
bias = helper.make_node("Add", ["A", "B"], ["add0_out"], "add0")
dropout_12 = helper.make_node("Dropout", ["add0_out", "ratio_const", "training_mode"], ["dropout_out", "mask"], "dropout0")
residual = helper.make_node("Add", ["dropout_out", "R"], ["C"], "add1")
graph = helper.make_graph(
[bias, dropout_12, residual],
"Bias_Dropout_Fusion", #name
[A, B, R],
[C],
[ratio, training_mode])
model = helper.make_model(graph, producer_name='onnx-example', **kwargs)
onnx.save(model, 'bias_dropout_residual_fusion1.onnx')
# Create the model (ModelProto)
bias = helper.make_node("Add", ["B", "A"], ["add0_out"], "add0")
dropout_12 = helper.make_node("Dropout", ["add0_out", "ratio_const", "training_mode"], ["dropout_out", "mask"], "dropout0")
residual = helper.make_node("Add", ["R", "dropout_out"], ["C"], "add1")
graph = helper.make_graph(
[bias, dropout_12, residual],
"Bias_Dropout_Fusion", #name
[A, B, R],
[C],
[ratio, training_mode])
model = helper.make_model(graph, producer_name='onnx-example', **kwargs)
onnx.save(model, 'bias_dropout_residual_fusion2.onnx')
# Create the model (ModelProto)
R_mismatch = helper.make_tensor_value_info('R', TensorProto.FLOAT, [3072])
bias = helper.make_node("Add", ["B", "A"], ["add0_out"], "add0")
dropout_12 = helper.make_node("Dropout", ["add0_out", "ratio_const", "training_mode"], ["dropout_out", "mask"], "dropout0")
residual = helper.make_node("Add", ["R", "dropout_out"], ["C"], "add1")
graph = helper.make_graph(
[bias, dropout_12, residual],
"Bias_Dropout_Fusion", #name
[A, B, R_mismatch],
[C],
[ratio, training_mode])
model = helper.make_model(graph, producer_name='onnx-example', **kwargs)
onnx.save(model, 'bias_dropout_residual_fusion_mismatch.onnx')
# Create the model (ModelProto)
bias = helper.make_node("Add", ["B", "A"], ["add0_out"], "add0")
trainable_dropout = helper.make_node("TrainableDropout", ["add0_out", "ratio_const"], ["dropout_out", "mask"], "dropout0")
residual = helper.make_node("Add", ["R", "dropout_out"], ["C"], "add1")
graph = helper.make_graph(
[bias, trainable_dropout, residual],
"Bias_Dropout_Fusion", #name
[A, B, R],
[C],
[ratio])
model = helper.make_model(graph, producer_name='onnx-example', **kwargs)
onnx.save(model, 'bias_trainabledropout_residual_fusion.onnx')

View file

@ -1010,6 +1010,54 @@ Example 4:
"Constrain indices to integer types")
.SetDoc(R"DOC(SoftmaxCrossEntropyLossGrad)DOC");
ONNX_CONTRIB_OPERATOR_SCHEMA(BiasDropout)
.SetDomain(kMSDomain)
.SinceVersion(1)
.SetDoc("BiasDropout")
.Attr("seed", "(Optional) Seed to the random generator, if not specified we will auto generate one.", AttributeProto::INT, OPTIONAL_VALUE)
.AllowUncheckedAttributes()
.Input(0, "data", "The input data as Tensor.", "T")
.Input(1, "bias", "The bias input, a vector with the same shape as last dim of data", "T")
.Input(2, "residual", "The residual input, must have the same shape as data", "T", OpSchema::Optional)
.Input(3, "ratio",
"The ratio of random dropout, with value in [0, 1). If this input was not set, "
"or if it was set to 0, the output would be a simple copy of the input. "
"If it's non-zero, output will be a random dropout of input, which is typically "
"the case during training.",
"T1",
OpSchema::Optional)
.Input(4, "training_mode",
"If set to true then it indicates dropout is being used for "
"training. It is an optional value hence unless specified explicitly, it is false. "
"If it is false, ratio is ignored and the operation mimics inference mode where nothing "
"will be dropped from the input data and if mask is requested as output it will contain "
"all ones.",
"T2",
OpSchema::Optional)
.Output(0, "output", "The output.", "T")
.Output(1, "mask", "The output mask of dropout.", "T2", OpSchema::Optional)
.TypeConstraint(
"T",
{"tensor(float16)", "tensor(float)", "tensor(double)"},
"Constrain input and output types to float tensors.")
.TypeConstraint(
"T1",
{"tensor(float16)", "tensor(float)", "tensor(double)"},
"Constrain input 'ratio' types to float tensors.")
.TypeConstraint(
"T2",
{"tensor(bool)"},
"Constrain output 'mask' types to boolean tensors.")
.TypeAndShapeInferenceFunction([](ONNX_NAMESPACE::InferenceContext& ctx) {
propagateShapeAndTypeFromFirstInput(ctx);
if (ctx.getNumOutputs() == 2) {
updateOutputElemType(ctx, 1, ONNX_NAMESPACE::TensorProto::BOOL);
if (hasNInputShapes(ctx, 1)) {
propagateShapeFromInputToOutput(ctx, 0, 1);
}
}
});
ONNX_CONTRIB_OPERATOR_SCHEMA(TrainableDropout)
.SetDomain(kOnnxDomain)
.SinceVersion(9)

View file

@ -0,0 +1,200 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include "core/optimizer/initializer.h"
#include "orttraining/core/optimizer/bias_dropout_fusion.h"
#include "core/graph/graph_utils.h"
#include <deque>
using namespace ONNX_NAMESPACE;
using namespace ::onnxruntime::common;
namespace onnxruntime {
void FuseResidualAddIfAny(Graph& graph, const Node& dropout_node,
std::vector<NodeArg*>& dropout_input,
std::vector<NodeArg*>& dropout_output,
std::vector<std::reference_wrapper<Node>>& nodes_to_fuse) {
bool has_residual_add = false;
for (auto last_node_itr = dropout_node.OutputNodesBegin(); last_node_itr != dropout_node.OutputNodesEnd(); ++last_node_itr) {
const Node& last_node = (*last_node_itr);
if (graph_utils::IsSupportedOptypeVersionAndDomain(last_node, "Add", {7}) &&
last_node.GetExecutionProviderType() == dropout_node.GetExecutionProviderType()) {
const TensorShapeProto* input1_shape = last_node.InputDefs()[0]->Shape();
const TensorShapeProto* input2_shape = last_node.InputDefs()[1]->Shape();
if (input1_shape == nullptr ||
input2_shape == nullptr ||
input1_shape->dim_size() < 1 ||
input2_shape->dim_size() < 1 ||
input1_shape->dim_size() != input2_shape->dim_size()) {
continue;
}
// Inputs of Residual Add must match in shape
bool match = true;
for (int i = 0; i < input1_shape->dim_size(); ++i) {
match &= ONNX_NAMESPACE::operator==(input1_shape->dim(i), input2_shape->dim(i));
}
if (!match) {
continue;
}
// dropout's output is not part of of graph output
if (!graph.GetNodeOutputsInGraphOutputs(dropout_node).empty()) {
continue;
}
Node& residual_add_node = *graph.GetNode(last_node.Index());
const std::string& dropout_output_name = dropout_node.OutputDefs()[0]->Name();
if (dropout_output_name == residual_add_node.InputDefs()[0]->Name()) {
dropout_input.push_back(residual_add_node.MutableInputDefs()[1]); // residual
} else if (dropout_output_name == residual_add_node.InputDefs()[1]->Name()) {
dropout_input.push_back(residual_add_node.MutableInputDefs()[0]); // residual
}
dropout_output[0] = residual_add_node.MutableOutputDefs()[0];
nodes_to_fuse.push_back(residual_add_node);
has_residual_add = true;
break;
}
}
if (!has_residual_add) {
NodeArg& dummy = graph.GetOrCreateNodeArg("", nullptr);
dropout_input.push_back(&dummy); // add a dummy residual
}
}
Status BiasDropoutFusion::ApplyImpl(Graph& graph, bool& modified, int graph_level, const logging::Logger& logger) const {
GraphViewer graph_viewer(graph);
const auto& node_topology_list = graph_viewer.GetNodesInTopologicalOrder();
for (auto node_index : node_topology_list) {
auto* node_ptr = graph.GetNode(node_index);
if (nullptr == node_ptr)
continue; // node was removed
auto& node = *node_ptr;
ORT_RETURN_IF_ERROR(Recurse(node, modified, graph_level, logger));
std::vector<std::reference_wrapper<Node>> nodes_to_fuse;
// matching for bias Add node
if (!graph_utils::IsSupportedOptypeVersionAndDomain(node, "Add", {7}) ||
!graph_utils::IsSupportedProvider(node, GetCompatibleExecutionProviders()) ||
node.GetOutputEdgesCount() != 1) {
continue;
}
std::vector<NodeArg*> dropout_input, dropout_output;
const TensorShapeProto* input1_shape = node.MutableInputDefs()[0]->Shape();
const TensorShapeProto* input2_shape = node.MutableInputDefs()[1]->Shape();
if (input1_shape == nullptr ||
input2_shape == nullptr ||
input1_shape->dim_size() < 1 ||
input2_shape->dim_size() < 1) {
continue;
}
int last_dim_shape1 = input1_shape->dim_size() - 1;
int last_dim_shape2 = input2_shape->dim_size() - 1;
if (!utils::HasDimValue(input1_shape->dim(last_dim_shape1)) ||
!utils::HasDimValue(input2_shape->dim(last_dim_shape2)) ||
input1_shape->dim(last_dim_shape1).dim_value() != input2_shape->dim(last_dim_shape2).dim_value()) {
continue;
}
if (input1_shape->dim_size() == 1) {
dropout_input.push_back(node.MutableInputDefs()[1]); // dropout input
dropout_input.push_back(node.MutableInputDefs()[0]); // bias
} else if (input2_shape->dim_size() == 1) {
dropout_input.push_back(node.MutableInputDefs()[0]); // dropout input
dropout_input.push_back(node.MutableInputDefs()[1]); // bias
} else {
continue;
}
Node& add_node = node;
nodes_to_fuse.push_back(add_node);
// matching for Dropout node
auto next_node_itr = node.OutputNodesBegin();
if (next_node_itr == node.OutputNodesEnd()) {
continue;
}
const Node& next_node = (*next_node_itr);
if (!(graph_utils::IsSupportedOptypeVersionAndDomain(next_node, "Dropout", {12}, kOnnxDomain) ||
graph_utils::IsSupportedOptypeVersionAndDomain(next_node, "TrainableDropout", {9}, kOnnxDomain)) ||
next_node.GetExecutionProviderType() != node.GetExecutionProviderType()) {
continue;
}
if (!graph.GetNodeOutputsInGraphOutputs(node).empty()) {
continue;
}
Node& dropout_node = *graph.GetNode(next_node.Index());
nodes_to_fuse.push_back(dropout_node);
dropout_output.push_back(dropout_node.MutableOutputDefs()[0]);
dropout_output.push_back(dropout_node.MutableOutputDefs()[1]);
FuseResidualAddIfAny(graph, dropout_node, dropout_input, dropout_output, nodes_to_fuse);
if (dropout_node.InputDefs().size() > 1) {
dropout_input.push_back(dropout_node.MutableInputDefs()[1]); // ratio
}
// populate training_mode
bool is_trainable_dropout = (dropout_node.OpType() == "TrainableDropout");
if (is_trainable_dropout) {
// Create training_mode initializer
ONNX_NAMESPACE::TensorProto training_mode_initializer;
training_mode_initializer.set_name(graph.GenerateNodeArgName("training_mode"));
training_mode_initializer.set_data_type(ONNX_NAMESPACE::TensorProto_DataType_BOOL);
const bool data = true;
training_mode_initializer.set_raw_data(&data, sizeof(bool));
NodeArg& training_mode_node_arg = graph_utils::AddInitializer(graph, training_mode_initializer);
dropout_input.push_back(&training_mode_node_arg);
} else {
if (dropout_node.InputDefs().size() > 2) {
dropout_input.push_back(dropout_node.MutableInputDefs()[2]);
}
}
const std::string op_type = "BiasDropout";
Node& dropout_add_fusion_node = graph.AddNode(graph.GenerateNodeName(op_type),
op_type,
"fused Add and Dropout",
dropout_input,
dropout_output,
{},
kMSDomain);
// Get attribute "seed" from "Dropout" node if available.
NodeAttributes dropout_attrs = dropout_node.GetAttributes();
NodeAttributes::const_iterator seed = dropout_attrs.find("seed");
if (seed != dropout_attrs.end()) {
dropout_add_fusion_node.AddAttribute("seed", seed->second);
}
// Assign provider to this new node. Provider should be same as the provider for old node.
dropout_add_fusion_node.SetExecutionProviderType(dropout_node.GetExecutionProviderType());
// delete bias_add_node, dropout_node and optionally residual_add_node
for (Node& n : nodes_to_fuse) {
graph_utils::RemoveNodeOutputEdges(graph, n);
graph.RemoveNode(n.Index());
}
modified = true;
}
return Status::OK();
}
} // namespace onnxruntime

View file

@ -0,0 +1,24 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#pragma once
#include "core/optimizer/graph_transformer.h"
namespace onnxruntime {
/**
@Class BiasDropoutFusion
Fuse Add + Dropout + optional Add to BiasDropoutFusion
*/
class BiasDropoutFusion : public GraphTransformer {
public:
BiasDropoutFusion(const std::unordered_set<std::string>& compatible_execution_providers = {}) noexcept
: GraphTransformer("BiasDropoutFusion", compatible_execution_providers) {}
Status ApplyImpl(Graph& graph, bool& modified, int graph_level, const logging::Logger& logger) const override;
};
} // namespace onnxruntime

View file

@ -6,6 +6,7 @@
#include "orttraining/core/framework/distributed_run_context.h"
#include "orttraining/core/optimizer/insert_output_rewriter.h"
#include "orttraining/core/optimizer/megatron_transformer.h"
#include "orttraining/core/optimizer/bias_dropout_fusion.h"
#include "orttraining/core/optimizer/nonzero_shape_setter.h"
#include "core/optimizer/identity_elimination.h"
#include "core/optimizer/slice_elimination.h"
@ -143,6 +144,7 @@ std::vector<std::unique_ptr<GraphTransformer>> GenerateTransformers(TransformerL
transformers.emplace_back(onnxruntime::make_unique<MatMulAddFusion>(l1_execution_providers));
transformers.emplace_back(onnxruntime::make_unique<FreeDimensionOverrideTransformer>(free_dimension_overrides));
transformers.emplace_back(onnxruntime::make_unique<MatmulTransposeFusion>(l1_execution_providers));
transformers.emplace_back(onnxruntime::make_unique<BiasDropoutFusion>(l1_execution_providers));
rule_transformer = optimizer_utils::GenerateRuleBasedGraphTransformer(level, transformers_and_rules_to_enable, l1_execution_providers);
} break;

View file

@ -10,11 +10,13 @@
#include "gtest/gtest.h"
#include "core/optimizer/rule_based_graph_transformer.h"
#include "core/optimizer/utils.h"
#include "orttraining/core/optimizer/bias_dropout_fusion.h"
#include "orttraining/core/optimizer/gist_encode_decode.h"
#include "orttraining/core/optimizer/nonzero_shape_setter.h"
#include "orttraining/core/optimizer/megatron_transformer.h"
#include "test/optimizer/graph_transform_test_fixture.h"
#include "test/util/include/default_providers.h"
#include "test/util/include/asserts.h"
#include "orttraining/test/optimizer/horizontal_parallel_test_utils.h"
#include <random>
@ -45,6 +47,33 @@ TEST_F(GraphTransformationTests, GistEncodeDecode) {
ASSERT_TRUE(op_to_count["GistBinarizeEncoder"] == op_to_count["GistBinarizeEncoder"]);
}
static void TestBiasDropoutFusion(const PathString& file_path, const logging::Logger& logger, const int add_count = 0) {
std::shared_ptr<Model> p_model;
ASSERT_TRUE(Model::Load(file_path, p_model, nullptr, logger).IsOK());
Graph& graph = p_model->MainGraph();
onnxruntime::GraphTransformerManager graph_transformation_mgr{5};
graph_transformation_mgr.Register(onnxruntime::make_unique<BiasDropoutFusion>(), TransformerLevel::Level2);
auto ret = graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level2, logger);
ASSERT_STATUS_OK(ret);
std::map<std::string, int> op_to_count = CountOpsInGraph(graph);
ASSERT_EQ(op_to_count["Add"], add_count);
ASSERT_EQ(op_to_count["Dropout"], 0);
ASSERT_EQ(op_to_count["TrainableDropout"], 0);
ASSERT_EQ(op_to_count["BiasDropout"], 1);
}
TEST_F(GraphTransformationTests, BiasDropoutFusionTest) {
TestBiasDropoutFusion(MODEL_FOLDER "fusion/bias_dropout_fusion1.onnx", *logger_);
TestBiasDropoutFusion(MODEL_FOLDER "fusion/bias_dropout_fusion2.onnx", *logger_);
TestBiasDropoutFusion(MODEL_FOLDER "fusion/bias_dropout_residual_fusion1.onnx", *logger_);
TestBiasDropoutFusion(MODEL_FOLDER "fusion/bias_dropout_residual_fusion2.onnx", *logger_);
TestBiasDropoutFusion(MODEL_FOLDER "fusion/bias_dropout_residual_fusion_mismatch.onnx", *logger_, 1);
TestBiasDropoutFusion(MODEL_FOLDER "fusion/bias_trainabledropout_residual_fusion.onnx", *logger_);
}
Node* GetNodeByName(Graph& graph, std::string node_name) {
GraphViewer graph_viewer(graph);
const auto& node_topology_list = graph_viewer.GetNodesInTopologicalOrder();

View file

@ -14,6 +14,7 @@
#include "gtest/gtest.h"
#include "test/common/tensor_op_test_utils.h"
#include "test/providers/provider_test_utils.h"
#include "test/util/include/default_providers.h"
@ -33,7 +34,7 @@ const Tensor& FetchTensor(const OrtValue& ort_value) {
return ort_value.Get<Tensor>();
}
void RunDropoutTest(const char* op, const bool use_mask, const std::vector<int64_t>& input_shape, float ratio = -1,
void RunDropoutTest(const char* op, const bool use_mask, const std::vector<int64_t>& input_shape, float ratio = -1.0f,
bool training_mode = true, bool use_float16_ratio = false) {
OpTester t{op, k_dropout_opset_version, kOnnxDomain};
@ -45,13 +46,21 @@ void RunDropoutTest(const char* op, const bool use_mask, const std::vector<int64
t.AddAttribute("seed", seed);
t.AddInput("data", input_shape, input);
if (ratio == -1) {
ratio = 0.5; // default.
t.AddInput("ratio", {}, {ratio});
} else if (use_float16_ratio) {
t.AddInput("ratio", {}, {MLFloat16(0)});
if (ratio == -1.0f) {
if (use_float16_ratio) {
t.AddMissingOptionalInput<MLFloat16>();
} else {
t.AddMissingOptionalInput<float>();
}
// set ratio to default value
ratio = 0.5f;
} else {
t.AddInput("ratio", {}, {ratio});
if (use_float16_ratio) {
t.AddInput("ratio", {}, {MLFloat16(math::floatToHalf(ratio))});
} else {
t.AddInput("ratio", {}, {ratio});
}
}
if (strcmp(op, "TrainableDropout") != 0 && training_mode) {
@ -73,12 +82,12 @@ void RunDropoutTest(const char* op, const bool use_mask, const std::vector<int64
const auto& output_tensor = FetchTensor(fetches[0]);
auto output_span = output_tensor.DataAsSpan<float>();
const auto num_output_zeros = std::count(output_span.begin(), output_span.end(), 0.0f);
const auto num_dropped_values = std::count(output_span.begin(), output_span.end(), 0.0f);
if (ratio == 1.0f) {
ASSERT_EQ(num_output_zeros, static_cast<size_t>(output_span.size())) << "provider: " << provider_type;
ASSERT_EQ(num_dropped_values, static_cast<size_t>(output_span.size())) << "provider: " << provider_type;
} else {
ASSERT_NEAR(static_cast<float>(num_output_zeros) / static_cast<size_t>(output_span.size()), ratio, 0.1f)
ASSERT_NEAR(static_cast<float>(num_dropped_values) / static_cast<size_t>(output_span.size()), ratio, 0.1f)
<< "provider: " << provider_type;
for (decltype(output_span.size()) i = 0; i < output_span.size(); ++i) {
@ -96,7 +105,7 @@ void RunDropoutTest(const char* op, const bool use_mask, const std::vector<int64
ASSERT_EQ(mask_span.size(), output_span.size()) << "provider: " << provider_type;
const auto num_mask_zeros = std::count(mask_span.begin(), mask_span.end(), false);
ASSERT_EQ(num_output_zeros, num_mask_zeros) << "provider: " << provider_type;
ASSERT_EQ(num_dropped_values, num_mask_zeros) << "provider: " << provider_type;
for (decltype(mask_span.size()) i = 0; i < mask_span.size(); ++i) {
ASSERT_TRUE(
@ -109,16 +118,17 @@ void RunDropoutTest(const char* op, const bool use_mask, const std::vector<int64
t.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, nullptr, ExecutionMode::ORT_SEQUENTIAL, output_verifier);
}
} // namespace
// Dropout
TEST(DropoutTest, Basic) {
RunDropoutTest("Dropout", false, {10, 10, 10}, 0.75);
RunDropoutTest("Dropout", false, {10, 10, 10}, 0.75f);
}
TEST(DropoutTest, Mask) {
RunDropoutTest("Dropout", true, {1000}, 0.25);
RunDropoutTest("Dropout", true, {1000}, 0.25f);
}
TEST(DropoutTest, RatioLimit) {
@ -130,11 +140,11 @@ TEST(DropoutTest, EmptyRatio) {
}
TEST(TrainableDropoutTest, Basic) {
RunDropoutTest("TrainableDropout", false, {10, 10, 10}, 0.75);
RunDropoutTest("TrainableDropout", false, {10, 10, 10}, 0.75f);
}
TEST(TrainableDropoutTest, Mask) {
RunDropoutTest("TrainableDropout", true, {1000}, 0.25);
RunDropoutTest("TrainableDropout", true, {1000}, 0.25f);
}
TEST(TrainableDropoutTest, RatioLimit) {
@ -142,9 +152,132 @@ TEST(TrainableDropoutTest, RatioLimit) {
}
TEST(TrainableDropoutTest, EmptyRatio) {
RunDropoutTest("TrainableDropout", true, {1000}, -1);
RunDropoutTest("TrainableDropout", true, {1000});
}
// BiasDropout kernel is only implemented for CUDA
#ifdef USE_CUDA
namespace {
void RunBiasDropoutTest(const bool use_mask, const std::vector<int64_t>& input_shape, float ratio = -1.0f,
bool training_mode = true, bool use_float16_ratio = false, bool has_residual = true) {
OpTester t{"BiasDropout", 1, kMSDomain};
const int64_t seed = 42;
t.AddAttribute("seed", seed);
const auto input_size = std::accumulate(
input_shape.begin(), input_shape.end(), static_cast<int64_t>(1), std::multiplies<>{});
const std::vector<float> input = ValueRange(input_size, 1.0f, 1.0f);
t.AddInput("data", input_shape, input);
std::vector<int64_t> bias_shape{input_shape.back()};
const auto bias_size = input_shape.back();
const std::vector<float> bias = ValueRange(bias_size, 2.0f, 1.0f);
t.AddInput("bias", bias_shape, bias);
float residual_value = 0.0f;
if (has_residual) {
residual_value = 1.0f;
const auto residual_size = input_size;
const std::vector<float> residual(residual_size, residual_value);
t.AddInput("residual", input_shape, residual);
} else {
t.AddMissingOptionalInput<float>();
}
if (ratio == -1.0f) {
if (use_float16_ratio) {
t.AddMissingOptionalInput<MLFloat16>();
} else {
t.AddMissingOptionalInput<float>();
}
// set ratio to default value
ratio = 0.5f;
} else {
if (use_float16_ratio) {
t.AddInput("ratio", {}, {MLFloat16(math::floatToHalf(ratio))});
} else {
t.AddInput("ratio", {}, {ratio});
}
}
if (training_mode) {
t.AddInput("training_mode", {}, {true});
}
t.AddOutput<float>("output", input_shape, input); // we'll do our own output verification
std::unique_ptr<bool[]> mask_buffer{};
if (use_mask) {
mask_buffer = onnxruntime::make_unique<bool[]>(input_size);
t.AddOutput<bool>("mask", input_shape, mask_buffer.get(), input_size);
} else {
t.AddMissingOptionalOutput<bool>();
}
auto output_verifier = [&](const std::vector<OrtValue>& fetches, const std::string& provider_type) {
ASSERT_GE(fetches.size(), 1);
const auto& output_tensor = FetchTensor(fetches[0]);
auto output_span = output_tensor.DataAsSpan<float>();
const auto num_dropped_values = std::count(output_span.begin(), output_span.end(), residual_value);
if (ratio == 1.0f) {
ASSERT_EQ(num_dropped_values, static_cast<size_t>(output_span.size())) << "provider: " << provider_type;
} else {
ASSERT_NEAR(static_cast<float>(num_dropped_values) / static_cast<size_t>(output_span.size()), ratio, 0.1f)
<< "provider: " << provider_type;
for (decltype(output_span.size()) i = 0; i < output_span.size(); ++i) {
if (output_span[i] == residual_value) continue;
const auto expected_value = (bias[i % bias_size] + i + 1.0f) / (1 - ratio) + residual_value;
ASSERT_NEAR(output_span[i], expected_value, 0.01f)
<< "unexpected output value at index " << i << ", provider: " << provider_type;
}
}
if (use_mask) {
ASSERT_GE(fetches.size(), 2);
const auto& mask_tensor = FetchTensor(fetches[1]);
auto mask_span = mask_tensor.DataAsSpan<bool>();
ASSERT_EQ(mask_span.size(), output_span.size()) << "provider: " << provider_type;
const auto num_mask_zeros = std::count(mask_span.begin(), mask_span.end(), false);
ASSERT_EQ(num_dropped_values, num_mask_zeros) << "provider: " << provider_type;
for (decltype(mask_span.size()) i = 0; i < mask_span.size(); ++i) {
ASSERT_TRUE(
(mask_span[i] && output_span[i] != residual_value) || (!mask_span[i] && output_span[i] == residual_value))
<< "output and mask mismatch at index " << i << ", output[i]: " << output_span[i]
<< ", mask[i]: " << mask_span[i] << ", provider: " << provider_type;
}
}
};
t.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, nullptr, ExecutionMode::ORT_SEQUENTIAL, output_verifier);
}
} // namespace
TEST(BiasDropoutTest, Basic) {
RunBiasDropoutTest(false, {10, 10, 10}, 0.75f);
}
TEST(BiasDropoutTest, BasicWithoutResidual) {
RunBiasDropoutTest(false, {10, 10, 10}, 0.75f, true, false, false);
}
TEST(BiasDropoutTest, Mask) {
RunBiasDropoutTest(true, {3, 5, 768}, 0.25f);
}
TEST(BiasDropoutTest, RatioLimit) {
RunBiasDropoutTest(true, {4, 8, 1024}, 0.0f, false);
}
TEST(BiasDropoutTest, EmptyRatio) {
RunBiasDropoutTest(true, {2, 7, 1024});
}
#endif
namespace {
void RunDropoutGradTest(const char* op, float ratio, const std::vector<int64_t>& input_dims, bool default_ratio = true) {
const auto input_shape = TensorShape(input_dims);

View file

@ -52,33 +52,10 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1
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_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, GatherGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, MLFloat16_MLFloat16, TrainableDropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, MLFloat16_float, TrainableDropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, MLFloat16_double, TrainableDropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float_MLFloat16, TrainableDropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float_float, TrainableDropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float_double, TrainableDropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, double_MLFloat16, TrainableDropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, double_float, TrainableDropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, double_double, TrainableDropout);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16, TrainableDropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float, TrainableDropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_double, TrainableDropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_MLFloat16, TrainableDropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float, TrainableDropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_double, TrainableDropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_MLFloat16, TrainableDropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_float, TrainableDropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double, TrainableDropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16, DropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float, DropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_double, DropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_MLFloat16, DropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float, DropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_double, DropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_MLFloat16, DropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_float, DropoutGrad);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double, DropoutGrad);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BiasDropout);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, TrainableDropout);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, TrainableDropoutGrad);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, DropoutGrad);
// TODO: decprecate GatherND-1 after updating training models to opset-12
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, int64_t, GatherND);
@ -141,133 +118,111 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, Mega
Status RegisterCudaTrainingKernels(KernelRegistry& kernel_registry) {
static const BuildKernelCreateInfoFn function_table[] = {
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, View)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, Group)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, SGDOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, View)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, Group)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, SGDOptimizer)>,
// Adam
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_int64_t_float_float_float_float, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_int64_t_float_MLFloat16_float_float, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_int64_t_float_MLFloat16_float_float, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_int64_t_float_float_MLFloat16_MLFloat16, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_int64_t_float_float_MLFloat16_float, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_int64_t_float_MLFloat16_MLFloat16_MLFloat16, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_int64_t_float_MLFloat16_MLFloat16_float, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_int64_t_float_MLFloat16_MLFloat16_MLFloat16, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_int64_t_float_MLFloat16_MLFloat16_float, AdamOptimizer)>,
// Adam
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_int64_t_float_float_float_float, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_int64_t_float_MLFloat16_float_float, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_int64_t_float_MLFloat16_float_float, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_int64_t_float_float_MLFloat16_MLFloat16, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_int64_t_float_float_MLFloat16_float, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_int64_t_float_MLFloat16_MLFloat16_MLFloat16, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_int64_t_float_MLFloat16_MLFloat16_float, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_int64_t_float_MLFloat16_MLFloat16_MLFloat16, AdamOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_int64_t_float_MLFloat16_MLFloat16_float, AdamOptimizer)>,
// Lamb
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float_float_float_float, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float_MLFloat16_float_MLFloat16, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float_MLFloat16_float_float, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double_double_double_double, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float_MLFloat16_MLFloat16_MLFloat16, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float_MLFloat16_MLFloat16_float, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float_MLFloat16_float_MLFloat16, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float_MLFloat16_float_float, LambOptimizer)>,
// Lamb
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float_float_float_float, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float_MLFloat16_float_MLFloat16, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float_MLFloat16_float_float, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double_double_double_double, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float_MLFloat16_MLFloat16_MLFloat16, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float_MLFloat16_MLFloat16_float, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float_MLFloat16_float_MLFloat16, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float_MLFloat16_float_float, LambOptimizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float, InPlaceAccumulator)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_MLFloat16, InPlaceAccumulator)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16, InPlaceAccumulator)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float, InPlaceAccumulator)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float, InPlaceAccumulator)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_MLFloat16, InPlaceAccumulator)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16, InPlaceAccumulator)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float, InPlaceAccumulator)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, ZeroGradient)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, ZeroGradient)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, MLFloat16_MLFloat16, TrainableDropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, MLFloat16_float, TrainableDropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, MLFloat16_double, TrainableDropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float_MLFloat16, TrainableDropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float_float, TrainableDropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float_double, TrainableDropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, double_MLFloat16, TrainableDropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, double_float, TrainableDropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, double_double, TrainableDropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16, TrainableDropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float, TrainableDropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_double, TrainableDropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_MLFloat16, TrainableDropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float, TrainableDropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_double, TrainableDropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_MLFloat16, TrainableDropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_float, TrainableDropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double, TrainableDropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16, DropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float, DropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_double, DropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_MLFloat16, DropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float, DropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_double, DropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_MLFloat16, DropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_float, DropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double, DropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, ZeroGradient)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, ZeroGradient)>,
// TODO: decprecate GatherND-1 after updating training models to opset-12
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, int64_t, GatherND)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, int64_t, GatherNDGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, SoftmaxCrossEntropy)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, SoftmaxCrossEntropyGrad)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float, int32_t, SparseSoftmaxCrossEntropy)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float, int64_t, SparseSoftmaxCrossEntropy)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float, int32_t, SparseSoftmaxCrossEntropyGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float, int64_t, SparseSoftmaxCrossEntropyGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, SoftmaxGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, SoftmaxGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, SoftmaxGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, float, int64_t, SoftmaxCrossEntropyLoss)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, int64_t, SoftmaxCrossEntropyLossGrad)>,
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_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, GatherGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, DivGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, DivGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, DivGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, GeluGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, GeluGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, GeluGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, FastGeluGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, FastGeluGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, FastGeluGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BiasGeluGrad_dX)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BiasFastGeluGrad_dX)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, IsFinite)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, IsFinite)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, IsFinite)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, bool, All)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, IsAllFinite)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, IsAllFinite)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, IsAllFinite)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, MixedPrecisionScale)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, MixedPrecisionScale)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float, ReduceAllL2)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float, ReduceAllL2)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_MLFloat16, ReduceAllL2)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16, ReduceAllL2)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float, LayerNormalizationGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_float, LayerNormalizationGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float, LayerNormalizationGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, SliceGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, GatherElementsGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BiasDropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, TrainableDropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, TrainableDropoutGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, DropoutGrad)>,
// P2P communication operators.
// TODO: decprecate GatherND-1 after updating training models to opset-12
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, int64_t, GatherND)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, int64_t, GatherNDGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, SoftmaxCrossEntropy)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, SoftmaxCrossEntropyGrad)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float, int32_t, SparseSoftmaxCrossEntropy)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float, int64_t, SparseSoftmaxCrossEntropy)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float, int32_t, SparseSoftmaxCrossEntropyGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, float, int64_t, SparseSoftmaxCrossEntropyGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, SoftmaxGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, SoftmaxGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, SoftmaxGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, float, int64_t, SoftmaxCrossEntropyLoss)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, int64_t, SoftmaxCrossEntropyLossGrad)>,
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_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, GatherGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, DivGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, DivGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, DivGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, GeluGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, GeluGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, GeluGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, FastGeluGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, FastGeluGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, FastGeluGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BiasGeluGrad_dX)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, BiasFastGeluGrad_dX)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, IsFinite)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, IsFinite)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, IsFinite)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, bool, All)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, IsAllFinite)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, IsAllFinite)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, IsAllFinite)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, MixedPrecisionScale)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, MixedPrecisionScale)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float, ReduceAllL2)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float, ReduceAllL2)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_MLFloat16, ReduceAllL2)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16, ReduceAllL2)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float, LayerNormalizationGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_float, LayerNormalizationGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float, LayerNormalizationGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, SliceGrad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, GatherElementsGrad)>,
// P2P communication operators.
#if defined(USE_NCCL) || defined(USE_HOROVOD)
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, Send)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, Recv)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, Send)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, Recv)>,
#endif
#ifdef USE_HOROVOD
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, HorovodAllReduce)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, HorovodBarrier)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, HorovodAllReduce)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, HorovodBarrier)>,
#endif
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, RecordEvent)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, WaitEvent)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, RecordEvent)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, WaitEvent)>,
#ifdef USE_NCCL
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, NcclAllReduce)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, NcclAllGather)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, NcclReduceScatter)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MegatronF)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MegatronG)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, NcclAllReduce)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, NcclAllGather)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, NcclReduceScatter)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MegatronF)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MegatronG)>,
#endif
};

View file

@ -4,104 +4,187 @@
#include "core/framework/random_seed.h"
#include "orttraining/training_ops/cuda/nn/dropout.h"
#include "core/providers/cuda/nn/dropout.h"
#include "core/providers/cuda/cuda_common.h"
#include "core/providers/common.h"
namespace onnxruntime {
namespace cuda {
#define REGISTER_TRAINABLE_KERNEL_TYPED(T1, T2) \
ONNX_OPERATOR_TYPED_KERNEL_EX( \
TrainableDropout, \
kOnnxDomain, \
9, \
T1##_##T2, \
kCudaExecutionProvider, \
KernelDefBuilder() \
.TypeConstraint("T", DataTypeImpl::GetTensorType<T1>()) \
.TypeConstraint("T1", DataTypeImpl::GetTensorType<T2>()) \
.InputMemoryType<OrtMemTypeCPUInput>(1), \
Dropout<T1, T2, true>);
// Temporary for backward compatibility, will eventually get rid of TrainableDropout when PyTorch exporter will move to
// opset-12.
ONNX_OPERATOR_KERNEL_EX(
TrainableDropout,
kOnnxDomain,
9,
kCudaExecutionProvider,
KernelDefBuilder()
.TypeConstraint("T", DataTypeImpl::AllIEEEFloatTensorTypes())
.TypeConstraint("T1", DataTypeImpl::AllIEEEFloatTensorTypes())
.InputMemoryType<OrtMemTypeCPUInput>(1),
Dropout<true>);
#define REGISTER_GRADIENT_KERNEL(OpName) \
ONNX_OPERATOR_KERNEL_EX( \
OpName, \
kMSDomain, \
1, \
kCudaExecutionProvider, \
KernelDefBuilder() \
.TypeConstraint("T", DataTypeImpl::AllIEEEFloatTensorTypes()) \
.TypeConstraint("T1", DataTypeImpl::AllIEEEFloatTensorTypes()) \
.TypeConstraint("T2", DataTypeImpl::GetTensorType<bool>()) \
.InputMemoryType<OrtMemTypeCPUInput>(2), \
DropoutGrad);
REGISTER_GRADIENT_KERNEL(DropoutGrad)
// Temporary for backward compatibility, will eventually get rid of TrainableDropout when PyTorch exporter will move to
// opset-12.
REGISTER_TRAINABLE_KERNEL_TYPED(MLFloat16, MLFloat16)
REGISTER_TRAINABLE_KERNEL_TYPED(MLFloat16, float)
REGISTER_TRAINABLE_KERNEL_TYPED(MLFloat16, double)
REGISTER_TRAINABLE_KERNEL_TYPED(float, MLFloat16)
REGISTER_TRAINABLE_KERNEL_TYPED(float, float)
REGISTER_TRAINABLE_KERNEL_TYPED(float, double)
REGISTER_TRAINABLE_KERNEL_TYPED(double, MLFloat16)
REGISTER_TRAINABLE_KERNEL_TYPED(double, float)
REGISTER_TRAINABLE_KERNEL_TYPED(double, double)
REGISTER_GRADIENT_KERNEL(TrainableDropoutGrad)
#define REGISTER_GRADIENT_KERNEL_TYPED(OpName, T1, T2) \
ONNX_OPERATOR_TYPED_KERNEL_EX( \
OpName, \
kMSDomain, \
1, \
T1##_##T2, \
kCudaExecutionProvider, \
KernelDefBuilder() \
.TypeConstraint("T", DataTypeImpl::GetTensorType<T1>()) \
.TypeConstraint("T1", DataTypeImpl::GetTensorType<T2>()) \
.TypeConstraint("T2", DataTypeImpl::GetTensorType<bool>()) \
.InputMemoryType<OrtMemTypeCPUInput>(2) \
.InputMemoryType<OrtMemTypeCPUInput>(3), \
DropoutGrad<T1, T2>);
template <typename T>
struct DropoutGradComputeImpl {
void operator()(const int64_t N,
const Tensor& dY,
const bool* mask_data,
const float ratio_data,
Tensor& dX) const {
typedef typename ToCudaType<T>::MappedType CudaT;
REGISTER_GRADIENT_KERNEL_TYPED(DropoutGrad, MLFloat16, MLFloat16)
REGISTER_GRADIENT_KERNEL_TYPED(DropoutGrad, MLFloat16, float)
REGISTER_GRADIENT_KERNEL_TYPED(DropoutGrad, MLFloat16, double)
REGISTER_GRADIENT_KERNEL_TYPED(DropoutGrad, float, MLFloat16)
REGISTER_GRADIENT_KERNEL_TYPED(DropoutGrad, float, float)
REGISTER_GRADIENT_KERNEL_TYPED(DropoutGrad, float, double)
REGISTER_GRADIENT_KERNEL_TYPED(DropoutGrad, double, MLFloat16)
REGISTER_GRADIENT_KERNEL_TYPED(DropoutGrad, double, float)
REGISTER_GRADIENT_KERNEL_TYPED(DropoutGrad, double, double)
// Temporary for backward compatibility, will eventually get rid of TrainableDropout when PyTorch exporter will move to
// opset-12.
REGISTER_GRADIENT_KERNEL_TYPED(TrainableDropoutGrad, MLFloat16, MLFloat16)
REGISTER_GRADIENT_KERNEL_TYPED(TrainableDropoutGrad, MLFloat16, float)
REGISTER_GRADIENT_KERNEL_TYPED(TrainableDropoutGrad, MLFloat16, double)
REGISTER_GRADIENT_KERNEL_TYPED(TrainableDropoutGrad, float, MLFloat16)
REGISTER_GRADIENT_KERNEL_TYPED(TrainableDropoutGrad, float, float)
REGISTER_GRADIENT_KERNEL_TYPED(TrainableDropoutGrad, float, double)
REGISTER_GRADIENT_KERNEL_TYPED(TrainableDropoutGrad, double, MLFloat16)
REGISTER_GRADIENT_KERNEL_TYPED(TrainableDropoutGrad, double, float)
REGISTER_GRADIENT_KERNEL_TYPED(TrainableDropoutGrad, double, double)
template <typename T1, typename T2>
Status DropoutGrad<T1, T2>::ComputeInternal(OpKernelContext* context) const {
typedef typename ToCudaType<T1>::MappedType CudaT;
const CudaT* dY_data = reinterpret_cast<const CudaT*>(dY.template Data<T>());
CudaT* dX_data = reinterpret_cast<CudaT*>(dX.template MutableData<T>());
DropoutGradientKernelImpl<CudaT>(N, dY_data, mask_data, ratio_data, dX_data);
}
};
Status DropoutGrad::ComputeInternal(OpKernelContext* context) const {
auto dY = context->Input<Tensor>(0);
const TensorShape& shape = dY->Shape();
auto dY_data = reinterpret_cast<const CudaT*>(dY->template Data<T1>());
const int64_t N = shape.Size();
auto mask = context->Input<Tensor>(1);
ORT_ENFORCE(mask->Shape().Size() == N);
const bool* mask_data = mask->template Data<bool>();
//Get the ratio_data
float ratio_data = default_ratio_;
auto ratio = context->Input<Tensor>(2);
if (ratio) {
utils::MLTypeCallDispatcher<GetRatioDataImpl, float, MLFloat16, double> t_disp(ratio->GetElementType());
t_disp.Invoke(ratio, ratio_data);
}
auto dX = context->Output(0, shape);
auto dX_data = reinterpret_cast<CudaT*>(dX->template MutableData<T1>());
float ratio_data;
auto ratio = context->Input<Tensor>(2);
static_assert(std::is_same<T2, MLFloat16>::value || std::is_same<T2, float>::value || std::is_same<T2, double>::value,
"T2 must be float16 or float or double");
if (ratio) {
ratio_data = static_cast<float>(*(ratio->template Data<T2>()));
} else {
ratio_data = default_ratio_;
}
ORT_ENFORCE(ratio_data >= 0.0f && ratio_data < 1.0f);
const bool* mask_data = mask->template Data<bool>();
DropoutGradientKernelImpl(N, dY_data, mask_data, ratio_data, dX_data);
utils::MLTypeCallDispatcher<DropoutGradComputeImpl, float, MLFloat16, double> t_disp(dY->GetElementType());
t_disp.Invoke(N, *dY, mask_data, ratio_data, *dX);
return Status::OK();
}
ONNX_OPERATOR_KERNEL_EX(
BiasDropout,
kMSDomain,
1,
kCudaExecutionProvider,
KernelDefBuilder()
.TypeConstraint("T", DataTypeImpl::AllIEEEFloatTensorTypes())
.TypeConstraint("T1", DataTypeImpl::AllIEEEFloatTensorTypes())
.TypeConstraint("T2", DataTypeImpl::GetTensorType<bool>())
.InputMemoryType<OrtMemTypeCPUInput>(3)
.InputMemoryType<OrtMemTypeCPUInput>(4),
BiasDropout);
template <typename T>
struct BiasDropoutComputeImpl {
Status operator()(const cudaDeviceProp& prop,
const int64_t N,
const fast_divmod fdm_dim,
const float ratio_data,
PhiloxGenerator& generator,
const Tensor& X,
const Tensor& bias,
const Tensor* residual,
Tensor& Y,
bool* mask_data) const {
typedef typename ToCudaType<T>::MappedType CudaT;
const CudaT* X_data = reinterpret_cast<const CudaT*>(X.template Data<T>());
const CudaT* bias_data = reinterpret_cast<const CudaT*>(bias.template Data<T>());
const CudaT* residual_data = nullptr;
if (residual) {
if (residual->Shape() != X.Shape()) {
return Status(common::ONNXRUNTIME, common::FAIL, "Residual input shape does not match X input shape.");
}
residual_data = reinterpret_cast<const CudaT*>(residual->template Data<T>());
}
CudaT* Y_data = reinterpret_cast<CudaT*>(Y.template MutableData<T>());
BiasDropoutKernelImpl<CudaT>(prop, N, fdm_dim, ratio_data, generator, X_data, bias_data, residual_data, Y_data, mask_data);
return Status::OK();
}
};
Status BiasDropout::ComputeInternal(OpKernelContext* context) const {
//Get X_data
const Tensor* X = context->Input<Tensor>(0);
ORT_RETURN_IF_NOT(X, "X Input is not available.");
const TensorShape& x_shape = X->Shape();
const int64_t N = x_shape.Size();
//Get bias_data
const Tensor* bias = context->Input<Tensor>(1);
if (bias == nullptr) return Status(common::ONNXRUNTIME, common::FAIL, "Bias input of BiasDropout is not available.");
const TensorShape& bias_shape = bias->Shape();
if (bias_shape.NumDimensions() != 1) {
return Status(common::ONNXRUNTIME, common::FAIL, "Bias input is not a 1D tensor.");
}
const int64_t dim = bias_shape[0];
if (dim != x_shape.GetDims().back()) {
return Status(common::ONNXRUNTIME, common::FAIL, "Bias' dimension doesn't match input's last dimension.");
}
//Get residual_data
const Tensor* residual = context->Input<Tensor>(2);
//Get Y_data
auto Y = context->Output(0, x_shape);
//Get mask_data
auto mask = context->Output(1, x_shape);
//Get the ratio_data
float ratio_data = default_ratio_;
auto ratio = context->Input<Tensor>(3);
if (ratio) {
utils::MLTypeCallDispatcher<GetRatioDataImpl, float, MLFloat16, double> t_disp(ratio->GetElementType());
t_disp.Invoke(ratio, ratio_data);
}
//Check for inference mode.
const Tensor* training_mode = context->Input<Tensor>(4);
bool is_training_mode = (training_mode != nullptr) && training_mode->Data<bool>();
if (!is_training_mode) {
ratio_data = 0.0f;
}
IAllocatorUniquePtr<bool> temp_mask_buffer{}; // buffer to use if mask is not provided
bool* const mask_data = [this, N, mask, &temp_mask_buffer]() {
if (mask) return mask->MutableData<bool>();
temp_mask_buffer = GetScratchBuffer<bool>(N);
return temp_mask_buffer.get();
}();
const fast_divmod fdm_dim(gsl::narrow_cast<int>(dim));
PhiloxGenerator& generator = generator_ ? *generator_ : PhiloxGenerator::Default();
utils::MLTypeCallDispatcherRet<Status, BiasDropoutComputeImpl, float, MLFloat16, double> t_disp(X->GetElementType());
return t_disp.Invoke(GetDeviceProp(), N, fdm_dim, ratio_data, generator, *X, *bias, residual, *Y, mask_data);
}
} // namespace cuda
} // namespace onnxruntime

View file

@ -9,16 +9,31 @@
namespace onnxruntime {
namespace cuda {
template <typename T1, typename T2>
class DropoutGrad final : public CudaKernel {
public:
DropoutGrad(const OpKernelInfo& info) : CudaKernel(info), default_ratio_(0.5) {
DropoutGrad(const OpKernelInfo& info) : CudaKernel(info) {
}
Status ComputeInternal(OpKernelContext* context) const override;
private:
const float default_ratio_;
static constexpr float default_ratio_ = 0.5f;
};
class BiasDropout final : public CudaKernel {
public:
BiasDropout(const OpKernelInfo& info) : CudaKernel(info) {
int64_t seed = 0;
if (info.GetAttr<int64_t>("seed", &seed).IsOK()) {
generator_ = onnxruntime::make_unique<PhiloxGenerator>(static_cast<uint64_t>(seed));
}
}
Status ComputeInternal(OpKernelContext* context) const override;
private:
mutable std::unique_ptr<PhiloxGenerator> generator_;
static constexpr float default_ratio_ = 0.5f;
};
} // namespace cuda

View file

@ -24,15 +24,21 @@
namespace onnxruntime {
namespace cuda {
template <typename T>
template <typename T, int NumThreadsPerBlock, int NumElementsPerThread>
__global__ void DropoutGradientKernel(
const int64_t N,
const T* dY_data,
const bool* mask_data,
const T scale,
const float scale,
T* dX_data) {
CALCULATE_ELEMENTWISE_INDEX_OR_EXIT(id, N);
dX_data[id] = dY_data[id] * T(mask_data[id]) * scale;
CUDA_LONG id = NumElementsPerThread * NumThreadsPerBlock * blockIdx.x + threadIdx.x;
#pragma unroll
for (int i = 0; i < NumElementsPerThread; i++) {
if (id < N) {
dX_data[id] = T(float(dY_data[id]) * mask_data[id] * scale);
id += NumThreadsPerBlock;
}
}
}
template <typename T>
@ -48,8 +54,9 @@ void DropoutGradientKernelImpl(
}
} else {
const float scale = 1.f / (1.f - ratio);
const int blocksPerGrid = (N + GridDim::maxThreadsPerBlock - 1) / GridDim::maxThreadsPerBlock;
DropoutGradientKernel<T><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(N, dY_data, mask_data, T(scale), dX_data);
const int blocksPerGrid = static_cast<int>(CeilDiv(N, GridDim::maxThreadsPerBlock * GridDim::maxElementsPerThread));
DropoutGradientKernel<T, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread>
<<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(N, dY_data, mask_data, scale, dX_data);
}
}
@ -65,5 +72,103 @@ SPECIALIZED_DROPOUT_GRAD_IMPL(float)
SPECIALIZED_DROPOUT_GRAD_IMPL(double)
SPECIALIZED_DROPOUT_GRAD_IMPL(half)
constexpr int UNROLL = 4;
template <typename T, bool has_residual>
__global__ void BiasDropoutKernel(
const int64_t N,
const fast_divmod fdm_dim,
const float ratio,
const std::pair<uint64_t, uint64_t> seeds,
const T* X_data,
const T* bias_data,
const T* residual_data,
T* Y_data,
bool* mask_data) {
const float p = 1.0f - ratio;
const float scale = 1.0f / p;
CUDA_LONG idx = blockDim.x * blockIdx.x + threadIdx.x;
CUDA_LONG step_size = gridDim.x * blockDim.x * UNROLL;
CUDA_LONG rounded_size = ((N - 1) / step_size + 1) * step_size;
curandStatePhilox4_32_10_t state;
curand_init(seeds.first, idx, seeds.second, &state);
// We ensure every thread generates the same number of random numbers (by rounding
// up the size) and at the same timestep (by syncing threads).
// From CUDA curand documentation:
// The Philox_4x32_10 algorithm is closely tied to the thread and block count.
// Each thread computes 4 random numbers in the same time thus the most efficient
// use of Philox_4x32_10 is to generate a multiple of 4 times number of threads.
for (CUDA_LONG id = idx; id < rounded_size; id += step_size) {
float4 rand = curand_uniform4(&state);
#pragma unroll
for (CUDA_LONG i = 0; i < UNROLL; i++) {
CUDA_LONG li = id + gridDim.x * blockDim.x * i;
if (li < N) {
int offset = fdm_dim.mod(li);
float bias = float(bias_data[offset]);
mask_data[li] = (&rand.x)[i] < p;
float output_data = (float(X_data[li]) + bias) * mask_data[li] * scale;
if (has_residual) {
output_data += float(residual_data[li]);
}
Y_data[li] = T(output_data);
}
}
__syncthreads();
}
}
template <typename T>
void BiasDropoutKernelImpl(
const cudaDeviceProp& prop,
const int64_t N,
const fast_divmod fdm_dim,
const float ratio,
PhiloxGenerator& generator,
const T* X_data,
const T* bias_data,
const T* residual_data,
T* Y_data,
bool* mask_data) {
const int block_size = 256;
const int blocks_per_sm = prop.maxThreadsPerMultiProcessor / block_size;
const int grid_size = std::min(prop.multiProcessorCount * blocks_per_sm, static_cast<int>(CeilDiv(N, block_size * UNROLL)));
// Compute the number of random numbers generated by each thread, and increment philox generator offset by that amount.
const uint64_t counter_offset = static_cast<uint64_t>(((N - 1) / (block_size * grid_size * UNROLL) + 1) * UNROLL);
auto seeds = generator.NextPhiloxSeeds(counter_offset);
if (residual_data == nullptr) {
BiasDropoutKernel<T, false><<<grid_size, block_size, 0>>>(N, fdm_dim, ratio, seeds, X_data, bias_data, residual_data, Y_data, mask_data);
} else {
BiasDropoutKernel<T, true><<<grid_size, block_size, 0>>>(N, fdm_dim, ratio, seeds, X_data, bias_data, residual_data, Y_data, mask_data);
}
}
#define SPECIALIZED_BIAS_DROPOUT_IMPL(T) \
template void BiasDropoutKernelImpl( \
const cudaDeviceProp& prop, \
const int64_t N, \
const fast_divmod fdm_dim, \
const float ratio, \
PhiloxGenerator& generator, \
const T* X_data, \
const T* bias_data, \
const T* residual_data, \
T* Y_data, \
bool* mask_data);
SPECIALIZED_BIAS_DROPOUT_IMPL(float)
SPECIALIZED_BIAS_DROPOUT_IMPL(double)
SPECIALIZED_BIAS_DROPOUT_IMPL(half)
} // namespace cuda
} // namespace onnxruntime

View file

@ -16,5 +16,18 @@ void DropoutGradientKernelImpl(
const float ratio,
T* dX_data);
template <typename T>
void BiasDropoutKernelImpl(
const cudaDeviceProp& prop,
const int64_t N,
const fast_divmod fdm_dim,
const float ratio,
PhiloxGenerator& generator,
const T* X_data,
const T* bias_data,
const T* residual_data,
T* Y_data,
bool* mask_data);
} // namespace cuda
} // namespace onnxruntime