diff --git a/orttraining/orttraining/test/gradient/gradient_ops_test.cc b/orttraining/orttraining/test/gradient/gradient_ops_test.cc index 2f78fd10d9..9bc844d46c 100644 --- a/orttraining/orttraining/test/gradient/gradient_ops_test.cc +++ b/orttraining/orttraining/test/gradient/gradient_ops_test.cc @@ -1374,7 +1374,6 @@ TEST(GradientCheckerTest, FastGeluGrad) { void TestBiasGeluGrad(const std::string& op_type, const std::string& domain, int opset_version) { const TensorShape input_shape({2, 3, 4}); const TensorShape bias_shape({4}); - const float error_tolerance = 1e-3f; GradientChecker gradient_checker; OpDef op_def{op_type, domain, opset_version}; @@ -1383,7 +1382,7 @@ void TestBiasGeluGrad(const std::string& op_type, const std::string& domain, int ASSERT_STATUS_OK(gradient_checker.ComputeGradientError( op_def, {input_shape, bias_shape}, {input_shape}, &max_error)); - EXPECT_IS_TINIER_THAN(max_error, error_tolerance); + EXPECT_IS_TINY(max_error); } TEST(GradientCheckerTest, FastGeluGrad_Bias) { diff --git a/orttraining/orttraining/test/training_ops/cpu/activation/activation_op_test.cc b/orttraining/orttraining/test/training_ops/cpu/activation/activation_op_test.cc index 9e2a9c2aec..6ec19d8dc0 100644 --- a/orttraining/orttraining/test/training_ops/cpu/activation/activation_op_test.cc +++ b/orttraining/orttraining/test/training_ops/cpu/activation/activation_op_test.cc @@ -173,11 +173,13 @@ void TestBiasGeluGradBroadcastBias(const std::string& op, int opset_version, con TEST(BiasGeluGradDxTest, BroadcastBias) { TestBiasGeluGradBroadcastBias("BiasGeluGrad_dX", 1, kMSDomain, {2, 3, 4, 5}, GeluGrad); TestBiasGeluGradBroadcastBias("BiasGeluGrad_dX", 1, kMSDomain, {2, 4, 3072}, GeluGrad); + TestBiasGeluGradBroadcastBias("BiasGeluGrad_dX", 1, kMSDomain, {2, 16384}, GeluGrad); } TEST(BiasFastGeluGradDxTest, BroadcastBias) { TestBiasGeluGradBroadcastBias("BiasFastGeluGrad_dX", 1, kMSDomain, {2, 3, 4, 5}, GeluApproximationGrad); TestBiasGeluGradBroadcastBias("BiasFastGeluGrad_dX", 1, kMSDomain, {2, 4, 3072}, GeluApproximationGrad); + TestBiasGeluGradBroadcastBias("BiasFastGeluGrad_dX", 1, kMSDomain, {2, 16384}, GeluApproximationGrad); } } // namespace test 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 bffe0b5ca5..da5d621f1d 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 @@ -11,20 +11,44 @@ namespace onnxruntime { namespace cuda { -template +template __global__ void BiasGeluGradDxKernel(int64_t bias_size, const T* dY, const T* X, const T* B, T* dX) { - const int64_t input_base_idx = bias_size * blockIdx.x + num_consecutive_elements_per_group * threadIdx.x; - const int64_t bias_base_idx = num_consecutive_elements_per_group * threadIdx.x; - const int64_t group_stride = num_consecutive_elements_per_group * blockDim.x; + const auto num_elements_per_block = num_elements_per_thread * blockDim.x; + const auto input_base_idx = bias_size * blockIdx.y + num_elements_per_block * blockIdx.x + threadIdx.x; + const auto bias_base_idx = num_elements_per_block * blockIdx.x + threadIdx.x; + const auto element_stride = blockDim.x; + T reg_dY[num_elements_per_thread]; + T reg_X[num_elements_per_thread]; + T reg_B[num_elements_per_thread]; + + { + auto input_idx = input_base_idx; + auto bias_idx = bias_base_idx; #pragma unroll - for (int group_idx = 0; group_idx < num_groups_per_thread; ++group_idx) { -#pragma unroll - for (int element_idx = 0; element_idx < num_consecutive_elements_per_group; ++element_idx) { - const auto offset = group_stride * group_idx + element_idx; - const auto input_idx = input_base_idx + offset, bias_idx = bias_base_idx + offset; + for (int element_idx = 0; element_idx < num_elements_per_thread; ++element_idx) { if (bias_idx < bias_size) { - dX[input_idx] = ComputeGeluGradScalar(dY[input_idx], X[input_idx] + B[bias_idx], GeluComputationMode{}); + reg_dY[element_idx] = dY[input_idx]; + reg_X[element_idx] = X[input_idx]; + reg_B[element_idx] = B[bias_idx]; + + input_idx += element_stride; + bias_idx += element_stride; + } + } + } + + { + auto input_idx = input_base_idx; + auto bias_idx = bias_base_idx; +#pragma unroll + for (int element_idx = 0; element_idx < num_elements_per_thread; ++element_idx) { + if (bias_idx < bias_size) { + dX[input_idx] = ComputeGeluGradScalar( + reg_dY[element_idx], reg_X[element_idx] + reg_B[element_idx], GeluComputationMode{}); + + input_idx += element_stride; + bias_idx += element_stride; } } } @@ -34,16 +58,19 @@ template void LaunchBiasGeluGradDxKernel( int64_t input_size, int64_t bias_size, const T* dY, const T* X, const T* B, T* dX) { - // each block handles bias_size elements - // there are input_size / bias_size blocks - constexpr int num_consecutive_elements_per_group = 4; - constexpr int num_groups_per_thread = 4; + // 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 auto num_threads_per_block = + std::min(CeilDiv(bias_size, num_elements_per_thread), static_cast(GridDim::maxThreadsPerBlock)); + const auto grid_width = CeilDiv(bias_size, num_elements_per_thread * num_threads_per_block); + const auto grid_height = input_size / bias_size; - const auto num_threads_per_block = CeilDiv(bias_size, num_consecutive_elements_per_group * num_groups_per_thread); - const auto num_blocks_per_grid = input_size / bias_size; + const dim3 grid_dim{static_cast(grid_width), static_cast(grid_height)}; - BiasGeluGradDxKernel - <<>>(bias_size, dY, X, B, dX); + BiasGeluGradDxKernel + <<>>(bias_size, dY, X, B, dX); } // explicit instantiations