diff --git a/aten/src/ATen/native/cuda/layer_norm_kernel.cu b/aten/src/ATen/native/cuda/layer_norm_kernel.cu index 59f5e4a041117..f20b9fdfafee2 100644 --- a/aten/src/ATen/native/cuda/layer_norm_kernel.cu +++ b/aten/src/ATen/native/cuda/layer_norm_kernel.cu @@ -1387,6 +1387,51 @@ void cuComputeGradGammaBeta( } } +// Two-pass gamma/beta backward: cuComputePartGradGammaBeta reduces +// dgamma/dbeta partial sums across part_size row-blocks, then +// cuComputeGradGammaBeta finishes the reduction across those partial sums. +template +void LaunchTwoPassGammaBetaBackwardCUDAKernel( + const T* dY_data, + const T* X_data, + const Tensor& X, + const Tensor& gamma, + const T_ACC* mean_data, + const T_ACC* rstd_data, + int64_t M, + int64_t N, + T* dgamma_data, + T* dbeta_data, + cudaStream_t cuda_stream) { + const int warp_size = at::cuda::warp_size(); + const int part_size = warp_size; + const dim3 threads2(warp_size, 4, 1); + const dim3 blocks2((N + threads2.x - 1) / threads2.x, part_size, 1); + const int nshared2_a = 2 * sizeof(T_ACC) * threads2.y * threads2.y * (threads2.x + 1); + const int nshared2_b = threads2.x * threads2.y * sizeof(T_ACC); + const int nshared2 = nshared2_a > nshared2_b ? nshared2_a : nshared2_b; + + const auto part_grad_dtype = at::toAccumulateType(X.scalar_type(), true); + Tensor part_grad_gamma = at::empty({part_size, N}, gamma.options().dtype(part_grad_dtype)); + Tensor part_grad_beta = at::native::empty_like(part_grad_gamma); + + cuComputePartGradGammaBeta<<>>( + dY_data, X_data, M, N, mean_data, rstd_data, + part_grad_gamma.template data_ptr(), + part_grad_beta.template data_ptr()); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + + const dim3 threads3(warp_size, 8, 1); // Optimization for ROCm + const dim3 blocks3((N + threads3.x - 1) / threads3.x, 1, 1); + const int nshared3 = threads3.x * threads3.y * sizeof(T_ACC); + + cuComputeGradGammaBeta<<>>( + part_grad_gamma.template data_ptr(), + part_grad_beta.template data_ptr(), + part_size, M, N, dgamma_data, dbeta_data); + C10_CUDA_KERNEL_LAUNCH_CHECK(); +} + template __global__ void cuComputeGradInput( const T* __restrict__ dout, @@ -1621,12 +1666,19 @@ void LayerNormBackwardKernelImplInternal( dbeta_data); C10_CUDA_KERNEL_LAUNCH_CHECK(); } else { - // Use the optimized tiled kernel adapted for wavefront-64. - // This replaces the legacy two-pass cuComputePartGradGammaBeta + - // cuComputeGradGammaBeta approach with a single-pass tiled reduction - // that has coalesced memory access and adaptive tile sizing. - LaunchGammaBetaBackwardCUDAKernel( - dY_data, X_data, mean_data, rstd_data, M, N, dgamma, dbeta, cuda_stream); + constexpr int block_dim_x = 64; // matches LaunchGammaBetaBackwardCUDAKernel's ROCm tile width + const int sm_count = at::cuda::getCurrentDeviceProperties()->multiProcessorCount; + // LaunchGammaBetaBackwardCUDAKernel already special-cases M >> N (huge M, small N) + // hich benchmarks ~2-6x faster than the two-pass kernel below for that regime. + const bool use_tiled_huge_M_kernel = + M > 64 * 1024 && N / block_dim_x < sm_count / 2; + if (use_tiled_huge_M_kernel) { + LaunchGammaBetaBackwardCUDAKernel( + dY_data, X_data, mean_data, rstd_data, M, N, dgamma, dbeta, cuda_stream); + } else { + LaunchTwoPassGammaBetaBackwardCUDAKernel( + dY_data, X_data, X, gamma, mean_data, rstd_data, M, N, dgamma_data, dbeta_data, cuda_stream); + } } #else LaunchGammaBetaBackwardCUDAKernel(