From 20c47ce91c0fe6ce01ad37f4b8386a19a699753b Mon Sep 17 00:00:00 2001 From: jingyanwangms <47403504+jingyanwangms@users.noreply.github.com> Date: Mon, 12 Oct 2020 11:22:12 -0700 Subject: [PATCH] Simplified layer norm changes (#5028) * t5 layer norm changes * add t5 layer norm kernel * use template for t5 layer norm * template definition changes * no build error * add CPU cuda kernel * first unit test * other forward unit tests * add T5LayerNormGrad * Add c++ transform and test for T5 LN * fix and some debug prints * fix cuda error * rename from t5 to simplified * PR comments * revert change on invertible LM code path * remove duplicate forward computation * add GradientCheckerTest.SimplifiedLayerNormGrad * change back macro * Fix SimplifiedLayerNorm Gradient * merge with Sherlockss changes * changed cuda kernel * reapply cpu kernel changes Co-authored-by: Jingyan Wang Co-authored-by: aishwarya bhandare Co-authored-by: Sherlock Huang --- .../contrib_ops/cpu/cpu_contrib_kernels.cc | 4 + onnxruntime/contrib_ops/cpu/layer_norm.cc | 63 ++- onnxruntime/contrib_ops/cpu/layer_norm.h | 2 +- .../contrib_ops/cuda/cuda_contrib_kernels.cc | 6 + onnxruntime/contrib_ops/cuda/layer_norm.cc | 40 +- onnxruntime/contrib_ops/cuda/layer_norm.h | 2 +- .../contrib_ops/cuda/layer_norm_impl.cu | 81 ++-- .../contrib_ops/cuda/layer_norm_impl.h | 2 +- .../core/graph/contrib_ops/contrib_defs.cc | 48 ++ .../core/optimizer/graph_transformer_utils.cc | 1 + .../core/optimizer/layer_norm_fusion.cc | 171 +++++++ .../core/optimizer/layer_norm_fusion.h | 17 + .../test/contrib_ops/layer_norm_test.cc | 44 +- .../test/optimizer/graph_transform_test.cc | 33 ++ .../transform/fusion/layer_norm_t5.onnx | Bin 0 -> 419 bytes .../transform/fusion/layer_norm_t5_gen.py | 48 ++ .../core/graph/gradient_builder.cc | 8 + .../orttraining/core/graph/gradient_builder.h | 1 + .../core/graph/gradient_builder_registry.cc | 1 + .../core/graph/training_op_defs.cc | 24 + .../core/optimizer/graph_transformer_utils.cc | 1 + .../test/gradient/gradient_ops_test.cc | 17 + .../test/training_ops/cuda/layer_norm_test.cc | 50 +- .../training_ops/cpu/cpu_training_kernels.cc | 4 + .../training_ops/cpu/nn/layer_norm.cc | 77 +-- .../training_ops/cpu/nn/layer_norm.h | 2 +- .../cuda/cuda_training_kernels.cc | 6 + .../training_ops/cuda/nn/layer_norm.cc | 51 +- .../training_ops/cuda/nn/layer_norm.h | 2 +- .../training_ops/cuda/nn/layer_norm_impl.cu | 440 +++--------------- .../training_ops/cuda/nn/layer_norm_impl.h | 2 +- 31 files changed, 742 insertions(+), 506 deletions(-) create mode 100644 onnxruntime/test/testdata/transform/fusion/layer_norm_t5.onnx create mode 100644 onnxruntime/test/testdata/transform/fusion/layer_norm_t5_gen.py diff --git a/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc b/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc index 4de159d651..275cd2b6ec 100644 --- a/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc +++ b/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc @@ -76,6 +76,8 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSNchwcDomai class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSNchwcDomain, 1, float, Upsample); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, float, LayerNormalization); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, double, LayerNormalization); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, float, SimplifiedLayerNormalization); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, double, SimplifiedLayerNormalization); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, SkipLayerNormalization); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, double, SkipLayerNormalization); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, Inverse); @@ -184,6 +186,8 @@ Status RegisterCpuContribKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/onnxruntime/contrib_ops/cpu/layer_norm.cc b/onnxruntime/contrib_ops/cpu/layer_norm.cc index 65998deb84..c3c26e4684 100644 --- a/onnxruntime/contrib_ops/cpu/layer_norm.cc +++ b/onnxruntime/contrib_ops/cpu/layer_norm.cc @@ -21,27 +21,36 @@ namespace contrib { kCpuExecutionProvider, \ KernelDefBuilder() \ .TypeConstraint("T", DataTypeImpl::GetTensorType()), \ - LayerNorm); - + LayerNorm); \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + SimplifiedLayerNormalization, \ + kOnnxDomain, \ + 1, \ + T, \ + kCpuExecutionProvider, \ + KernelDefBuilder() \ + .TypeConstraint("T", DataTypeImpl::GetTensorType()), \ + LayerNorm); + REGISTER_KERNEL_TYPED(float) REGISTER_KERNEL_TYPED(double) -template -LayerNorm::LayerNorm(const OpKernelInfo& op_kernel_info) +template +LayerNorm::LayerNorm(const OpKernelInfo& op_kernel_info) : OpKernel(op_kernel_info) { ORT_ENFORCE(op_kernel_info.GetAttr("axis", &axis_).IsOK()); ORT_ENFORCE(op_kernel_info.GetAttr("epsilon", &epsilon_).IsOK()); } -template -Status LayerNorm::Compute(OpKernelContext* p_ctx) const { +template +Status LayerNorm::Compute(OpKernelContext* p_ctx) const { // Inputs const Tensor* X = p_ctx->Input(0); const Tensor* scale = p_ctx->Input(1); const Tensor* bias = p_ctx->Input(2); auto X_data = X->template Data(); auto scale_data = scale->template Data(); - auto bias_data = bias->template Data(); + auto bias_data = simplified ? nullptr : bias->template Data(); const TensorShape& x_shape = X->Shape(); const int64_t axis = HandleNegativeAxis(axis_, x_shape.NumDimensions()); @@ -67,19 +76,23 @@ Status LayerNorm::Compute(OpKernelContext* p_ctx) const { T* mean_data = nullptr; BufferUniquePtr mean_data_buf_ptr; - Tensor* mean = p_ctx->Output(1, TensorShape(mean_inv_std_var_dim)); - if (mean != nullptr) { - mean_data = mean->template MutableData(); - } else { - auto mean_data_buf = alloc->Alloc(SafeInt(sizeof(T)) * norm_count); - mean_data_buf_ptr = BufferUniquePtr(mean_data_buf, BufferDeleter(alloc)); - mean_data = static_cast(mean_data_buf_ptr.get()); + int output_index = 1; + + if (!simplified) { + Tensor* mean = p_ctx->Output(output_index++, TensorShape(mean_inv_std_var_dim)); + if (mean != nullptr) { + mean_data = mean->template MutableData(); + } else { + auto mean_data_buf = alloc->Alloc(SafeInt(sizeof(T)) * norm_count); + mean_data_buf_ptr = BufferUniquePtr(mean_data_buf, BufferDeleter(alloc)); + mean_data = static_cast(mean_data_buf_ptr.get()); + } } T* inv_std_var_data = nullptr; BufferUniquePtr inv_std_var_data_buf_ptr; - Tensor* inv_std_var = p_ctx->Output(2, TensorShape(mean_inv_std_var_dim)); + Tensor* inv_std_var = p_ctx->Output(output_index, TensorShape(mean_inv_std_var_dim)); if (inv_std_var != nullptr) { inv_std_var_data = inv_std_var->template MutableData(); } else { @@ -102,13 +115,23 @@ Status LayerNorm::Compute(OpKernelContext* p_ctx) const { } mean = mean / norm_size; - mean_square = sqrt(mean_square / norm_size - mean * mean + epsilon_); - - for (int64_t h = 0; h < norm_size; h++) { - p_output[h] = (p_input[h] - mean) / mean_square * scale_data[h] + bias_data[h]; + if (simplified) { + mean_square = sqrt(mean_square / norm_size + epsilon_); + } else { + mean_square = sqrt(mean_square / norm_size - mean * mean + epsilon_); } - mean_data[task_idx] = mean; + for (int64_t h = 0; h < norm_size; h++) { + if (simplified) { + p_output[h] = p_input[h] / mean_square * scale_data[h]; + } else { + p_output[h] = (p_input[h] - mean) / mean_square * scale_data[h] + bias_data[h]; + } + } + + if (mean_data != nullptr) { + mean_data[task_idx] = mean; + } inv_std_var_data[task_idx] = 1 / mean_square; }, 0); diff --git a/onnxruntime/contrib_ops/cpu/layer_norm.h b/onnxruntime/contrib_ops/cpu/layer_norm.h index 89fa853611..d42669e150 100644 --- a/onnxruntime/contrib_ops/cpu/layer_norm.h +++ b/onnxruntime/contrib_ops/cpu/layer_norm.h @@ -10,7 +10,7 @@ namespace onnxruntime { namespace contrib { -template +template class LayerNorm final : public OpKernel { public: LayerNorm(const OpKernelInfo& op_kernel_info); diff --git a/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc b/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc index 773ea9ec7e..6dc5fa4f1d 100644 --- a/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc +++ b/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc @@ -67,6 +67,9 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, float_float, LayerNormalization); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, double_double, LayerNormalization); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, MLFloat16_float, LayerNormalization); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, float_float, SimplifiedLayerNormalization); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, double_double, SimplifiedLayerNormalization); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, MLFloat16_float, SimplifiedLayerNormalization); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, Inverse); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, int8_t_MLFloat16, QuantizeLinear); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, uint8_t_MLFloat16, QuantizeLinear); @@ -141,6 +144,9 @@ Status RegisterCudaContribKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/onnxruntime/contrib_ops/cuda/layer_norm.cc b/onnxruntime/contrib_ops/cuda/layer_norm.cc index de89bd424d..685fd20da8 100644 --- a/onnxruntime/contrib_ops/cuda/layer_norm.cc +++ b/onnxruntime/contrib_ops/cuda/layer_norm.cc @@ -21,22 +21,32 @@ namespace cuda { KernelDefBuilder() \ .TypeConstraint("T", DataTypeImpl::GetTensorType()) \ .TypeConstraint("U", DataTypeImpl::GetTensorType()), \ - LayerNorm); + LayerNorm); \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + SimplifiedLayerNormalization, \ + kOnnxDomain, \ + 1, \ + T##_##U, \ + kCudaExecutionProvider, \ + KernelDefBuilder() \ + .TypeConstraint("T", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("U", DataTypeImpl::GetTensorType()), \ + LayerNorm); REGISTER_KERNEL_TYPED(float, float) REGISTER_KERNEL_TYPED(double, double) REGISTER_KERNEL_TYPED(MLFloat16, float) -template -LayerNorm::LayerNorm(const OpKernelInfo& op_kernel_info) : CudaKernel(op_kernel_info) { +template +LayerNorm::LayerNorm(const OpKernelInfo& op_kernel_info) : CudaKernel(op_kernel_info) { ORT_ENFORCE(op_kernel_info.GetAttr("axis", &axis_).IsOK()); float tmp_epsilon; ORT_ENFORCE(op_kernel_info.GetAttr("epsilon", &tmp_epsilon).IsOK()); epsilon_ = tmp_epsilon; } -template -Status LayerNorm::ComputeInternal(OpKernelContext* ctx) const { +template +Status LayerNorm::ComputeInternal(OpKernelContext* ctx) const { typedef typename ToCudaType::MappedType CudaT; typedef typename ToCudaType::MappedType CudaU; //Inputs @@ -46,7 +56,7 @@ Status LayerNorm::ComputeInternal(OpKernelContext* ctx) const { auto X_data = reinterpret_cast(X->template Data()); auto scale_data = reinterpret_cast(scale->template Data()); - auto bias_data = reinterpret_cast(bias->template Data()); + auto bias_data = simplified ? nullptr: reinterpret_cast(bias->template Data()); const TensorShape& x_shape = X->Shape(); const int64_t axis = HandleNegativeAxis(axis_, x_shape.NumDimensions()); @@ -69,19 +79,23 @@ Status LayerNorm::ComputeInternal(OpKernelContext* ctx) const { mean_inv_std_var_dim.emplace_back(1); } } - Tensor* mean = ctx->Output(1, TensorShape(mean_inv_std_var_dim)); - Tensor* var = ctx->Output(2, TensorShape(mean_inv_std_var_dim)); - CudaU* mean_data = nullptr; - if (mean != nullptr) { - mean_data = reinterpret_cast(mean->template MutableData()); - } + int output_index = 1; + CudaU* mean_data = nullptr; + if (!simplified) { + Tensor* mean = ctx->Output(output_index++, TensorShape(mean_inv_std_var_dim)); + if (mean != nullptr) { + mean_data = reinterpret_cast(mean->template MutableData()); + } + } + + Tensor* var = ctx->Output(output_index, TensorShape(mean_inv_std_var_dim)); CudaU* inv_var_data = nullptr; if (var != nullptr) { inv_var_data = reinterpret_cast(var->template MutableData()); } - HostApplyLayerNorm(GetDeviceProp(), Y_data, mean_data, inv_var_data, X_data, n1, n2, epsilon_, scale_data, bias_data); + HostApplyLayerNorm(GetDeviceProp(), Y_data, mean_data, inv_var_data, X_data, n1, n2, epsilon_, scale_data, bias_data); return Status::OK(); } diff --git a/onnxruntime/contrib_ops/cuda/layer_norm.h b/onnxruntime/contrib_ops/cuda/layer_norm.h index e7bbbc3ca5..efcf23e89b 100644 --- a/onnxruntime/contrib_ops/cuda/layer_norm.h +++ b/onnxruntime/contrib_ops/cuda/layer_norm.h @@ -12,7 +12,7 @@ namespace cuda { using namespace onnxruntime::cuda; -template +template class LayerNorm final : public CudaKernel { public: LayerNorm(const OpKernelInfo& op_kernel_info); diff --git a/onnxruntime/contrib_ops/cuda/layer_norm_impl.cu b/onnxruntime/contrib_ops/cuda/layer_norm_impl.cu index 4251c21033..30c2f6386c 100644 --- a/onnxruntime/contrib_ops/cuda/layer_norm_impl.cu +++ b/onnxruntime/contrib_ops/cuda/layer_norm_impl.cu @@ -32,7 +32,7 @@ namespace cuda { using namespace onnxruntime::cuda; -template +template __device__ void cuWelfordOnlineSum( const U curr, U& mu, @@ -42,11 +42,15 @@ __device__ void cuWelfordOnlineSum( U delta = curr - mu; U lmean = mu + delta / count; mu = lmean; - U delta2 = curr - lmean; - sigma2 = sigma2 + delta * delta2; + if (simplified) { + sigma2 = sigma2 + curr * curr; + } else { + U delta2 = curr - lmean; + sigma2 = sigma2 + delta * delta2; + } } -template +template __device__ void cuChanOnlineSum( const U muB, const U sigma2B, @@ -63,14 +67,18 @@ __device__ void cuChanOnlineSum( nA = nA / nX; nB = nB / nX; mu = nA * mu + nB * muB; - sigma2 = sigma2 + sigma2B + delta * delta * nA * nB * nX; + if (simplified) { + sigma2 = sigma2 + sigma2B; + } else { + sigma2 = sigma2 + sigma2B + delta * delta * nA * nB * nX; + } } else { mu = U(0); sigma2 = U(0); } } -template +template __device__ void cuWelfordMuSigma2( const T* __restrict__ vals, const int n1, @@ -99,12 +107,12 @@ __device__ void cuWelfordMuSigma2( for (; l + 3 < n2; l += 4 * numx) { for (int k = 0; k < 4; ++k) { U curr = static_cast(lvals[l + k]); - cuWelfordOnlineSum(curr, mu, sigma2, count); + cuWelfordOnlineSum(curr, mu, sigma2, count); } } for (; l < n2; ++l) { U curr = static_cast(lvals[l]); - cuWelfordOnlineSum(curr, mu, sigma2, count); + cuWelfordOnlineSum(curr, mu, sigma2, count); } // intra-warp reductions #pragma unroll @@ -112,7 +120,7 @@ __device__ void cuWelfordMuSigma2( U muB = WARP_SHFL_DOWN(mu, stride); U countB = WARP_SHFL_DOWN(count, stride); U sigma2B = WARP_SHFL_DOWN(sigma2, stride); - cuChanOnlineSum(muB, sigma2B, countB, mu, sigma2, count); + cuChanOnlineSum(muB, sigma2B, countB, mu, sigma2, count); } // threadIdx.x == 0 has correct values for each warp @@ -134,7 +142,7 @@ __device__ void cuWelfordMuSigma2( U muB = ubuf[2 * threadIdx.y]; U sigma2B = ubuf[2 * threadIdx.y + 1]; U countB = ibuf[threadIdx.y]; - cuChanOnlineSum(muB, sigma2B, countB, mu, sigma2, count); + cuChanOnlineSum(muB, sigma2B, countB, mu, sigma2, count); } __syncthreads(); } @@ -154,7 +162,7 @@ __device__ void cuWelfordMuSigma2( } } -template <> +template __device__ void cuWelfordMuSigma2( const half* __restrict__ vals, const int n1, @@ -185,7 +193,7 @@ __device__ void cuWelfordMuSigma2( // first thread consumes first point if (thrx == 0) { float curr = static_cast(lvals[0]); - cuWelfordOnlineSum(curr, mu, sigma2, count); + cuWelfordOnlineSum(curr, mu, sigma2, count); } ++l; } @@ -193,13 +201,13 @@ __device__ void cuWelfordMuSigma2( for (; l + 7 < n2; l += 8 * numx) { for (int k = 0; k < 8; k += 2) { float2 curr = __half22float2(*((__half2*)(lvals + l + k))); - cuWelfordOnlineSum(static_cast(curr.x), mu, sigma2, count); - cuWelfordOnlineSum(static_cast(curr.y), mu, sigma2, count); + cuWelfordOnlineSum(static_cast(curr.x), mu, sigma2, count); + cuWelfordOnlineSum(static_cast(curr.y), mu, sigma2, count); } } for (; l < n2; ++l) { float curr = static_cast(lvals[l]); - cuWelfordOnlineSum(curr, mu, sigma2, count); + cuWelfordOnlineSum(curr, mu, sigma2, count); } // intra-warp reductions #pragma unroll @@ -207,7 +215,7 @@ __device__ void cuWelfordMuSigma2( float muB = WARP_SHFL_DOWN(mu, stride); float countB = WARP_SHFL_DOWN(count, stride); float sigma2B = WARP_SHFL_DOWN(sigma2, stride); - cuChanOnlineSum(muB, sigma2B, countB, mu, sigma2, count); + cuChanOnlineSum(muB, sigma2B, countB, mu, sigma2, count); } // threadIdx.x == 0 has correct values for each warp @@ -229,7 +237,7 @@ __device__ void cuWelfordMuSigma2( float muB = ubuf[2 * threadIdx.y]; float sigma2B = ubuf[2 * threadIdx.y + 1]; float countB = ibuf[threadIdx.y]; - cuChanOnlineSum(muB, sigma2B, countB, mu, sigma2, count); + cuChanOnlineSum(muB, sigma2B, countB, mu, sigma2, count); } __syncthreads(); } @@ -297,7 +305,7 @@ struct SharedMemory { }; } // namespace -template +template __global__ void cuApplyLayerNorm( T* __restrict__ output_vals, U* __restrict__ mean, @@ -316,21 +324,20 @@ __global__ void cuApplyLayerNorm( SharedMemory shared; U* buf = shared.getPointer(); U mu, sigma2; - cuWelfordMuSigma2(vals, n1, n2, i1, mu, sigma2, buf); + cuWelfordMuSigma2(vals, n1, n2, i1, mu, sigma2, buf); const T* lvals = vals + i1 * n2; T* ovals = output_vals + i1 * n2; U c_invvar = rsqrt(sigma2 + epsilon); const int numx = blockDim.x * blockDim.y; const int thrx = threadIdx.x + threadIdx.y * blockDim.x; - if (gamma != NULL && beta != NULL) { - for (int i = thrx; i < n2; i += numx) { - U curr = static_cast(lvals[i]); - ovals[i] = gamma[i] * static_cast(c_invvar * (curr - mu)) + beta[i]; - } - } else { - for (int i = thrx; i < n2; i += numx) { - U curr = static_cast(lvals[i]); - ovals[i] = static_cast(c_invvar * (curr - mu)); + for (int i = thrx; i < n2; i += numx) { + U curr = static_cast(lvals[i]); + T gamma_i = (gamma != NULL) ? gamma[i]: (T)1; + T beta_i = (beta != NULL) ? beta[i] : (T) 0; + if (simplified) { + ovals[i] = gamma_i * static_cast(c_invvar * curr); + } else { + ovals[i] = gamma_i * static_cast(c_invvar * (curr - mu)) + beta_i; } } if (threadIdx.x == 0 && threadIdx.y == 0) { @@ -340,7 +347,7 @@ __global__ void cuApplyLayerNorm( } } -template +template void HostApplyLayerNorm( const cudaDeviceProp& prop, T* output, @@ -360,7 +367,7 @@ void HostApplyLayerNorm( const dim3 blocks(1, std::min((uint64_t)n1, maxGridY), 1); int nshared = threads.y > 1 ? threads.y * sizeof(U) + (threads.y / 2) * sizeof(U) : 0; - cuApplyLayerNorm<<>>( + cuApplyLayerNorm<<>>( output, mean, invvar, @@ -370,13 +377,17 @@ void HostApplyLayerNorm( gamma, beta); } -#define LAYERNORM_LINEAR_IMPL(T, U) \ - template void HostApplyLayerNorm(const cudaDeviceProp& prop, T* output, U* mean, U* invvar, const T* input, int64_t n1, int64_t n2, \ +#define LAYERNORM_LINEAR_IMPL(T, U, simplified) \ + template void HostApplyLayerNorm(const cudaDeviceProp& prop, T* output, U* mean, U* invvar, const T* input, int64_t n1, int64_t n2, \ double epsilon, const T* gamma, const T* beta); -LAYERNORM_LINEAR_IMPL(float, float) -LAYERNORM_LINEAR_IMPL(half, float) -LAYERNORM_LINEAR_IMPL(double, double) +LAYERNORM_LINEAR_IMPL(float, float, true) +LAYERNORM_LINEAR_IMPL(half, float, true) +LAYERNORM_LINEAR_IMPL(double, double, true) +LAYERNORM_LINEAR_IMPL(float, float, false) +LAYERNORM_LINEAR_IMPL(half, float, false) +LAYERNORM_LINEAR_IMPL(double, double, false) + //LAYERNORM_LINEAR_IMPL(half, half) } // namespace cuda diff --git a/onnxruntime/contrib_ops/cuda/layer_norm_impl.h b/onnxruntime/contrib_ops/cuda/layer_norm_impl.h index 70bc24eba2..82c169143d 100644 --- a/onnxruntime/contrib_ops/cuda/layer_norm_impl.h +++ b/onnxruntime/contrib_ops/cuda/layer_norm_impl.h @@ -29,7 +29,7 @@ namespace onnxruntime { namespace contrib { namespace cuda { -template +template void HostApplyLayerNorm( const cudaDeviceProp& prop, T* output, diff --git a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc index 5010a057dd..5635965995 100644 --- a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc @@ -2854,6 +2854,54 @@ Example 4: } }); + ONNX_CONTRIB_OPERATOR_SCHEMA(SimplifiedLayerNormalization) + .SetDomain(kOnnxDomain) + .SinceVersion(1) + .SetSupportLevel(OpSchema::SupportType::EXPERIMENTAL) + .SetDoc("SimplifiedLayerNormalization") + .Attr("axis", + "The first normalization dimension: normalization will be performed along dimensions axis : rank(inputs).", + AttributeProto::INT, static_cast(-1)) + .Attr("epsilon", + "The epsilon value to use to avoid division by zero.", + AttributeProto::FLOAT, 1e-5f) + .AllowUncheckedAttributes() + .Input(0, "X", "Input data tensor from the previous layer.", "T") + .Input(1, "scale", "Scale tensor.", "T") + .Output(0, "Y", "Output data tensor.", "T") + .Output(1, "inv_std_var", "Saved inverse standard variance used during training to speed up gradient computation.", "U", OpSchema::Optional) + .TypeConstraint( + "T", + {"tensor(float16)", "tensor(float)", "tensor(double)", "tensor(bfloat16)"}, + "Constrain input and output types (except mean and inv_std_var) to float tensors.") + .TypeConstraint( + "U", + {"tensor(float)"}, + "Constrain mean and inv_std_var to be float tensors.") + .TypeAndShapeInferenceFunction([](ONNX_NAMESPACE::InferenceContext& ctx) { + propagateShapeAndTypeFromFirstInput(ctx); + propagateElemTypeFromInputToOutput(ctx, 0, 0); + if (!hasNInputShapes(ctx, 1)) { + return; + } + auto& input_shape = ctx.getInputType(0)->tensor_type().shape(); + int64_t input_ndim = input_shape.dim_size(); + int64_t axis = -1; + auto axis_proto = ctx.getAttribute("axis"); + if (axis_proto) { + axis = axis_proto->i(); + } + if (axis < 0) { + axis += input_ndim; + } + + if (ctx.getNumOutputs() > 1) { + auto saved_inv_std_var_shape = ctx.getOutputType(1)->mutable_tensor_type()->mutable_shape(); + saved_inv_std_var_shape->CopyFrom(input_shape); + saved_inv_std_var_shape->mutable_dim(static_cast(axis))->set_dim_value(1); + } + }); + // Register the NCHWc schemas if supported by the platform. if (MlasNchwcGetBlockSize() > 1) { RegisterNchwcSchemas(); diff --git a/onnxruntime/core/optimizer/graph_transformer_utils.cc b/onnxruntime/core/optimizer/graph_transformer_utils.cc index 96c9e2aa46..2963caaf51 100644 --- a/onnxruntime/core/optimizer/graph_transformer_utils.cc +++ b/onnxruntime/core/optimizer/graph_transformer_utils.cc @@ -143,6 +143,7 @@ std::vector> GenerateTransformers(TransformerL std::unordered_set cpu_cuda_execution_providers = {onnxruntime::kCpuExecutionProvider, onnxruntime::kCudaExecutionProvider}; transformers.emplace_back(onnxruntime::make_unique(cpu_cuda_execution_providers)); transformers.emplace_back(onnxruntime::make_unique(cpu_cuda_execution_providers)); + transformers.emplace_back(onnxruntime::make_unique(cpu_cuda_execution_providers)); transformers.emplace_back(onnxruntime::make_unique(cpu_cuda_execution_providers)); transformers.emplace_back(onnxruntime::make_unique(cpu_cuda_execution_providers)); diff --git a/onnxruntime/core/optimizer/layer_norm_fusion.cc b/onnxruntime/core/optimizer/layer_norm_fusion.cc index 4a3ed60e10..c3e53531dc 100644 --- a/onnxruntime/core/optimizer/layer_norm_fusion.cc +++ b/onnxruntime/core/optimizer/layer_norm_fusion.cc @@ -329,4 +329,175 @@ Status LayerNormFusion::ApplyImpl(Graph& graph, bool& modified, int graph_level, } return Status::OK(); } + +/** +Layer Normalization will fuse LayerNormalization into one node : + +X --> Pow --> ReduceMean --> Add --> Sqrt --> Div --> Mul +| ^ +| | ++----------------------------------------------+ +*/ +Status SimplifiedLayerNormFusion::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(); + std::vector> nodes_to_remove; + for (auto node_index : node_topology_list) { + nodes_to_remove.clear(); + auto* p_pow = graph.GetNode(node_index); + if (p_pow == nullptr) + continue; // we removed the node as part of an earlier fusion + + Node& pow_node = *p_pow; + ORT_RETURN_IF_ERROR(Recurse(pow_node, modified, graph_level, logger)); + + if (!graph_utils::IsSupportedOptypeVersionAndDomain(pow_node, "Pow", {7, 12, 13}) || + !graph_utils::IsSupportedProvider(pow_node, GetCompatibleExecutionProviders()) || + !optimizer_utils::CheckOutputEdges(graph, pow_node, 1) || + !graph.GetNodeOutputsInGraphOutputs(pow_node).empty() || + !IsSupportedDataType(pow_node)) { + continue; + } + nodes_to_remove.push_back(pow_node); + + const Node* p_reduce_mean = nullptr; + + p_reduce_mean = graph_utils::FirstChildByType(pow_node, "ReduceMean"); + if (p_reduce_mean == nullptr) { + continue; + } + Node& reduce_mean_node = *graph.GetNode(p_reduce_mean->Index()); + if (!graph_utils::IsSupportedOptypeVersionAndDomain(reduce_mean_node, "ReduceMean", {1, 11, 13}) || + reduce_mean_node.GetExecutionProviderType() != pow_node.GetExecutionProviderType() || + !optimizer_utils::CheckOutputEdges(graph, reduce_mean_node, 1) || + !IsSupportedDataType(reduce_mean_node) || + reduce_mean_node.GetInputEdgesCount() == 0) { + continue; + } + nodes_to_remove.push_back(reduce_mean_node); + + const Node* p_add = graph_utils::FirstChildByType(reduce_mean_node, "Add"); + if (p_add == nullptr) { + continue; + } + Node& add_node = *graph.GetNode(p_add->Index()); + if (!graph_utils::IsSupportedOptypeVersionAndDomain(add_node, "Add", {7, 13}) || + add_node.GetExecutionProviderType() != pow_node.GetExecutionProviderType() || + !optimizer_utils::CheckOutputEdges(graph, add_node, 1) || + !IsSupportedDataType(add_node)) { + continue; + } + nodes_to_remove.push_back(add_node); + + const Node* p_sqrt = graph_utils::FirstChildByType(add_node, "Sqrt"); + if (p_sqrt == nullptr) { + continue; + } + Node& sqrt_node = *graph.GetNode(p_sqrt->Index()); + + if (!graph_utils::IsSupportedOptypeVersionAndDomain(sqrt_node, "Sqrt", {6, 13}) || + sqrt_node.GetExecutionProviderType() != pow_node.GetExecutionProviderType() || + !optimizer_utils::CheckOutputEdges(graph, sqrt_node, 1) || + !IsSupportedDataType(sqrt_node) || + sqrt_node.GetInputEdgesCount() == 0) { + continue; + } + nodes_to_remove.push_back(sqrt_node); + + const Node* p_div = graph_utils::FirstChildByType(sqrt_node, "Div"); + if (p_div == nullptr) { + continue; + } + Node& div_node = *graph.GetNode(p_div->Index()); + if (!graph_utils::IsSupportedOptypeVersionAndDomain(div_node, "Div", {7, 13}) || + div_node.GetExecutionProviderType() != pow_node.GetExecutionProviderType() || + !optimizer_utils::CheckOutputEdges(graph, div_node, 1) || + !IsSupportedDataType(div_node)) { + continue; + } + nodes_to_remove.push_back(div_node); + + const NodeArg* p_div_input = div_node.MutableInputDefs()[0]; + const NodeArg* p_pow_input = pow_node.MutableInputDefs()[0]; + + if (p_div_input == nullptr || p_pow_input == nullptr || + p_div_input != p_pow_input) { + continue; + } + + // div --> mul + Node& mul_node = *graph.GetNode(div_node.OutputNodesBegin()->Index()); + if (!graph_utils::IsSupportedOptypeVersionAndDomain(mul_node, "Mul", {7, 13}) || + mul_node.GetExecutionProviderType() != pow_node.GetExecutionProviderType() || + !IsSupportedDataType(mul_node)) { + continue; + } + nodes_to_remove.push_back(mul_node); + + // get axes attributes + const onnxruntime::NodeAttributes& attributes = reduce_mean_node.GetAttributes(); + std::vector axes_values; + if (attributes.find("axes") != attributes.end()) { + axes_values = RetrieveValues(attributes.at("axes")); + } + + // Get the inputs for the new LayerNormalization node. + // scale and bias could be multi-dims; we only support it for training at the moment + // because SkipLayerNorm kernel, for example, has dependency on single dim size + NodeArg* scale = nullptr; + for (size_t i = 0; i < mul_node.MutableInputDefs().size(); i++) { + if (graph_utils::NodeArgIsConstant(graph, *(mul_node.MutableInputDefs()[i])) || + graph_utils::IsGraphInput(graph, mul_node.MutableInputDefs()[i])) { +#ifdef ENABLE_TRAINING + if (axes_values.empty() || + mul_node.MutableInputDefs()[i]->Shape()->dim_size() == static_cast(axes_values.size())) { + scale = mul_node.MutableInputDefs()[i]; + } +#else + // Scale must be 1d. + if (mul_node.MutableInputDefs()[i]->Shape()->dim_size() == 1) { + scale = mul_node.MutableInputDefs()[i]; + } +#endif + } + } + + if (scale == nullptr) { + continue; + } + + const std::vector layer_norm_input_defs{pow_node.MutableInputDefs()[0], scale}; + Node& layer_norm_node = graph.AddNode(graph.GenerateNodeName("SimplifiedLayerNormalization"), + "SimplifiedLayerNormalization", + "fused LayerNorm subgraphs ", + layer_norm_input_defs, + {}, {}, kOnnxDomain); + + // Get constant "epsilon" from "Add" node if available. Else, default value will be used. + const ONNX_NAMESPACE::TensorProto* tensor_proto = graph_utils::GetConstantInitializer(graph, add_node.MutableInputDefs()[1]->Name()); + if (tensor_proto != nullptr && + tensor_proto->data_type() == ONNX_NAMESPACE::TensorProto_DataType_FLOAT) { + Initializer initializer{*tensor_proto, graph.ModelPath()}; + layer_norm_node.AddAttribute("epsilon", initializer.data()[0]); + } else { + layer_norm_node.AddAttribute("epsilon", DEFAULT_LAYERNORM_EPSILON); + } + + // Assign provider to this new node. Provider should be same as the provider for old node. + layer_norm_node.SetExecutionProviderType(reduce_mean_node.GetExecutionProviderType()); + + // move input edges to add (first in list) across to the layer_norm_node. + // move output definitions and output edges from mul_node (last in list) to layer_norm_node. + // remove all the other nodes. + graph_utils::FinalizeNodeFusion(graph, nodes_to_remove, layer_norm_node); + +#ifdef ENABLE_TRAINING + // add one extra output def, so we have 2 output defs that match what gradient builder expected + layer_norm_node.MutableOutputDefs().push_back(&graph.GetOrCreateNodeArg(graph.GenerateNodeArgName("saved_inv_std_var"), nullptr)); +#endif + + modified = true; + } + return Status::OK(); +} } // namespace onnxruntime diff --git a/onnxruntime/core/optimizer/layer_norm_fusion.h b/onnxruntime/core/optimizer/layer_norm_fusion.h index cb1b2a31ce..ee67471de8 100644 --- a/onnxruntime/core/optimizer/layer_norm_fusion.h +++ b/onnxruntime/core/optimizer/layer_norm_fusion.h @@ -24,4 +24,21 @@ class LayerNormFusion : public GraphTransformer { Status ApplyImpl(Graph& graph, bool& modified, int graph_level, const logging::Logger& logger) const override; }; +/** +@Class SimplifiedLayerNormFusion + +Rewrite graph fusing Layer Normalization subgraph to a single LayerNormalization node. + +The formula corresponding to LayerNorm activation subgraph: +(x ) / sqrt(var(x, axis)) * scale, where x is the input, and var() is given by mean(x^2, axis). + +*/ +class SimplifiedLayerNormFusion : public GraphTransformer { + public: + SimplifiedLayerNormFusion(const std::unordered_set& compatible_execution_providers = {}) noexcept + : GraphTransformer("SimplifiedLayerNormFusion", compatible_execution_providers) {} + + Status ApplyImpl(Graph& graph, bool& modified, int graph_level, const logging::Logger& logger) const override; +}; + } // namespace onnxruntime diff --git a/onnxruntime/test/contrib_ops/layer_norm_test.cc b/onnxruntime/test/contrib_ops/layer_norm_test.cc index 41037ec6cb..98e409cf1d 100644 --- a/onnxruntime/test/contrib_ops/layer_norm_test.cc +++ b/onnxruntime/test/contrib_ops/layer_norm_test.cc @@ -10,6 +10,8 @@ namespace test { constexpr auto k_epsilon_default = 1e-5f; constexpr auto k_random_data_min = -10.0f; constexpr auto k_random_data_max = 10.0f; +const std::string SIMPLIFIED_LAYER_NORM_OP = "SimplifiedLayerNormalization"; +const std::string LAYER_NORM_OP = "LayerNormalization"; // The dimensions are split at the specified axis into N (before, exclusive) and M (after, inclusive). static Status SplitDims( @@ -24,6 +26,7 @@ static Status SplitDims( } static void TestLayerNorm(const std::vector& x_dims, + const std::string& op, optional epsilon, int64_t axis = -1, int64_t keep_dims = 1) { @@ -38,8 +41,8 @@ static void TestLayerNorm(const std::vector& x_dims, ASSERT_NE(keep_dims, 0); const std::vector& stats_dims = keep_dims ? n_and_ones_dims : n_dims; - - CompareOpTester test("LayerNormalization"); + + CompareOpTester test(op.c_str()); test.AddAttribute("axis", axis); test.AddAttribute("keep_dims", keep_dims); if (epsilon.has_value()) { @@ -54,14 +57,18 @@ static void TestLayerNorm(const std::vector& x_dims, test.AddInput("X", n_x_m_dims, X_data); test.AddInput("scale", m_dims, scale_data, true); - test.AddInput("B", m_dims, B_data, true); + if (op.compare(SIMPLIFIED_LAYER_NORM_OP) != 0) { + test.AddInput("B", m_dims, B_data, true); + } std::vector Y_data = FillZeros(n_x_m_dims); std::vector mean_data = FillZeros(stats_dims); std::vector var_data = FillZeros(stats_dims); test.AddOutput("output", n_x_m_dims, Y_data); - test.AddOutput("mean", stats_dims, mean_data); + if (op.compare(SIMPLIFIED_LAYER_NORM_OP) != 0) { + test.AddOutput("mean", stats_dims, mean_data); + } test.AddOutput("var", stats_dims, var_data); test.CompareWithCPU(kCudaExecutionProvider); @@ -69,23 +76,44 @@ static void TestLayerNorm(const std::vector& x_dims, TEST(CudaKernelTest, LayerNorm_SmallSizeTensor) { const std::vector X_dims{4, 20, 128}; - TestLayerNorm(X_dims, k_epsilon_default); + TestLayerNorm(X_dims, LAYER_NORM_OP, k_epsilon_default); } TEST(CudaKernelTest, LayerNorm_SmallSizeTensor_IntermediateAxis) { const std::vector X_dims{4, 20, 8, 16}; const int64_t axis = -2; - TestLayerNorm(X_dims, k_epsilon_default, axis); + TestLayerNorm(X_dims, LAYER_NORM_OP, k_epsilon_default, axis); } TEST(CudaKernelTest, LayerNorm_MidSizeTensor) { std::vector X_dims{8, 80, 768}; - TestLayerNorm(X_dims, k_epsilon_default); + TestLayerNorm(X_dims, LAYER_NORM_OP, k_epsilon_default); } TEST(CudaKernelTest, LayerNorm_LargeSizeTensor) { std::vector X_dims{16, 512, 1024}; - TestLayerNorm(X_dims, k_epsilon_default); + TestLayerNorm(X_dims, LAYER_NORM_OP, k_epsilon_default); +} + +TEST(CudaKernelTest, SimplifiedLayerNorm_SmallSizeTensor) { + const std::vector X_dims{4, 20, 128}; + TestLayerNorm(X_dims, SIMPLIFIED_LAYER_NORM_OP, k_epsilon_default); +} + +TEST(CudaKernelTest, SimplifiedLayerNorm_SmallSizeTensor_IntermediateAxis) { + const std::vector X_dims{4, 20, 8, 16}; + const int64_t axis = -2; + TestLayerNorm(X_dims, SIMPLIFIED_LAYER_NORM_OP, k_epsilon_default, axis); +} + +TEST(CudaKernelTest, SimplifiedLayerNorm_MidSizeTensor) { + std::vector X_dims{8, 80, 768}; + TestLayerNorm(X_dims, SIMPLIFIED_LAYER_NORM_OP, k_epsilon_default); +} + +TEST(CudaKernelTest, SimplifiedLayerNorm_LargeSizeTensor) { + std::vector X_dims{16, 512, 1024}; + TestLayerNorm(X_dims, SIMPLIFIED_LAYER_NORM_OP, k_epsilon_default); } #endif } // namespace test diff --git a/onnxruntime/test/optimizer/graph_transform_test.cc b/onnxruntime/test/optimizer/graph_transform_test.cc index b2ef4453c7..3fbc4b8437 100644 --- a/onnxruntime/test/optimizer/graph_transform_test.cc +++ b/onnxruntime/test/optimizer/graph_transform_test.cc @@ -2481,6 +2481,39 @@ TEST_F(GraphTransformationTests, LayerNormWithSubDupFusionTest) { } } +TEST_F(GraphTransformationTests, SimplifiedLayerNormFusionTest) { + auto model_uri = MODEL_FOLDER "fusion/layer_norm_t5.onnx"; + std::shared_ptr p_model; + ASSERT_STATUS_OK(Model::Load(model_uri, p_model, nullptr, *logger_)); + 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_TRUE(ret.IsOK()); + + Model::Save(*p_model, "t5_fused.onnx"); + std::map op_to_count = CountOpsInGraph(graph); + ASSERT_TRUE(op_to_count["Div"] == 0); + ASSERT_TRUE(op_to_count["Add"] == 0); + ASSERT_TRUE(op_to_count["ReduceMean"] == 0); + ASSERT_TRUE(op_to_count["Pow"] == 0); + ASSERT_TRUE(op_to_count["Sqrt"] == 0); + ASSERT_TRUE(op_to_count["SimplifiedLayerNormalization"] == 1); + + for (const Node& node : graph.Nodes()) { + if (node.OpType() == "SimplifiedLayerNormalization") { + // LayerNormalization should have two inputs. + EXPECT_EQ(node.InputDefs().size(), 2u) << "LayerNormalization number of inputs does not equal to 2. Got:" << node.InputDefs().size(); + // LayerNormalization input "scale" and "bias" should have the same dimension. + const TensorShapeProto* scale_shape = node.InputDefs()[1]->Shape(); + EXPECT_EQ(scale_shape->dim_size(), 1) << "LayerNormalization scale should be 1D. Got: " << scale_shape->dim_size(); + } else { + EXPECT_TRUE(false) << "Unexpected node " << node.Name(); + } + } +} + static void TestSkipLayerNormFusion(const std::basic_string& file_path, int add_count, int ln_count, int skip_ln_count, logging::Logger* logger) { std::shared_ptr p_model; diff --git a/onnxruntime/test/testdata/transform/fusion/layer_norm_t5.onnx b/onnxruntime/test/testdata/transform/fusion/layer_norm_t5.onnx new file mode 100644 index 0000000000000000000000000000000000000000..19614b9f36297a31424bf28670ab5d937ce8ff85 GIT binary patch literal 419 zcmZutO-lnY5ZyGdDKpZlWwUP2G=HZ@qZ)NBBeh z3H~RY?Yi}1$V(pY4Kt77?&%wa4v+!xshx+(Bmg7Y%`aw17ZXZ`F n1bCc?CN#K*=pfiV8pHpLfPv#RgB{q4tvNPHoZ8Gj{ + NodeDef(OpDef{"SimplifiedLayerNormalizationGrad", kMSDomain, 1}, + {GO(0), I(0), I(1), O(1)}, + {GI(0), GI(1)}, + {SrcNodeAttributes()})}; +} + IMPLEMENT_GRADIENT_BUILDER(GetBatchNormalizationGradient) { auto attributes = SrcNodeAttributes(); if (attributes.find("epsilon") != attributes.end()) { diff --git a/orttraining/orttraining/core/graph/gradient_builder.h b/orttraining/orttraining/core/graph/gradient_builder.h index 02282a643f..92f679753b 100644 --- a/orttraining/orttraining/core/graph/gradient_builder.h +++ b/orttraining/orttraining/core/graph/gradient_builder.h @@ -56,6 +56,7 @@ DECLARE_GRADIENT_BUILDER(GetGeluGradient) DECLARE_GRADIENT_BUILDER(GetBiasGeluGradient) DECLARE_GRADIENT_BUILDER(GetFastGeluGradient) DECLARE_GRADIENT_BUILDER(GetLayerNormalizationGradient) +DECLARE_GRADIENT_BUILDER(GetSimplifiedLayerNormalizationGradient) DECLARE_GRADIENT_BUILDER(GetBatchNormalizationGradient) DECLARE_GRADIENT_BUILDER(GetMegatronFGradient) DECLARE_GRADIENT_BUILDER(GetMegatronGGradient) diff --git a/orttraining/orttraining/core/graph/gradient_builder_registry.cc b/orttraining/orttraining/core/graph/gradient_builder_registry.cc index 11ef5ec6b8..9ac3f51419 100644 --- a/orttraining/orttraining/core/graph/gradient_builder_registry.cc +++ b/orttraining/orttraining/core/graph/gradient_builder_registry.cc @@ -87,6 +87,7 @@ void GradientBuilderRegistry::RegisterGradientBuilders() { REGISTER_GRADIENT_BUILDER("BiasGelu", GetBiasGeluGradient); REGISTER_GRADIENT_BUILDER("FastGelu", GetFastGeluGradient); REGISTER_GRADIENT_BUILDER("LayerNormalization", GetLayerNormalizationGradient); + REGISTER_GRADIENT_BUILDER("SimplifiedLayerNormalization", GetSimplifiedLayerNormalizationGradient); REGISTER_GRADIENT_BUILDER("BatchNormalization", GetBatchNormalizationGradient); REGISTER_GRADIENT_BUILDER("MegatronF", GetMegatronFGradient); REGISTER_GRADIENT_BUILDER("MegatronG", GetMegatronGGradient); diff --git a/orttraining/orttraining/core/graph/training_op_defs.cc b/orttraining/orttraining/core/graph/training_op_defs.cc index 6ec1b4d3e4..e37a801c25 100644 --- a/orttraining/orttraining/core/graph/training_op_defs.cc +++ b/orttraining/orttraining/core/graph/training_op_defs.cc @@ -1616,6 +1616,30 @@ Example 4: "U", {"tensor(float)", "tensor(bfloat16)"}, "Constrain mean and inv_std_var to float tensors."); + + ONNX_CONTRIB_OPERATOR_SCHEMA(SimplifiedLayerNormalizationGrad) + .SetDomain(kMSDomain) + .SinceVersion(1) + .SetSupportLevel(OpSchema::SupportType::EXPERIMENTAL) + .SetDoc("SimplifiedLayerNormalizationGrad") + .Attr("axis", + "The first normalization dimension: normalization will be performed along dimensions axis : rank(inputs).", + AttributeProto::INT, static_cast(-1)) + .AllowUncheckedAttributes() + .Input(0, "Y_grad", "The gradient tensor from output.", "T") + .Input(1, "X", "Input data tensor from the forward path", "T") + .Input(2, "scale", "Scale tensor.", "T") + .Input(3, "inv_std_var", "inverse std variance of X.", "U") + .Output(0, "X_grad", "Gradient of the input.", "T") + .Output(1, "scale_grad", "Gradient of the scale.", "T") + .TypeConstraint( + "T", + {"tensor(float16)", "tensor(float)", "tensor(double)", "tensor(bfloat16)"}, + "Constrain input and output types (except mean and inv_std_var) to float tensors.") + .TypeConstraint( + "U", + {"tensor(float)"}, + "Constrain mean and inv_std_var to float tensors."); ONNX_CONTRIB_OPERATOR_SCHEMA(InvertibleLayerNormalizationGrad) .SetDomain(kMSDomain) diff --git a/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc b/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc index 4c515e20b4..90237901a1 100644 --- a/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc +++ b/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc @@ -77,6 +77,7 @@ std::vector> GeneratePreTrainingTransformers( transformers.emplace_back(onnxruntime::make_unique(compatible_eps)); transformers.emplace_back(onnxruntime::make_unique(compatible_eps)); + transformers.emplace_back(onnxruntime::make_unique(compatible_eps)); transformers.emplace_back(onnxruntime::make_unique(compatible_eps)); #ifdef USE_CUDA diff --git a/orttraining/orttraining/test/gradient/gradient_ops_test.cc b/orttraining/orttraining/test/gradient/gradient_ops_test.cc index 15a3ed19f5..7c59539c1b 100644 --- a/orttraining/orttraining/test/gradient/gradient_ops_test.cc +++ b/orttraining/orttraining/test/gradient/gradient_ops_test.cc @@ -1744,6 +1744,23 @@ TEST(GradientCheckerTest, LayerNormGrad) { EXPECT_IS_TINIER_THAN(max_error, error_tolerance); } } + +TEST(GradientCheckerTest, SimplifiedLayerNormGrad) { + GradientChecker gradient_checker; + { + TensorShape shape({2, 3, 8}); + TensorInfo x_info{shape, true}; + TensorInfo scale_info{{8}, true}; + TensorInfo var_info{{2, 3, 1}, false}; + + float max_error; + float error_tolerance = 1e-2f; + + OpDef op_def{"SimplifiedLayerNormalization"}; + gradient_checker.ComputeGradientError(op_def, {x_info, scale_info}, {shape, var_info}, &max_error); + EXPECT_IS_TINIER_THAN(max_error, error_tolerance); + } +} #endif TEST(GradientUtilsTest, InPlaceAccumulatorFloat32) { diff --git a/orttraining/orttraining/test/training_ops/cuda/layer_norm_test.cc b/orttraining/orttraining/test/training_ops/cuda/layer_norm_test.cc index cbd23c43f6..8717e307b2 100644 --- a/orttraining/orttraining/test/training_ops/cuda/layer_norm_test.cc +++ b/orttraining/orttraining/test/training_ops/cuda/layer_norm_test.cc @@ -11,6 +11,8 @@ namespace test { constexpr auto k_epsilon_default = 1e-5f; constexpr auto k_random_data_min = -10.0f; constexpr auto k_random_data_max = 10.0f; +const std::string SIMPLIFIED_LAYER_NORM_GRAD_OP = "SimplifiedLayerNormalizationGrad"; +const std::string LAYER_NORM_GRAD_OP = "LayerNormalizationGrad"; // The dimensions are split at the specified axis into N (before, exclusive) and M (after, inclusive). static Status SplitDims( @@ -26,6 +28,7 @@ static Status SplitDims( static void TestLayerNormGrad( const std::vector& x_dims, + const std::string& op, int64_t axis = -1, double error_tolerance = 1e-4) { const std::vector& n_x_m_dims = x_dims; @@ -35,7 +38,7 @@ static void TestLayerNormGrad( const auto N = std::accumulate(n_dims.begin(), n_dims.end(), static_cast(1), std::multiplies<>{}); const auto M = std::accumulate(m_dims.begin(), m_dims.end(), static_cast(1), std::multiplies<>{}); - CompareOpTester test{"LayerNormalizationGrad", 1, kMSDomain}; + CompareOpTester test{op.c_str(), 1, kMSDomain}; test.AddAttribute("axis", axis); @@ -55,14 +58,20 @@ static void TestLayerNormGrad( EigenRowVectorArrayMap mean{mean_data.data(), N}; EigenRowVectorArrayMap inv_std_var{inv_std_var_data.data(), N}; - mean = X.colwise().mean(); - inv_std_var = ((X.colwise().squaredNorm() / X.rows()) - mean.square() + k_epsilon_default).rsqrt(); + if (op.compare(SIMPLIFIED_LAYER_NORM_GRAD_OP) != 0) { + mean = X.colwise().mean(); + inv_std_var = ((X.colwise().squaredNorm() / X.rows()) - mean.square() + k_epsilon_default).rsqrt(); + } else { + inv_std_var = ((X.colwise().squaredNorm() / X.rows()) + k_epsilon_default).rsqrt(); + } } test.AddInput("Y_grad", n_x_m_dims, Y_grad_data); test.AddInput("X", n_x_m_dims, X_data); test.AddInput("scale", m_dims, scale_data, true); - test.AddInput("mean", n_dims, mean_data); + if (op.compare(SIMPLIFIED_LAYER_NORM_GRAD_OP) != 0) { + test.AddInput("mean", n_dims, mean_data); + } test.AddInput("inv_std_var", n_dims, inv_std_var_data); const auto X_grad_data = FillZeros(n_x_m_dims); @@ -71,30 +80,53 @@ static void TestLayerNormGrad( test.AddOutput("X_grad", n_x_m_dims, X_grad_data); test.AddOutput("scale_grad_data", m_dims, scale_grad_data); - test.AddOutput("bias_grad_data", m_dims, bias_grad_data); + if (op.compare(SIMPLIFIED_LAYER_NORM_GRAD_OP) != 0) { + test.AddOutput("bias_grad_data", m_dims, bias_grad_data); + } test.CompareWithCPU(kCudaExecutionProvider, error_tolerance); } TEST(CudaKernelTest, LayerNormGrad_SmallSizeTensor) { const std::vector X_dims{4, 20, 128}; - TestLayerNormGrad(X_dims); + TestLayerNormGrad(X_dims, LAYER_NORM_GRAD_OP); } TEST(CudaKernelTest, LayerNormGrad_SmallSizeTensor_IntermediateAxis) { const std::vector X_dims{4, 20, 16, 8}; const int64_t axis = -2; - TestLayerNormGrad(X_dims, axis); + TestLayerNormGrad(X_dims, LAYER_NORM_GRAD_OP, axis); } TEST(CudaKernelTest, LayerNormGrad_MidSizeTensor) { const std::vector X_dims{8, 80, 768}; - TestLayerNormGrad(X_dims); + TestLayerNormGrad(X_dims, LAYER_NORM_GRAD_OP); } TEST(CudaKernelTest, LayerNormGrad_LargeSizeTensor) { const std::vector X_dims{16, 512, 1024}; - TestLayerNormGrad(X_dims, -1, 5e-3); + TestLayerNormGrad(X_dims, LAYER_NORM_GRAD_OP, -1, 5e-3); +} + +TEST(CudaKernelTest, SimplifiedLayerNormGrad_SmallSizeTensor) { + const std::vector X_dims{4, 20, 128}; + TestLayerNormGrad(X_dims, SIMPLIFIED_LAYER_NORM_GRAD_OP); +} + +TEST(CudaKernelTest, SimplifiedLayerNormGrad_SmallSizeTensor_IntermediateAxis) { + const std::vector X_dims{4, 20, 16, 8}; + const int64_t axis = -2; + TestLayerNormGrad(X_dims, SIMPLIFIED_LAYER_NORM_GRAD_OP, axis); +} + +TEST(CudaKernelTest, SimplifiedLayerNormGrad_MidSizeTensor) { + const std::vector X_dims{8, 80, 768}; + TestLayerNormGrad(X_dims, SIMPLIFIED_LAYER_NORM_GRAD_OP); +} + +TEST(CudaKernelTest, SimplifiedLayerNormGrad_LargeSizeTensor) { + const std::vector X_dims{16, 512, 1024}; + TestLayerNormGrad(X_dims, SIMPLIFIED_LAYER_NORM_GRAD_OP, -1, 5e-3); } static void TestInvertibleLayerNormGrad( diff --git a/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc b/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc index 7147050a50..e18ad4bd4d 100644 --- a/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc +++ b/orttraining/orttraining/training_ops/cpu/cpu_training_kernels.cc @@ -85,6 +85,8 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, GistB class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, GistBinarizeDecoder); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, LayerNormalizationGrad); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, double, LayerNormalizationGrad); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, SimplifiedLayerNormalizationGrad); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, double, SimplifiedLayerNormalizationGrad); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, InvertibleLayerNormalizationGrad); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, double, InvertibleLayerNormalizationGrad); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SliceGrad); @@ -181,6 +183,8 @@ Status RegisterCpuTrainingKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/orttraining/orttraining/training_ops/cpu/nn/layer_norm.cc b/orttraining/orttraining/training_ops/cpu/nn/layer_norm.cc index e55fe5e411..30e1cfc405 100644 --- a/orttraining/orttraining/training_ops/cpu/nn/layer_norm.cc +++ b/orttraining/orttraining/training_ops/cpu/nn/layer_norm.cc @@ -20,9 +20,18 @@ namespace contrib { kCpuExecutionProvider, \ KernelDefBuilder() \ .TypeConstraint("T", DataTypeImpl::GetTensorType()), \ - LayerNormGrad); \ + LayerNormGrad); \ ONNX_OPERATOR_TYPED_KERNEL_EX( \ - InvertibleLayerNormalizationGrad, \ + SimplifiedLayerNormalizationGrad, \ + kMSDomain, \ + 1, \ + T, \ + kCpuExecutionProvider, \ + KernelDefBuilder() \ + .TypeConstraint("T", DataTypeImpl::GetTensorType()), \ + LayerNormGrad); \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + InvertibleLayerNormalizationGrad, \ kMSDomain, \ 1, \ T, \ @@ -36,41 +45,46 @@ REGISTER_KERNEL_TYPED(double) #undef REGISTER_KERNEL_TYPED -template -LayerNormGrad::LayerNormGrad(const OpKernelInfo& op_kernel_info) +template +LayerNormGrad::LayerNormGrad(const OpKernelInfo& op_kernel_info) : OpKernel{op_kernel_info} { ORT_ENFORCE(op_kernel_info.GetAttr("axis", &axis_).IsOK()); } -template -Status LayerNormGrad::Compute(OpKernelContext* op_kernel_context) const { - const Tensor* Y_grad = op_kernel_context->Input(0); - const Tensor* X = op_kernel_context->Input(1); - const Tensor* scale = op_kernel_context->Input(2); - const Tensor* mean = op_kernel_context->Input(3); - const Tensor* inv_std_var = op_kernel_context->Input(4); - +template +Status LayerNormGrad::Compute(OpKernelContext* op_kernel_context) const { + int input_index = 0; + const Tensor* Y_grad = op_kernel_context->Input(input_index++); + const Tensor* X = op_kernel_context->Input(input_index++); const auto& X_shape = X->Shape(); const auto axis = HandleNegativeAxis(axis_, X_shape.NumDimensions()); const auto N = X_shape.SizeToDimension(axis); const auto M = X_shape.SizeFromDimension(axis); ORT_ENFORCE(M != 1); + + const Tensor* scale = op_kernel_context->Input(input_index++); + const Tensor* mean; + if (!simplified) { + mean = op_kernel_context->Input(input_index++); + } + const Tensor* inv_std_var = op_kernel_context->Input(input_index); + const auto& scale_shape = scale->Shape(); Tensor* X_grad = op_kernel_context->Output(0, X_shape); Tensor* scale_grad = op_kernel_context->Output(1, scale_shape); - Tensor* bias_grad = op_kernel_context->Output(2, scale_shape); + Tensor* bias_grad = (!simplified) ? op_kernel_context->Output(2, scale_shape) : nullptr; // Note: Eigen has column-major storage order by default ConstEigenArrayMap Y_grad_arr{Y_grad->Data(), M, N}; ConstEigenArrayMap X_arr{X->Data(), M, N}; ConstEigenVectorArrayMap scale_vec{scale->Data(), M}; - ConstEigenVectorArrayMap mean_vec{mean->Data(), N}; + ConstEigenVectorArrayMap mean_vec{simplified ? nullptr : mean->Data(), N}; ConstEigenVectorArrayMap inv_std_var_vec{inv_std_var->Data(), N}; EigenArrayMap X_grad_arr{X_grad->MutableData(), M, N}; EigenVectorArrayMap scale_grad_vec{scale_grad->MutableData(), M}; - EigenVectorArrayMap bias_grad_vec{bias_grad->MutableData(), M}; + EigenVectorArrayMap bias_grad_vec = (!simplified) ? EigenVectorArrayMap{bias_grad->MutableData(), M} : EigenVectorArrayMap{nullptr, 0}; using Array = Eigen::ArrayXX; using RowVector = Eigen::Array; @@ -80,29 +94,36 @@ Status LayerNormGrad::Compute(OpKernelContext* op_kernel_context) const { // B = Y_grad * scale * inv_std_var // C = Y_grad * scale * inv_std_var * (X - mean(X)) * inv_std_var + // Simplified Layer Norm + // A = Y_grad * X * inv_std_var + // B = Y_grad * scale * inv_std_var + // C = Y_grad * scale * inv_std_var * X * inv_std_var // A, B, and C are M x N - - Array X_mean_difference_over_std_var = + Array X_mean_difference_over_std_var; + if (simplified) { + X_mean_difference_over_std_var = + X_arr.rowwise() * inv_std_var_vec.cast().transpose(); + } else { + X_mean_difference_over_std_var = (X_arr.rowwise() - mean_vec.cast().transpose()).rowwise() * inv_std_var_vec.cast().transpose(); + } Array A = Y_grad_arr * X_mean_difference_over_std_var; Array B = (Y_grad_arr.colwise() * scale_vec).rowwise() * inv_std_var_vec.cast().transpose(); Array C = B * X_mean_difference_over_std_var; - // mean_B = mean(Y_grad * scale * inv_std_var) - RowVector mean_B = B.colwise().mean(); // 1 x N - - // mean_C = mean(Y_grad * scale * inv_std_var * (X - mean(X)) * inv_std_var) RowVector mean_C = C.colwise().mean(); // 1 x N - // X_grad = Y_grad * scale * inv_std_var - mean_B - (X - mean(X)) * inv_std_var * mean_C - // = B - mean_B - (X - mean(X)) * inv_std_var * mean_c - X_grad_arr = B.rowwise() - mean_B - X_mean_difference_over_std_var.rowwise() * mean_C; + if (simplified) { + X_grad_arr = B - X_mean_difference_over_std_var.rowwise() * mean_C; + } else { + RowVector mean_B = B.colwise().mean(); // 1 x N + X_grad_arr = B.rowwise() - mean_B - X_mean_difference_over_std_var.rowwise() * mean_C; + } - // bias_grad = sum(Y_grad) - bias_grad_vec = Y_grad_arr.rowwise().sum(); + if (!simplified) { + bias_grad_vec = Y_grad_arr.rowwise().sum(); + } - // scale_grad = sum(Y_grad * (X - mean(X)) * inv_std_var) - // = sum(A) scale_grad_vec = A.rowwise().sum(); return Status::OK(); diff --git a/orttraining/orttraining/training_ops/cpu/nn/layer_norm.h b/orttraining/orttraining/training_ops/cpu/nn/layer_norm.h index 4cd6207c26..185b36c69c 100644 --- a/orttraining/orttraining/training_ops/cpu/nn/layer_norm.h +++ b/orttraining/orttraining/training_ops/cpu/nn/layer_norm.h @@ -8,7 +8,7 @@ namespace onnxruntime { namespace contrib { -template +template class LayerNormGrad final : public OpKernel { public: LayerNormGrad(const OpKernelInfo& op_kernel_info); diff --git a/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc b/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc index ebcc68da51..9c01fb5154 100644 --- a/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc +++ b/orttraining/orttraining/training_ops/cuda/cuda_training_kernels.cc @@ -106,6 +106,9 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1 class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float, LayerNormalizationGrad); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double, LayerNormalizationGrad); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float, LayerNormalizationGrad); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float, SimplifiedLayerNormalizationGrad); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double, SimplifiedLayerNormalizationGrad); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float, SimplifiedLayerNormalizationGrad); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float_float, InvertibleLayerNormalizationGrad); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double_double, InvertibleLayerNormalizationGrad); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16_float, InvertibleLayerNormalizationGrad); @@ -233,6 +236,9 @@ Status RegisterCudaTrainingKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/orttraining/orttraining/training_ops/cuda/nn/layer_norm.cc b/orttraining/orttraining/training_ops/cuda/nn/layer_norm.cc index d7f8acfc8c..b6e9677c56 100644 --- a/orttraining/orttraining/training_ops/cuda/nn/layer_norm.cc +++ b/orttraining/orttraining/training_ops/cuda/nn/layer_norm.cc @@ -19,7 +19,7 @@ namespace cuda { KernelDefBuilder() \ .TypeConstraint("T", DataTypeImpl::GetTensorType()) \ .TypeConstraint("U", DataTypeImpl::GetTensorType()), \ - LayerNormGrad); \ + LayerNormGrad); \ ONNX_OPERATOR_TYPED_KERNEL_EX( \ InvertibleLayerNormalizationGrad, \ kMSDomain, \ @@ -29,31 +29,46 @@ namespace cuda { KernelDefBuilder() \ .TypeConstraint("T", DataTypeImpl::GetTensorType()) \ .TypeConstraint("U", DataTypeImpl::GetTensorType()), \ - InvertibleLayerNormGrad); + InvertibleLayerNormGrad); \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + SimplifiedLayerNormalizationGrad, \ + kMSDomain, \ + 1, \ + T##_##U, \ + kCudaExecutionProvider, \ + KernelDefBuilder() \ + .TypeConstraint("T", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("U", DataTypeImpl::GetTensorType()), \ + LayerNormGrad); + REGISTER_GRADIENT_KERNEL_TYPED(float, float) REGISTER_GRADIENT_KERNEL_TYPED(double, double) REGISTER_GRADIENT_KERNEL_TYPED(MLFloat16, float) -template -LayerNormGrad::LayerNormGrad(const OpKernelInfo& op_kernel_info) : CudaKernel(op_kernel_info) { +template +LayerNormGrad::LayerNormGrad(const OpKernelInfo& op_kernel_info) : CudaKernel(op_kernel_info) { ORT_ENFORCE(op_kernel_info.GetAttr("axis", &axis_).IsOK()); } -template -Status LayerNormGrad::ComputeInternal(OpKernelContext* p_op_kernel_context) const { +template +Status LayerNormGrad::ComputeInternal(OpKernelContext* p_op_kernel_context) const { typedef typename ToCudaType::MappedType CudaT; typedef typename ToCudaType::MappedType CudaU; // Inputs - const Tensor* Y_grad = p_op_kernel_context->Input(0); - const Tensor* X = p_op_kernel_context->Input(1); - const Tensor* scale = p_op_kernel_context->Input(2); - const Tensor* mean = p_op_kernel_context->Input(3); - const Tensor* inv_std_var = p_op_kernel_context->Input(4); + int input_index = 0; + const Tensor* Y_grad = p_op_kernel_context->Input(input_index++); + const Tensor* X = p_op_kernel_context->Input(input_index++); + const Tensor* scale = p_op_kernel_context->Input(input_index++); + const Tensor* mean; + if (!simplified) { + mean = p_op_kernel_context->Input(input_index++); + } + const Tensor* inv_std_var = p_op_kernel_context->Input(input_index); auto Y_grad_data = reinterpret_cast(Y_grad->template Data()); auto X_data = reinterpret_cast(X->template Data()); auto scale_data = reinterpret_cast(scale->template Data()); - auto mean_data = reinterpret_cast(mean->template Data()); + auto mean_data = simplified ? nullptr: reinterpret_cast(mean->template Data()); auto inv_std_var_data = reinterpret_cast(inv_std_var->template Data()); const TensorShape& x_shape = X->Shape(); @@ -65,17 +80,19 @@ Status LayerNormGrad::ComputeInternal(OpKernelContext* p_op_kernel_context // Outputs Tensor* X_grad = p_op_kernel_context->Output(0, x_shape); auto X_grad_data = reinterpret_cast(X_grad->template MutableData()); - Tensor* scale_grad = p_op_kernel_context->Output(1, scale->Shape()); - Tensor* bias_grad = p_op_kernel_context->Output(2, scale->Shape()); auto scale_grad_data = reinterpret_cast(scale_grad->template MutableData()); - auto bias_grad_data = reinterpret_cast(bias_grad->template MutableData()); + CudaT* bias_grad_data = nullptr; + if (!simplified) { + Tensor* bias_grad = p_op_kernel_context->Output(2, scale->Shape()); + bias_grad_data = reinterpret_cast(bias_grad->template MutableData()); + } const int part_size = 16; auto part_grad_gamma = GetScratchBuffer(part_size * n2); auto part_grad_beta = GetScratchBuffer(part_size * n2); - HostLayerNormGradient(GetDeviceProp(), Y_grad_data, X_data, reinterpret_cast(NULL), + HostLayerNormGradient(GetDeviceProp(), Y_grad_data, X_data, reinterpret_cast(NULL), scale_data, reinterpret_cast(NULL), mean_data, inv_std_var_data, n1, n2, X_grad_data, scale_grad_data, bias_grad_data, part_grad_gamma.get(), part_grad_beta.get(), part_size); @@ -124,7 +141,7 @@ Status InvertibleLayerNormGrad::ComputeInternal(OpKernelContext* p_op_kern auto part_grad_gamma = GetScratchBuffer(part_size * n2); auto part_grad_beta = GetScratchBuffer(part_size * n2); - HostLayerNormGradient(GetDeviceProp(), Y_grad_data, reinterpret_cast(NULL), Y_data, + HostLayerNormGradient(GetDeviceProp(), Y_grad_data, reinterpret_cast(NULL), Y_data, scale_data, bias_data, reinterpret_cast(NULL), inv_std_var_data, n1, n2, X_grad_data, scale_grad_data, bias_grad_data, part_grad_gamma.get(), part_grad_beta.get(), part_size); diff --git a/orttraining/orttraining/training_ops/cuda/nn/layer_norm.h b/orttraining/orttraining/training_ops/cuda/nn/layer_norm.h index ab092ed12a..2b4ce47858 100644 --- a/orttraining/orttraining/training_ops/cuda/nn/layer_norm.h +++ b/orttraining/orttraining/training_ops/cuda/nn/layer_norm.h @@ -15,7 +15,7 @@ class LayerNorm final : public CudaKernel { double epsilon_; }; -template +template class LayerNormGrad final : public CudaKernel { public: LayerNormGrad(const OpKernelInfo& op_kernel_info); diff --git a/orttraining/orttraining/training_ops/cuda/nn/layer_norm_impl.cu b/orttraining/orttraining/training_ops/cuda/nn/layer_norm_impl.cu index 11753e080b..23f206bb39 100644 --- a/orttraining/orttraining/training_ops/cuda/nn/layer_norm_impl.cu +++ b/orttraining/orttraining/training_ops/cuda/nn/layer_norm_impl.cu @@ -30,352 +30,42 @@ namespace cuda { using namespace onnxruntime::cuda; -template -__device__ void cuWelfordOnlineSum( - const U curr, - U& mu, - U& sigma2, - U& count) { - count = count + U(1); - U delta = curr - mu; - U lmean = mu + delta / count; - mu = lmean; - U delta2 = curr - lmean; - sigma2 = sigma2 + delta * delta2; -} - -template -__device__ void cuChanOnlineSum( - const U muB, - const U sigma2B, - const U countB, - U& mu, - U& sigma2, - U& count) { - U delta = muB - mu; - U nA = count; - U nB = countB; - count = count + countB; - U nX = count; - if (nX > U(0)) { - nA = nA / nX; - nB = nB / nX; - mu = nA * mu + nB * muB; - sigma2 = sigma2 + sigma2B + delta * delta * nA * nB * nX; - } else { - mu = U(0); - sigma2 = U(0); - } -} - -template -__device__ void cuWelfordMuSigma2( - const T* __restrict__ vals, - const int n1, - const int n2, - const int i1, - U& mu, - U& sigma2, - U* buf) { - // Assumptions: - // 1) blockDim.x == GPU_WARP_SIZE - // 2) Tensor is contiguous - // 3) 2*blockDim.y*sizeof(U)+blockDim.y*sizeof(int) shared memory available. - // - // compute variance and mean over n2 - U count = U(0); - mu = U(0); - sigma2 = U(0); - if (i1 < n1) { - // one warp normalizes one n1 index, - // synchronization is implicit - // initialize with standard Welford algorithm - const int numx = blockDim.x * blockDim.y; - const int thrx = threadIdx.x + threadIdx.y * blockDim.x; - const T* lvals = vals + i1 * n2; - int l = 4 * thrx; - for (; l + 3 < n2; l += 4 * numx) { - for (int k = 0; k < 4; ++k) { - U curr = static_cast(lvals[l + k]); - cuWelfordOnlineSum(curr, mu, sigma2, count); - } - } - for (; l < n2; ++l) { - U curr = static_cast(lvals[l]); - cuWelfordOnlineSum(curr, mu, sigma2, count); - } - // intra-warp reductions - for (int l = 0; l <= 4; ++l) { - int srcLaneB = (threadIdx.x + (1 << l)) & 31; - U muB = WARP_SHFL(mu, srcLaneB); - U countB = WARP_SHFL(count, srcLaneB); - U sigma2B = WARP_SHFL(sigma2, srcLaneB); - cuChanOnlineSum(muB, sigma2B, countB, mu, sigma2, count); - } - // threadIdx.x == 0 has correct values for each warp - // inter-warp reductions - if (blockDim.y > 1) { - U* ubuf = (U*)buf; - U* ibuf = (U*)(ubuf + blockDim.y); - for (int offset = blockDim.y / 2; offset > 0; offset /= 2) { - // upper half of warps write to shared - if (threadIdx.x == 0 && threadIdx.y >= offset && threadIdx.y < 2 * offset) { - const int wrt_y = threadIdx.y - offset; - ubuf[2 * wrt_y] = mu; - ubuf[2 * wrt_y + 1] = sigma2; - ibuf[wrt_y] = count; - } - __syncthreads(); - // lower half merges - if (threadIdx.x == 0 && threadIdx.y < offset) { - U muB = ubuf[2 * threadIdx.y]; - U sigma2B = ubuf[2 * threadIdx.y + 1]; - U countB = ibuf[threadIdx.y]; - cuChanOnlineSum(muB, sigma2B, countB, mu, sigma2, count); - } - __syncthreads(); - } - // threadIdx.x = 0 && threadIdx.y == 0 only thread that has correct values - if (threadIdx.x == 0 && threadIdx.y == 0) { - ubuf[0] = mu; - ubuf[1] = sigma2; - } - __syncthreads(); - mu = ubuf[0]; - sigma2 = ubuf[1] / U(n2); - // don't care about final value of count, we know count == n2 - } else { - mu = WARP_SHFL(mu, 0); - sigma2 = WARP_SHFL(sigma2 / U(n2), 0); - } - } -} - -template <> -__device__ void cuWelfordMuSigma2( - const half* __restrict__ vals, - const int n1, - const int n2, - const int i1, - float& mu, - float& sigma2, - float* buf) { - // Assumptions: - // 1) blockDim.x == GPU_WARP_SIZE - // 2) Tensor is contiguous - // 3) 2*blockDim.y*sizeof(U)+blockDim.y*sizeof(int) shared memory available. - // - // compute variance and mean over n2 - float count = 0.0f; - mu = float(0); - sigma2 = float(0); - if (i1 < n1) { - // one warp normalizes one n1 index, - // synchronization is implicit - // initialize with standard Welford algorithm - const int numx = blockDim.x * blockDim.y; - const int thrx = threadIdx.x + threadIdx.y * blockDim.x; - const half* lvals = vals + i1 * n2; - int l = 8 * thrx; - if ((((size_t)lvals) & 3) != 0) { - // 16 bit alignment - // first thread consumes first point - if (thrx == 0) { - float curr = static_cast(lvals[0]); - cuWelfordOnlineSum(curr, mu, sigma2, count); - } - ++l; - } - // at this point, lvals[l] are 32 bit aligned for all threads. - for (; l + 7 < n2; l += 8 * numx) { - for (int k = 0; k < 8; k += 2) { - float2 curr = __half22float2(*((__half2*)(lvals + l + k))); - cuWelfordOnlineSum(static_cast(curr.x), mu, sigma2, count); - cuWelfordOnlineSum(static_cast(curr.y), mu, sigma2, count); - } - } - for (; l < n2; ++l) { - float curr = static_cast(lvals[l]); - cuWelfordOnlineSum(curr, mu, sigma2, count); - } - // intra-warp reductions - for (int l = 0; l <= 4; ++l) { - int srcLaneB = (threadIdx.x + (1 << l)) & 31; - float muB = WARP_SHFL(mu, srcLaneB); - float countB = WARP_SHFL(count, srcLaneB); - float sigma2B = WARP_SHFL(sigma2, srcLaneB); - cuChanOnlineSum(muB, sigma2B, countB, mu, sigma2, count); - } - // threadIdx.x == 0 has correct values for each warp - // inter-warp reductions - if (blockDim.y > 1) { - float* ubuf = (float*)buf; - float* ibuf = (float*)(ubuf + blockDim.y); - for (int offset = blockDim.y / 2; offset > 0; offset /= 2) { - // upper half of warps write to shared - if (threadIdx.x == 0 && threadIdx.y >= offset && threadIdx.y < 2 * offset) { - const int wrt_y = threadIdx.y - offset; - ubuf[2 * wrt_y] = mu; - ubuf[2 * wrt_y + 1] = sigma2; - ibuf[wrt_y] = count; - } - __syncthreads(); - // lower half merges - if (threadIdx.x == 0 && threadIdx.y < offset) { - float muB = ubuf[2 * threadIdx.y]; - float sigma2B = ubuf[2 * threadIdx.y + 1]; - float countB = ibuf[threadIdx.y]; - cuChanOnlineSum(muB, sigma2B, countB, mu, sigma2, count); - } - __syncthreads(); - } - // threadIdx.x = 0 && threadIdx.y == 0 only thread that has correct values - if (threadIdx.x == 0 && threadIdx.y == 0) { - ubuf[0] = mu; - ubuf[1] = sigma2; - } - __syncthreads(); - mu = ubuf[0]; - sigma2 = ubuf[1] / float(n2); - // don't care about final value of count, we know count == n2 - } else { - mu = WARP_SHFL(mu, 0); - sigma2 = WARP_SHFL(sigma2 / float(n2), 0); - } - } -} - -template -__device__ U rsqrt(U v) { - return U(1) / sqrt(v); -} -template <> -__device__ float rsqrt(float v) { - return rsqrtf(v); -} -template <> -__device__ double rsqrt(double v) { - return rsqrt(v); -} - namespace { -// This is the un-specialized struct. Note that we prevent instantiation of this -// struct by putting an undefined symbol in the function body so it won't compile. -// template -// struct SharedMemory -// { -// // Ensure that we won't compile any un-specialized types -// __device__ T *getPointer() -// { -// extern __device__ void error(void); -// error(); -// return NULL; -// } -// }; -// https://github.com/NVIDIA/apex/issues/246 -template -struct SharedMemory; - -template <> -struct SharedMemory { - __device__ float* getPointer() { - extern __shared__ float s_float[]; - return s_float; - } -}; - -template <> -struct SharedMemory { - __device__ double* getPointer() { - extern __shared__ double s_double[]; - return s_double; - } -}; -} // namespace - -template -__global__ void cuApplyLayerNorm( - T* __restrict__ output_vals, - U* __restrict__ mean, - U* __restrict__ invvar, - const T* __restrict__ vals, - const int n1, - const int n2, - const U epsilon, - const T* __restrict__ gamma, - const T* __restrict__ beta) { - // Assumptions: - // 1) blockDim.x == GPU_WARP_SIZE - // 2) Tensors are contiguous - // - for (int i1 = blockIdx.y; i1 < n1; i1 += gridDim.y) { - SharedMemory shared; - U* buf = shared.getPointer(); - U mu, sigma2; - cuWelfordMuSigma2(vals, n1, n2, i1, mu, sigma2, buf); - const T* lvals = vals + i1 * n2; - T* ovals = output_vals + i1 * n2; - U c_invvar = rsqrt(sigma2 + epsilon); - const int numx = blockDim.x * blockDim.y; - const int thrx = threadIdx.x + threadIdx.y * blockDim.x; - if (gamma != NULL && beta != NULL) { - for (int i = thrx; i < n2; i += numx) { - U curr = static_cast(lvals[i]); - ovals[i] = gamma[i] * static_cast(c_invvar * (curr - mu)) + beta[i]; - } - } else { - for (int i = thrx; i < n2; i += numx) { - U curr = static_cast(lvals[i]); - ovals[i] = static_cast(c_invvar * (curr - mu)); - } + // This is the un-specialized struct. Note that we prevent instantiation of this + // struct by putting an undefined symbol in the function body so it won't compile. + // template + // struct SharedMemory + // { + // // Ensure that we won't compile any un-specialized types + // __device__ T *getPointer() + // { + // extern __device__ void error(void); + // error(); + // return NULL; + // } + // }; + // https://github.com/NVIDIA/apex/issues/246 + template + struct SharedMemory; + + template <> + struct SharedMemory { + __device__ float* getPointer() { + extern __shared__ float s_float[]; + return s_float; } - if (threadIdx.x == 0 && threadIdx.y == 0) { - mean[i1] = mu; - invvar[i1] = c_invvar; + }; + + template <> + struct SharedMemory { + __device__ double* getPointer() { + extern __shared__ double s_double[]; + return s_double; } - } -} + }; + } // namespace -template -void HostApplyLayerNorm( - const cudaDeviceProp& prop, - T* output, - U* mean, - U* invvar, - const T* input, - int64_t n1, - int64_t n2, - double epsilon, - const T* gamma, - const T* beta) { - const uint64_t maxGridY = prop.maxGridSize[1]; - const int warp_size = prop.warpSize; - ORT_ENFORCE(warp_size == GPU_WARP_SIZE); - - const dim3 threads(warp_size, 4, 1); - const dim3 blocks(1, std::min((uint64_t)n1, maxGridY), 1); - int nshared = - threads.y > 1 ? threads.y * sizeof(U) + (threads.y / 2) * sizeof(U) : 0; - cuApplyLayerNorm<<>>( - output, - mean, - invvar, - input, - n1, n2, - U(epsilon), - gamma, beta); -} - -#define LAYERNORM_LINEAR_IMPL(T, U) \ - template void HostApplyLayerNorm(const cudaDeviceProp& prop, T* output, U* mean, U* invvar, const T* input, int64_t n1, int64_t n2, \ - double epsilon, const T* gamma, const T* beta); - -LAYERNORM_LINEAR_IMPL(float, float) -LAYERNORM_LINEAR_IMPL(half, float) -LAYERNORM_LINEAR_IMPL(double, double) -//LAYERNORM_LINEAR_IMPL(half, half) - -template +template __device__ void cuLoadWriteStridedInputs( const int i1_block, const int thr_load_row_off, @@ -395,7 +85,7 @@ __device__ void cuLoadWriteStridedInputs( const U* __restrict__ invvar) { int i1 = i1_block + thr_load_row_off; if (i1 < i1_end) { - U curr_mean = use_mean ? mean[i1] : U(0); + U curr_mean = (use_mean && !simplified) ? mean[i1] : U(0); U curr_invvar = use_mean ? invvar[i1] : U(0); for (int k = 0; k < blockDim.y; ++k) { int i2 = i2_off + k; @@ -427,7 +117,7 @@ __device__ void cuLoadWriteStridedInputs( } } -template +template __device__ void cuLoadAddStridedInputs( const int i1_block, const int thr_load_row_off, @@ -447,7 +137,7 @@ __device__ void cuLoadAddStridedInputs( const U* __restrict__ invvar) { int i1 = i1_block + thr_load_row_off; if (i1 < i1_end) { - U curr_mean = use_mean ? mean[i1] : U(0); + U curr_mean = (use_mean && !simplified) ? mean[i1] : U(0); U curr_invvar = use_mean ? invvar[i1] : U(0); for (int k = 0; k < blockDim.y; ++k) { int i2 = i2_off + k; @@ -470,7 +160,7 @@ __device__ void cuLoadAddStridedInputs( } } -template +template __global__ void cuComputePartGradGammaBeta( const T* __restrict__ dout, const T* __restrict__ input, @@ -498,9 +188,9 @@ __global__ void cuComputePartGradGammaBeta( U* warp_buf2 = warp_buf1 + blockDim.y * blockDim.y * row_stride; // compute partial sums from strided inputs // do this to increase number of loads in flight - cuLoadWriteStridedInputs(i1_beg, thr_load_row_off, thr_load_col_off, i2_off, row_stride, warp_buf1, warp_buf2, input, output, dout, i1_end, n2, gamma, beta, mean, invvar); + cuLoadWriteStridedInputs(i1_beg, thr_load_row_off, thr_load_col_off, i2_off, row_stride, warp_buf1, warp_buf2, input, output, dout, i1_end, n2, gamma, beta, mean, invvar); for (int i1_block = i1_beg + blockDim.y * blockDim.y; i1_block < i1_end; i1_block += blockDim.y * blockDim.y) { - cuLoadAddStridedInputs(i1_block, thr_load_row_off, thr_load_col_off, i2_off, row_stride, warp_buf1, warp_buf2, input, output, dout, i1_end, n2, gamma, beta, mean, invvar); + cuLoadAddStridedInputs(i1_block, thr_load_row_off, thr_load_col_off, i2_off, row_stride, warp_buf1, warp_buf2, input, output, dout, i1_end, n2, gamma, beta, mean, invvar); } __syncthreads(); // inter-warp reductions @@ -539,7 +229,7 @@ __global__ void cuComputePartGradGammaBeta( } } -template +template __global__ void cuComputeGradGammaBeta( const U* part_grad_gamma, const U* part_grad_beta, @@ -584,12 +274,14 @@ __global__ void cuComputeGradGammaBeta( // write out fully summed gradients if (threadIdx.y == 0) { grad_gamma[i2] = sum_gamma; - grad_beta[i2] = sum_beta; + if (!simplified) { + grad_beta[i2] = sum_beta; + } } } } -template +template __global__ void cuComputeGradInput( const T* __restrict__ dout, const T* __restrict__ input, @@ -604,7 +296,7 @@ __global__ void cuComputeGradInput( for (int i1 = blockIdx.y; i1 < n1; i1 += gridDim.y) { U sum_loss1 = U(0); U sum_loss2 = U(0); - const U c_mean = use_mean ? mean[i1] : U(0); + const U c_mean = (use_mean && !simplified) ? mean[i1] : U(0); const U c_invvar = invvar[i1]; const T* k_input = use_mean ? input + i1 * n2 : nullptr; const T* k_output = use_mean ? nullptr: output + i1 * n2; @@ -700,6 +392,7 @@ __global__ void cuComputeGradInput( } } // all threads now have the two sums over l + // U sum_loss2 = X_mean_difference_over_std_var in cpu kernel U fH = (U)n2; U term1 = (U(1) / fH) * c_invvar; T* k_grad_input = grad_input + i1 * n2; @@ -707,7 +400,9 @@ __global__ void cuComputeGradInput( for (int l = thrx; l < n2; l += numx) { const U c_loss = static_cast(k_dout[l]); U f_grad_input = fH * c_loss * U(gamma[l]); - f_grad_input -= sum_loss1; + if (!simplified) { + f_grad_input -= sum_loss1; + } if (use_mean) { const U c_h = static_cast(k_input[l]); f_grad_input -= (c_h - c_mean) * c_invvar * sum_loss2; @@ -722,7 +417,9 @@ __global__ void cuComputeGradInput( for (int l = thrx; l < n2; l += numx) { const U c_loss = static_cast(k_dout[l]); U f_grad_input = fH * c_loss; - f_grad_input -= sum_loss1; + if (!simplified) { + f_grad_input -= sum_loss1; + } if (use_mean) { const U c_h = static_cast(k_input[l]); f_grad_input -= (c_h - c_mean) * c_invvar * sum_loss2; @@ -737,7 +434,7 @@ __global__ void cuComputeGradInput( } } -template +template void HostLayerNormGradient( const cudaDeviceProp& prop, const T* dout, @@ -763,9 +460,9 @@ void HostLayerNormGradient( const int nshared2_a = 2 * sizeof(U) * threads2.y * threads2.y * (threads2.x + 1); const int nshared2_b = threads2.x * threads2.y * sizeof(U); const int nshared2 = nshared2_a > nshared2_b ? nshared2_a : nshared2_b; - - if (mean == nullptr) { - cuComputePartGradGammaBeta<<>>( + if (mean == nullptr && !simplified) { + // use_mean == false, simplified == false -> Inverted Layer Norm + cuComputePartGradGammaBeta<<>>( dout, input, output, @@ -777,7 +474,9 @@ void HostLayerNormGradient( part_grad_gamma, part_grad_beta); } else { - cuComputePartGradGammaBeta<<>>( + // use_mean == true, simplified == false -> Layer Norm + // use_mean == true, simplified == true -> Simplified Layer Norm + cuComputePartGradGammaBeta<<>>( dout, input, output, @@ -789,26 +488,24 @@ void HostLayerNormGradient( part_grad_gamma, part_grad_beta); } - const dim3 threads3(warp_size, 8, 1); const dim3 blocks3((n2 + threads2.x - 1) / threads2.x, 1, 1); const int nshared3 = threads3.x * threads3.y * sizeof(U); - cuComputeGradGammaBeta<<>>( + cuComputeGradGammaBeta<<>>( part_grad_gamma, part_grad_beta, part_size, n1, n2, grad_gamma, grad_beta); - // compute grad_input const uint64_t maxGridY = prop.maxGridSize[1]; const dim3 blocks1(1, std::min((uint64_t)n1, maxGridY), 1); const dim3 threads1(warp_size, 4, 1); int nshared = threads1.y > 1 ? threads1.y * threads1.x * sizeof(U) : 0; - if (mean == nullptr) { - cuComputeGradInput<<>>( + if (mean == nullptr && !simplified) { + cuComputeGradInput<<>>( dout, input, output, @@ -819,7 +516,7 @@ void HostLayerNormGradient( n1, n2, grad_input); } else { - cuComputeGradInput<<>>( + cuComputeGradInput<<>>( dout, input, output, @@ -832,14 +529,17 @@ void HostLayerNormGradient( } } -#define LAYERNORMGRAD_IMPL(T, U) \ - template void HostLayerNormGradient(const cudaDeviceProp& prop, const T* dout, const T* input, const T* output, \ +#define LAYERNORMGRAD_IMPL(T, U, simplified) \ + template void HostLayerNormGradient(const cudaDeviceProp& prop, const T* dout, const T* input, const T* output, \ const T* gamma, const T* beta, const U* mean, const U* invvar, int64_t n1, int64_t n2, \ T* grad_input, T* grad_gamma, T* grad_beta, U* part_grad_gamma, U* part_grad_beta, const int part_size); -LAYERNORMGRAD_IMPL(float, float) -LAYERNORMGRAD_IMPL(double, double) -LAYERNORMGRAD_IMPL(half, float) +LAYERNORMGRAD_IMPL(float, float, true) +LAYERNORMGRAD_IMPL(double, double, true) +LAYERNORMGRAD_IMPL(half, float, true) +LAYERNORMGRAD_IMPL(float, float, false) +LAYERNORMGRAD_IMPL(double, double, false) +LAYERNORMGRAD_IMPL(half, float, false) } // namespace cuda } // namespace onnxruntime \ No newline at end of file diff --git a/orttraining/orttraining/training_ops/cuda/nn/layer_norm_impl.h b/orttraining/orttraining/training_ops/cuda/nn/layer_norm_impl.h index ea61ff71ab..a974382268 100644 --- a/orttraining/orttraining/training_ops/cuda/nn/layer_norm_impl.h +++ b/orttraining/orttraining/training_ops/cuda/nn/layer_norm_impl.h @@ -41,7 +41,7 @@ void HostApplyLayerNorm( const T* gamma, const T* beta); -template +template void HostLayerNormGradient( const cudaDeviceProp& prop, const T* dout,