From fe8d867efadf336bdd9dac9241d74a3ae6b45b6b Mon Sep 17 00:00:00 2001 From: Hubert Lu <55214931+hubertlu-tw@users.noreply.github.com> Date: Thu, 3 Mar 2022 08:07:15 -0800 Subject: [PATCH] Optimize BinaryElementWise and BiasGeluGrad kernels for AMD (#10594) * Optimize elementwise and biasgelugrad kernels for AMD * Clean up for BiasGeluGradDxKernel --- .../cuda/cu_inc/binary_elementwise_impl.cuh | 40 +++++++++++++------ .../cuda/activation/bias_gelu_grad_impl.cu | 13 ++++-- 2 files changed, 38 insertions(+), 15 deletions(-) diff --git a/onnxruntime/core/providers/cuda/cu_inc/binary_elementwise_impl.cuh b/onnxruntime/core/providers/cuda/cu_inc/binary_elementwise_impl.cuh index 069cf0658d..1f76a6c096 100644 --- a/onnxruntime/core/providers/cuda/cu_inc/binary_elementwise_impl.cuh +++ b/onnxruntime/core/providers/cuda/cu_inc/binary_elementwise_impl.cuh @@ -188,15 +188,24 @@ void BinaryElementWiseNoBroadcastImpl( size_t count) { if (count == 0) // special case where there's a dim value of 0 in the output shape return; + + #ifdef USE_ROCM + const int num_elements_per_thread = 2; + const int num_threads_per_block = 512; + #else + const int num_elements_per_thread = GridDim::maxElementsPerThread; + const int num_threads_per_block = GridDim::maxThreadsPerBlock; + #endif - int blocksPerGrid = static_cast(CeilDiv(count, GridDim::maxThreadsPerBlock * GridDim::maxElementsPerThread)); + int blocksPerGrid = static_cast(CeilDiv(count, num_threads_per_block * num_elements_per_thread)); CUDA_LONG N = static_cast(count); - _BinaryElementWiseSimple<<>>( + _BinaryElementWiseSimple<<>>( lhs_data, rhs_data, output_data, func, N); + } template @@ -216,32 +225,39 @@ void BinaryElementWiseImpl( if (count == 0) // special case where there's a dim value of 0 in the output shape return; - int blocksPerGrid = static_cast(CeilDiv(count, GridDim::maxThreadsPerBlock * GridDim::maxElementsPerThread)); + #ifdef USE_ROCM + const int num_elements_per_thread = 2; + const int num_threads_per_block = 512; + #else + const int num_elements_per_thread = GridDim::maxElementsPerThread; + const int num_threads_per_block = GridDim::maxThreadsPerBlock; + #endif + + int blocksPerGrid = static_cast(CeilDiv(count, num_threads_per_block * num_elements_per_thread)); CUDA_LONG N = static_cast(count); if (output_rank_or_simple_broadcast == static_cast(SimpleBroadcast::NoBroadcast)) { - _BinaryElementWiseSimple<<>>( + _BinaryElementWiseSimple<<>>( lhs_data, rhs_data, output_data, func, N); } else if (output_rank_or_simple_broadcast == static_cast(SimpleBroadcast::LeftScalar)) { - _BinaryElementWiseSimple<<>>( + _BinaryElementWiseSimple<<>>( lhs_data, rhs_data, output_data, func, N); } else if (output_rank_or_simple_broadcast == static_cast(SimpleBroadcast::RightScalar)) { - _BinaryElementWiseSimple<<>>( + _BinaryElementWiseSimple<<>>( lhs_data, rhs_data, output_data, func, N); } else if (output_rank_or_simple_broadcast == static_cast(SimpleBroadcast::RightPerChannelBatch1)) { - _BinaryElementWiseRhsPerChannelBatch1<<>>( + _BinaryElementWiseRhsPerChannelBatch1<<>>( lhs_data, rhs_data, fdm_H, @@ -249,7 +265,7 @@ void BinaryElementWiseImpl( func, N); } else if (output_rank_or_simple_broadcast == static_cast(SimpleBroadcast::RightPerChannelBatchN)) { - _BinaryElementWiseRhsPerChannelBatchN<<>>( + _BinaryElementWiseRhsPerChannelBatchN<<>>( lhs_data, rhs_data, fdm_H, @@ -259,7 +275,7 @@ void BinaryElementWiseImpl( N); } else { if (lhs_padded_strides && rhs_padded_strides && lhs_padded_strides->Size() && rhs_padded_strides->Size()) - _BinaryElementWise<<>>( + _BinaryElementWise<<>>( output_rank_or_simple_broadcast, *lhs_padded_strides, lhs_data, @@ -270,7 +286,7 @@ void BinaryElementWiseImpl( func, N); else if (lhs_padded_strides && lhs_padded_strides->Size()) - _BinaryElementWise<<>>( + _BinaryElementWise<<>>( output_rank_or_simple_broadcast, *lhs_padded_strides, lhs_data, @@ -281,7 +297,7 @@ void BinaryElementWiseImpl( func, N); else if (rhs_padded_strides && rhs_padded_strides->Size()) - _BinaryElementWise<<>>( + _BinaryElementWise<<>>( output_rank_or_simple_broadcast, TArray(), // lhs is not computed, so no need to deference lhs_padded_strides lhs_data, diff --git a/orttraining/orttraining/training_ops/cuda/activation/bias_gelu_grad_impl.cu b/orttraining/orttraining/training_ops/cuda/activation/bias_gelu_grad_impl.cu index 67c95872a9..03824c344b 100644 --- a/orttraining/orttraining/training_ops/cuda/activation/bias_gelu_grad_impl.cu +++ b/orttraining/orttraining/training_ops/cuda/activation/bias_gelu_grad_impl.cu @@ -62,9 +62,16 @@ void LaunchBiasGeluGradDxKernel( // given a 2D grid of blocks: // each grid row handles bias_size elements // there are input_size / bias_size rows - constexpr int num_elements_per_thread = GridDim::maxElementsPerThread; - const int num_threads_per_block = - std::min(static_cast(CeilDiv(bias_size, num_elements_per_thread)), static_cast(GridDim::maxThreadsPerBlock)); + + const int num_elements_per_thread = GridDim::maxElementsPerThread; + int max_threads_per_block = GridDim::maxThreadsPerBlock; + #ifdef USE_ROCM + // Optimization for ROCm MI100 + max_threads_per_block = 512; + #endif + + int num_threads_per_block = + std::min(static_cast(CeilDiv(bias_size, num_elements_per_thread)), static_cast(max_threads_per_block)); const auto grid_width = CeilDiv(bias_size, num_elements_per_thread * num_threads_per_block); const auto grid_height = input_size / bias_size;