mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
Optimize BinaryElementWise and BiasGeluGrad kernels for AMD (#10594)
* Optimize elementwise and biasgelugrad kernels for AMD * Clean up for BiasGeluGradDxKernel
This commit is contained in:
parent
4c20f6863d
commit
fe8d867efa
2 changed files with 38 additions and 15 deletions
|
|
@ -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<int>(CeilDiv(count, GridDim::maxThreadsPerBlock * GridDim::maxElementsPerThread));
|
||||
int blocksPerGrid = static_cast<int>(CeilDiv(count, num_threads_per_block * num_elements_per_thread));
|
||||
CUDA_LONG N = static_cast<CUDA_LONG>(count);
|
||||
_BinaryElementWiseSimple<true, true, T, T1, T2, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0, stream>>>(
|
||||
_BinaryElementWiseSimple<true, true, T, T1, T2, FuncT, num_threads_per_block, num_elements_per_thread><<<blocksPerGrid, num_threads_per_block, 0, stream>>>(
|
||||
lhs_data,
|
||||
rhs_data,
|
||||
output_data,
|
||||
func,
|
||||
N);
|
||||
|
||||
}
|
||||
|
||||
template <typename T, typename T1, typename T2, typename FuncT>
|
||||
|
|
@ -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<int>(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<int>(CeilDiv(count, num_threads_per_block * num_elements_per_thread));
|
||||
CUDA_LONG N = static_cast<CUDA_LONG>(count);
|
||||
if (output_rank_or_simple_broadcast == static_cast<int32_t>(SimpleBroadcast::NoBroadcast)) {
|
||||
_BinaryElementWiseSimple<true, true, T, T1, T2, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0, stream>>>(
|
||||
_BinaryElementWiseSimple<true, true, T, T1, T2, FuncT, num_threads_per_block, num_elements_per_thread><<<blocksPerGrid, num_threads_per_block, 0, stream>>>(
|
||||
lhs_data,
|
||||
rhs_data,
|
||||
output_data,
|
||||
func,
|
||||
N);
|
||||
} else if (output_rank_or_simple_broadcast == static_cast<int32_t>(SimpleBroadcast::LeftScalar)) {
|
||||
_BinaryElementWiseSimple<false, true, T, T1, T2, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0, stream>>>(
|
||||
_BinaryElementWiseSimple<false, true, T, T1, T2, FuncT, num_threads_per_block, num_elements_per_thread><<<blocksPerGrid, num_threads_per_block, 0, stream>>>(
|
||||
lhs_data,
|
||||
rhs_data,
|
||||
output_data,
|
||||
func,
|
||||
N);
|
||||
} else if (output_rank_or_simple_broadcast == static_cast<int32_t>(SimpleBroadcast::RightScalar)) {
|
||||
_BinaryElementWiseSimple<true, false, T, T1, T2, FuncT, GridDim::maxThreadsPerBlock,
|
||||
GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0, stream>>>(
|
||||
_BinaryElementWiseSimple<true, false, T, T1, T2, FuncT, num_threads_per_block, num_elements_per_thread><<<blocksPerGrid, num_threads_per_block, 0, stream>>>(
|
||||
lhs_data,
|
||||
rhs_data,
|
||||
output_data,
|
||||
func,
|
||||
N);
|
||||
} else if (output_rank_or_simple_broadcast == static_cast<int32_t>(SimpleBroadcast::RightPerChannelBatch1)) {
|
||||
_BinaryElementWiseRhsPerChannelBatch1<T, T1, T2, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0, stream>>>(
|
||||
_BinaryElementWiseRhsPerChannelBatch1<T, T1, T2, FuncT, num_threads_per_block, num_elements_per_thread><<<blocksPerGrid, num_threads_per_block, 0, stream>>>(
|
||||
lhs_data,
|
||||
rhs_data,
|
||||
fdm_H,
|
||||
|
|
@ -249,7 +265,7 @@ void BinaryElementWiseImpl(
|
|||
func,
|
||||
N);
|
||||
} else if (output_rank_or_simple_broadcast == static_cast<int32_t>(SimpleBroadcast::RightPerChannelBatchN)) {
|
||||
_BinaryElementWiseRhsPerChannelBatchN<T, T1, T2, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0, stream>>>(
|
||||
_BinaryElementWiseRhsPerChannelBatchN<T, T1, T2, FuncT, num_threads_per_block, num_elements_per_thread><<<blocksPerGrid, num_threads_per_block, 0, stream>>>(
|
||||
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<T, T1, T2, FuncT, true, true, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0, stream>>>(
|
||||
_BinaryElementWise<T, T1, T2, FuncT, true, true, num_threads_per_block, num_elements_per_thread><<<blocksPerGrid, num_threads_per_block, 0, stream>>>(
|
||||
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<T, T1, T2, FuncT, true, false, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0, stream>>>(
|
||||
_BinaryElementWise<T, T1, T2, FuncT, true, false, num_threads_per_block, num_elements_per_thread><<<blocksPerGrid, num_threads_per_block, 0, stream>>>(
|
||||
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<T, T1, T2, FuncT, false, true, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0, stream>>>(
|
||||
_BinaryElementWise<T, T1, T2, FuncT, false, true, num_threads_per_block, num_elements_per_thread><<<blocksPerGrid, num_threads_per_block, 0, stream>>>(
|
||||
output_rank_or_simple_broadcast,
|
||||
TArray<int64_t>(), // lhs is not computed, so no need to deference lhs_padded_strides
|
||||
lhs_data,
|
||||
|
|
|
|||
|
|
@ -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<int>(static_cast<int>(CeilDiv(bias_size, num_elements_per_thread)), static_cast<int>(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<int>(static_cast<int>(CeilDiv(bias_size, num_elements_per_thread)), static_cast<int>(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;
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue