diff --git a/onnxruntime/core/providers/cuda/cuda_execution_provider.cc b/onnxruntime/core/providers/cuda/cuda_execution_provider.cc index 18a0e1a7ed..d0cd9c3e1e 100644 --- a/onnxruntime/core/providers/cuda/cuda_execution_provider.cc +++ b/onnxruntime/core/providers/cuda/cuda_execution_provider.cc @@ -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, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, }; diff --git a/onnxruntime/core/providers/cuda/nn/dropout.cc b/onnxruntime/core/providers/cuda/nn/dropout.cc index 199b987c81..56b22e3593 100644 --- a/onnxruntime/core/providers/cuda/nn/dropout.cc +++ b/onnxruntime/core/providers/cuda/nn/dropout.cc @@ -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()) \ - .TypeConstraint("T1", DataTypeImpl::GetTensorType()) \ - .TypeConstraint("T2", DataTypeImpl::GetTensorType()) \ - .InputMemoryType(1) \ - .InputMemoryType(2), \ - Dropout); - -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()) + .InputMemoryType(1) + .InputMemoryType(2), + Dropout); } // namespace cuda } // namespace onnxruntime diff --git a/onnxruntime/core/providers/cuda/nn/dropout.h b/onnxruntime/core/providers/cuda/nn/dropout.h index a9fc195738..0ef8456f44 100644 --- a/onnxruntime/core/providers/cuda/nn/dropout.h +++ b/onnxruntime/core/providers/cuda/nn/dropout.h @@ -12,10 +12,36 @@ namespace onnxruntime { namespace cuda { -template +template +struct GetRatioDataImpl { + void operator()(const Tensor* ratio, float& ratio_data) const { + ratio_data = static_cast(*(ratio->template Data())); + ORT_ENFORCE(ratio_data >= 0.0f && ratio_data < 1.0f, "ratio_data is outside range [0, 1)"); + } +}; + +template +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::MappedType CudaT; + + const CudaT* X_data = reinterpret_cast(X.template Data()); + CudaT* Y_data = reinterpret_cast(Y.template MutableData()); + + DropoutKernelImpl(prop, N, ratio_data, generator, X_data, Y_data, mask_data); + } +}; + +template 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("seed", &seed).IsOK()) { generator_ = onnxruntime::make_unique(static_cast(seed)); @@ -26,48 +52,40 @@ class Dropout final : public CudaKernel { private: mutable std::unique_ptr generator_; - const float default_ratio_; + static constexpr float default_ratio_ = 0.5f; }; -template -Status Dropout::ComputeInternal(OpKernelContext* context) const { - typedef typename ToCudaType::MappedType CudaT; - +template +Status Dropout::ComputeInternal(OpKernelContext* context) const { //Get X_data const Tensor* X = context->Input(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(X->template Data()); const int64_t N = shape.Size(); //Get Y_data auto Y = context->Output(0, shape); - auto Y_data = reinterpret_cast(Y->template MutableData()); //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(1); - - static_assert(std::is_same::value || std::is_same::value || std::is_same::value, - "T2 must be float16 or float or double"); - if (ratio) { - ratio_data = static_cast(*(ratio->template Data())); - } else { - ratio_data = default_ratio_; + utils::MLTypeCallDispatcher 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(2); //Check for inference mode. if ((0 == ratio_data /*Backward compat with TrainableDropout*/) || (!trainable_dropout && (training_mode == nullptr || *(training_mode->Data()) == 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::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 t_disp(X->GetElementType()); + t_disp.Invoke(GetDeviceProp(), N, ratio_data, generator, *X, *Y, mask_data); return Status::OK(); } diff --git a/onnxruntime/core/providers/cuda/nn/dropout_impl.cu b/onnxruntime/core/providers/cuda/nn/dropout_impl.cu index d0ebc21851..0abe3f774a 100644 --- a/onnxruntime/core/providers/cuda/nn/dropout_impl.cu +++ b/onnxruntime/core/providers/cuda/nn/dropout_impl.cu @@ -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(CeilDiv(N, block_size))); + const int grid_size = std::min(prop.multiProcessorCount * blocks_per_sm, static_cast(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(((N - 1) / (block_size * grid_size * UNROLL) + 1) * UNROLL); diff --git a/onnxruntime/test/python/onnxruntime_test_ort_trainer.py b/onnxruntime/test/python/onnxruntime_test_ort_trainer.py index 52790dd40a..8566c31577 100644 --- a/onnxruntime/test/python/onnxruntime_test_ort_trainer.py +++ b/onnxruntime/test/python/onnxruntime_test_ort_trainer.py @@ -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") diff --git a/onnxruntime/test/testdata/transform/fusion/bias_dropout_fusion1.onnx b/onnxruntime/test/testdata/transform/fusion/bias_dropout_fusion1.onnx new file mode 100644 index 0000000000..40ce7de832 Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/bias_dropout_fusion1.onnx differ diff --git a/onnxruntime/test/testdata/transform/fusion/bias_dropout_fusion2.onnx b/onnxruntime/test/testdata/transform/fusion/bias_dropout_fusion2.onnx new file mode 100644 index 0000000000..9075927dd0 Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/bias_dropout_fusion2.onnx differ diff --git a/onnxruntime/test/testdata/transform/fusion/bias_dropout_residual_fusion1.onnx b/onnxruntime/test/testdata/transform/fusion/bias_dropout_residual_fusion1.onnx new file mode 100644 index 0000000000..ab104bf831 Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/bias_dropout_residual_fusion1.onnx differ diff --git a/onnxruntime/test/testdata/transform/fusion/bias_dropout_residual_fusion2.onnx b/onnxruntime/test/testdata/transform/fusion/bias_dropout_residual_fusion2.onnx new file mode 100644 index 0000000000..dec9a1423d Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/bias_dropout_residual_fusion2.onnx differ diff --git a/onnxruntime/test/testdata/transform/fusion/bias_dropout_residual_fusion_mismatch.onnx b/onnxruntime/test/testdata/transform/fusion/bias_dropout_residual_fusion_mismatch.onnx new file mode 100644 index 0000000000..12559ee280 Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/bias_dropout_residual_fusion_mismatch.onnx differ diff --git a/onnxruntime/test/testdata/transform/fusion/bias_dropout_residual_gen.py b/onnxruntime/test/testdata/transform/fusion/bias_dropout_residual_gen.py new file mode 100644 index 0000000000..3b1f9f01fc --- /dev/null +++ b/onnxruntime/test/testdata/transform/fusion/bias_dropout_residual_gen.py @@ -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') \ No newline at end of file diff --git a/onnxruntime/test/testdata/transform/fusion/bias_trainabledropout_residual_fusion.onnx b/onnxruntime/test/testdata/transform/fusion/bias_trainabledropout_residual_fusion.onnx new file mode 100644 index 0000000000..6d24631a1c Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/bias_trainabledropout_residual_fusion.onnx differ diff --git a/orttraining/orttraining/core/graph/gradient_schema_defs.cc b/orttraining/orttraining/core/graph/gradient_schema_defs.cc index 31056a5bc4..72702786e4 100644 --- a/orttraining/orttraining/core/graph/gradient_schema_defs.cc +++ b/orttraining/orttraining/core/graph/gradient_schema_defs.cc @@ -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) diff --git a/orttraining/orttraining/core/optimizer/bias_dropout_fusion.cc b/orttraining/orttraining/core/optimizer/bias_dropout_fusion.cc new file mode 100644 index 0000000000..d0e467c35a --- /dev/null +++ b/orttraining/orttraining/core/optimizer/bias_dropout_fusion.cc @@ -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 + +using namespace ONNX_NAMESPACE; +using namespace ::onnxruntime::common; +namespace onnxruntime { + +void FuseResidualAddIfAny(Graph& graph, const Node& dropout_node, + std::vector& dropout_input, + std::vector& dropout_output, + std::vector>& 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> 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 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 diff --git a/orttraining/orttraining/core/optimizer/bias_dropout_fusion.h b/orttraining/orttraining/core/optimizer/bias_dropout_fusion.h new file mode 100644 index 0000000000..ae16619541 --- /dev/null +++ b/orttraining/orttraining/core/optimizer/bias_dropout_fusion.h @@ -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& 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 diff --git a/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc b/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc index 8cbc56d970..a9ac30528a 100644 --- a/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc +++ b/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc @@ -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> GenerateTransformers(TransformerL transformers.emplace_back(onnxruntime::make_unique(l1_execution_providers)); transformers.emplace_back(onnxruntime::make_unique(free_dimension_overrides)); transformers.emplace_back(onnxruntime::make_unique(l1_execution_providers)); + transformers.emplace_back(onnxruntime::make_unique(l1_execution_providers)); rule_transformer = optimizer_utils::GenerateRuleBasedGraphTransformer(level, transformers_and_rules_to_enable, l1_execution_providers); } break; diff --git a/orttraining/orttraining/test/optimizer/graph_transform_test.cc b/orttraining/orttraining/test/optimizer/graph_transform_test.cc index b41b01609f..2f4e119ec0 100644 --- a/orttraining/orttraining/test/optimizer/graph_transform_test.cc +++ b/orttraining/orttraining/test/optimizer/graph_transform_test.cc @@ -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 @@ -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 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(), TransformerLevel::Level2); + auto ret = graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level2, logger); + ASSERT_STATUS_OK(ret); + + std::map 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(); diff --git a/orttraining/orttraining/test/training_ops/cpu/nn/dropout_op_test.cc b/orttraining/orttraining/test/training_ops/cpu/nn/dropout_op_test.cc index da3ba31a01..70fff78a5d 100644 --- a/orttraining/orttraining/test/training_ops/cpu/nn/dropout_op_test.cc +++ b/orttraining/orttraining/test/training_ops/cpu/nn/dropout_op_test.cc @@ -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(); } -void RunDropoutTest(const char* op, const bool use_mask, const std::vector& input_shape, float ratio = -1, +void RunDropoutTest(const char* op, const bool use_mask, const std::vector& 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(); + } else { + t.AddMissingOptionalInput(); + } + // 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(); - 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(output_span.size())) << "provider: " << provider_type; + ASSERT_EQ(num_dropped_values, static_cast(output_span.size())) << "provider: " << provider_type; } else { - ASSERT_NEAR(static_cast(num_output_zeros) / static_cast(output_span.size()), ratio, 0.1f) + ASSERT_NEAR(static_cast(num_dropped_values) / static_cast(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& 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(1), std::multiplies<>{}); + const std::vector input = ValueRange(input_size, 1.0f, 1.0f); + t.AddInput("data", input_shape, input); + + std::vector bias_shape{input_shape.back()}; + const auto bias_size = input_shape.back(); + const std::vector 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 residual(residual_size, residual_value); + t.AddInput("residual", input_shape, residual); + } else { + t.AddMissingOptionalInput(); + } + + if (ratio == -1.0f) { + if (use_float16_ratio) { + t.AddMissingOptionalInput(); + } else { + t.AddMissingOptionalInput(); + } + // 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("output", input_shape, input); // we'll do our own output verification + + std::unique_ptr mask_buffer{}; + if (use_mask) { + mask_buffer = onnxruntime::make_unique(input_size); + t.AddOutput("mask", input_shape, mask_buffer.get(), input_size); + } else { + t.AddMissingOptionalOutput(); + } + + auto output_verifier = [&](const std::vector& fetches, const std::string& provider_type) { + ASSERT_GE(fetches.size(), 1); + const auto& output_tensor = FetchTensor(fetches[0]); + auto output_span = output_tensor.DataAsSpan(); + + 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(output_span.size())) << "provider: " << provider_type; + } else { + ASSERT_NEAR(static_cast(num_dropped_values) / static_cast(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(); + 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& input_dims, bool default_ratio = true) { const auto input_shape = TensorShape(input_dims); diff --git a/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc b/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc index 922f87c708..e648e6fb14 100644 --- a/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc +++ b/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc @@ -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, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, - // Adam - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + // Adam + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, - // Lamb - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + // Lamb + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, - // TODO: decprecate GatherND-1 after updating training models to opset-12 - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - // BuildKernelCreateInfo, - BuildKernelCreateInfo, - // BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, - // P2P communication operators. + // TODO: decprecate GatherND-1 after updating training models to opset-12 + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + // BuildKernelCreateInfo, + BuildKernelCreateInfo, + // BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + +// P2P communication operators. #if defined(USE_NCCL) || defined(USE_HOROVOD) - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, #endif #ifdef USE_HOROVOD - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, #endif - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, #ifdef USE_NCCL - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, #endif }; diff --git a/orttraining/orttraining/training_ops/cuda/nn/dropout.cc b/orttraining/orttraining/training_ops/cuda/nn/dropout.cc index fd84b6f3bb..425e68827b 100644 --- a/orttraining/orttraining/training_ops/cuda/nn/dropout.cc +++ b/orttraining/orttraining/training_ops/cuda/nn/dropout.cc @@ -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()) \ - .TypeConstraint("T1", DataTypeImpl::GetTensorType()) \ - .InputMemoryType(1), \ - Dropout); +// 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(1), + Dropout); + +#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()) \ + .InputMemoryType(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()) \ - .TypeConstraint("T1", DataTypeImpl::GetTensorType()) \ - .TypeConstraint("T2", DataTypeImpl::GetTensorType()) \ - .InputMemoryType(2) \ - .InputMemoryType(3), \ - DropoutGrad); +template +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::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 -Status DropoutGrad::ComputeInternal(OpKernelContext* context) const { - typedef typename ToCudaType::MappedType CudaT; + const CudaT* dY_data = reinterpret_cast(dY.template Data()); + CudaT* dX_data = reinterpret_cast(dX.template MutableData()); + DropoutGradientKernelImpl(N, dY_data, mask_data, ratio_data, dX_data); + } +}; +Status DropoutGrad::ComputeInternal(OpKernelContext* context) const { auto dY = context->Input(0); const TensorShape& shape = dY->Shape(); - auto dY_data = reinterpret_cast(dY->template Data()); const int64_t N = shape.Size(); auto mask = context->Input(1); ORT_ENFORCE(mask->Shape().Size() == N); + const bool* mask_data = mask->template Data(); + + //Get the ratio_data + float ratio_data = default_ratio_; + auto ratio = context->Input(2); + if (ratio) { + utils::MLTypeCallDispatcher t_disp(ratio->GetElementType()); + t_disp.Invoke(ratio, ratio_data); + } auto dX = context->Output(0, shape); - auto dX_data = reinterpret_cast(dX->template MutableData()); - float ratio_data; - auto ratio = context->Input(2); - static_assert(std::is_same::value || std::is_same::value || std::is_same::value, - "T2 must be float16 or float or double"); - - if (ratio) { - ratio_data = static_cast(*(ratio->template Data())); - } else { - ratio_data = default_ratio_; - } - ORT_ENFORCE(ratio_data >= 0.0f && ratio_data < 1.0f); - - const bool* mask_data = mask->template Data(); - DropoutGradientKernelImpl(N, dY_data, mask_data, ratio_data, dX_data); + utils::MLTypeCallDispatcher 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()) + .InputMemoryType(3) + .InputMemoryType(4), + BiasDropout); + +template +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::MappedType CudaT; + + const CudaT* X_data = reinterpret_cast(X.template Data()); + const CudaT* bias_data = reinterpret_cast(bias.template Data()); + + 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(residual->template Data()); + } + + CudaT* Y_data = reinterpret_cast(Y.template MutableData()); + + BiasDropoutKernelImpl(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(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(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(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(3); + if (ratio) { + utils::MLTypeCallDispatcher t_disp(ratio->GetElementType()); + t_disp.Invoke(ratio, ratio_data); + } + + //Check for inference mode. + const Tensor* training_mode = context->Input(4); + bool is_training_mode = (training_mode != nullptr) && training_mode->Data(); + if (!is_training_mode) { + ratio_data = 0.0f; + } + + IAllocatorUniquePtr 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(); + temp_mask_buffer = GetScratchBuffer(N); + return temp_mask_buffer.get(); + }(); + + const fast_divmod fdm_dim(gsl::narrow_cast(dim)); + PhiloxGenerator& generator = generator_ ? *generator_ : PhiloxGenerator::Default(); + + utils::MLTypeCallDispatcherRet 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 diff --git a/orttraining/orttraining/training_ops/cuda/nn/dropout.h b/orttraining/orttraining/training_ops/cuda/nn/dropout.h index 12084bfc7f..a92a10d38f 100644 --- a/orttraining/orttraining/training_ops/cuda/nn/dropout.h +++ b/orttraining/orttraining/training_ops/cuda/nn/dropout.h @@ -9,16 +9,31 @@ namespace onnxruntime { namespace cuda { -template 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("seed", &seed).IsOK()) { + generator_ = onnxruntime::make_unique(static_cast(seed)); + } + } + + Status ComputeInternal(OpKernelContext* context) const override; + + private: + mutable std::unique_ptr generator_; + static constexpr float default_ratio_ = 0.5f; }; } // namespace cuda diff --git a/orttraining/orttraining/training_ops/cuda/nn/dropout_impl.cu b/orttraining/orttraining/training_ops/cuda/nn/dropout_impl.cu index df296f0fb1..4bf303d67e 100644 --- a/orttraining/orttraining/training_ops/cuda/nn/dropout_impl.cu +++ b/orttraining/orttraining/training_ops/cuda/nn/dropout_impl.cu @@ -24,15 +24,21 @@ namespace onnxruntime { namespace cuda { -template +template __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 @@ -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<<>>(N, dY_data, mask_data, T(scale), dX_data); + const int blocksPerGrid = static_cast(CeilDiv(N, GridDim::maxThreadsPerBlock * GridDim::maxElementsPerThread)); + DropoutGradientKernel + <<>>(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 +__global__ void BiasDropoutKernel( + const int64_t N, + const fast_divmod fdm_dim, + const float ratio, + const std::pair 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 +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(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(((N - 1) / (block_size * grid_size * UNROLL) + 1) * UNROLL); + auto seeds = generator.NextPhiloxSeeds(counter_offset); + + if (residual_data == nullptr) { + BiasDropoutKernel<<>>(N, fdm_dim, ratio, seeds, X_data, bias_data, residual_data, Y_data, mask_data); + } else { + BiasDropoutKernel<<>>(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 diff --git a/orttraining/orttraining/training_ops/cuda/nn/dropout_impl.h b/orttraining/orttraining/training_ops/cuda/nn/dropout_impl.h index b75ee462a6..09444662af 100644 --- a/orttraining/orttraining/training_ops/cuda/nn/dropout_impl.h +++ b/orttraining/orttraining/training_ops/cuda/nn/dropout_impl.h @@ -16,5 +16,18 @@ void DropoutGradientKernelImpl( const float ratio, T* dX_data); +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); + } // namespace cuda } // namespace onnxruntime